@@ -43,14 +43,27 @@ def extract_user_only_text(session_turns: Sequence[Dict[str, Any]]) -> str:
4343 return "\n " .join ([line for line in lines if line ])
4444
4545
46- def format_session_memory (session_id : str , session_date : str , session_turns : Sequence [Dict [str , Any ]]) -> str :
47- """Create a memory payload that preserves session metadata in plain text."""
48- user_text = extract_user_only_text (session_turns )
46+ def format_session_memory (session_id : str , session_date : str , session_turns : Sequence [Dict [str , Any ]], include_all_roles : bool = False ) -> str :
47+ """Create a memory payload that preserves session metadata in plain text.
48+
49+ When include_all_roles=True, includes both user and assistant turns
50+ for richer context in deferred enrichment mode.
51+ """
52+ if include_all_roles :
53+ all_text = []
54+ for turn in session_turns :
55+ role = turn .get ("role" , "user" )
56+ content = str (turn .get ("content" , "" )).strip ()
57+ if content :
58+ all_text .append (f"{ role } : { content } " )
59+ full_text = "\n " .join (all_text )
60+ else :
61+ full_text = extract_user_only_text (session_turns )
4962 return (
5063 f"Session ID: { session_id } \n "
5164 f"Session Date: { session_date } \n "
5265 f"{ HISTORY_HEADER } \n "
53- f"{ user_text } "
66+ f"{ full_text } "
5467 )
5568
5669
@@ -149,8 +162,13 @@ def build_memory(
149162 llm_model : Optional [str ] = None ,
150163 embedder_model : Optional [str ] = None ,
151164 full_potential : bool = True ,
165+ defer_enrichment : bool = False ,
152166) -> Memory :
153- """Build Engram Memory for LongMemEval. By default uses full potential (echo, categories, graph, scenes, profiles)."""
167+ """Build Engram Memory for LongMemEval. By default uses full potential (echo, categories, graph, scenes, profiles).
168+
169+ When defer_enrichment=True, ingestion uses 0 LLM calls (store fast), and
170+ enrichment is done in batch after all sessions are loaded.
171+ """
154172 vector_cfg : Dict [str , Any ] = {
155173 "collection_name" : "engram_longmemeval" ,
156174 "embedding_model_dims" : embedding_dims ,
@@ -174,8 +192,12 @@ def build_memory(
174192 graph = KnowledgeGraphConfig (enable_graph = full_potential ),
175193 scene = SceneConfig (use_llm_summarization = full_potential , enable_scenes = full_potential ),
176194 profile = ProfileConfig (use_llm_extraction = full_potential , enable_profiles = full_potential ),
177- enrichment = EnrichmentConfig (enable_unified = full_potential , max_batch_size = 10 ),
178- batch = BatchConfig (enable_batch = full_potential , max_batch_size = 50 ),
195+ enrichment = EnrichmentConfig (
196+ enable_unified = full_potential ,
197+ max_batch_size = 10 ,
198+ defer_enrichment = defer_enrichment ,
199+ ),
200+ batch = BatchConfig (enable_batch = full_potential and not defer_enrichment , max_batch_size = 50 ),
179201 )
180202 mem = Memory (config )
181203 # FullMemory features (categories, scenes, profiles) need FullSQLiteManager
@@ -234,6 +256,7 @@ def run_longmemeval(args: argparse.Namespace) -> Dict[str, Any]:
234256 if args .skip_abstention :
235257 selected = [entry for entry in selected if "_abs" not in str (entry .get ("question_id" , "" ))]
236258
259+ use_deferred = getattr (args , "defer_enrichment" , False )
237260 memory = build_memory (
238261 llm_provider = args .llm_provider ,
239262 embedder_provider = args .embedder_provider ,
@@ -243,6 +266,7 @@ def run_longmemeval(args: argparse.Namespace) -> Dict[str, Any]:
243266 llm_model = args .llm_model ,
244267 embedder_model = args .embedder_model ,
245268 full_potential = args .full_potential ,
269+ defer_enrichment = use_deferred ,
246270 )
247271
248272 hf_responder : Optional [HFResponder ] = None
@@ -276,7 +300,17 @@ def run_longmemeval(args: argparse.Namespace) -> Dict[str, Any]:
276300 # Build batch items for all sessions
277301 batch_items = []
278302 for sess_id , sess_date , sess_turns in zip (session_ids , session_dates , sessions ):
279- payload = format_session_memory (str (sess_id ), str (sess_date ), sess_turns or [])
303+ payload = format_session_memory (
304+ str (sess_id ), str (sess_date ), sess_turns or [],
305+ include_all_roles = use_deferred ,
306+ )
307+ # Build context_messages from session turns for deferred mode
308+ ctx_msgs = None
309+ if use_deferred and sess_turns :
310+ ctx_msgs = [
311+ {"role" : t .get ("role" , "user" ), "content" : str (t .get ("content" , "" )).strip ()}
312+ for t in sess_turns if str (t .get ("content" , "" )).strip ()
313+ ]
280314 batch_items .append ({
281315 "content" : payload ,
282316 "metadata" : {
@@ -285,17 +319,13 @@ def run_longmemeval(args: argparse.Namespace) -> Dict[str, Any]:
285319 "question_id" : question_id ,
286320 },
287321 "categories" : ["longmemeval" , "session" ],
322+ "_context_messages" : ctx_msgs ,
288323 })
289324
290325 # Use add_batch for fewer LLM calls; fallback to sequential on failure
291326 if batch_items :
292- try :
293- memory .add_batch (
294- items = batch_items ,
295- user_id = args .user_id ,
296- )
297- except Exception as e :
298- logger .warning ("Batch add failed for question %s, retrying sequentially: %s" , question_id , e )
327+ if use_deferred :
328+ # Deferred mode: sequential add with context_messages
299329 for item in batch_items :
300330 try :
301331 memory .add (
@@ -304,9 +334,36 @@ def run_longmemeval(args: argparse.Namespace) -> Dict[str, Any]:
304334 metadata = item ["metadata" ],
305335 categories = item ["categories" ],
306336 infer = False ,
337+ context_messages = item .get ("_context_messages" ),
307338 )
308339 except Exception as e2 :
309340 logger .warning ("Skipping session for question %s: %s" , question_id , e2 )
341+ else :
342+ try :
343+ memory .add_batch (
344+ items = batch_items ,
345+ user_id = args .user_id ,
346+ )
347+ except Exception as e :
348+ logger .warning ("Batch add failed for question %s, retrying sequentially: %s" , question_id , e )
349+ for item in batch_items :
350+ try :
351+ memory .add (
352+ messages = item ["content" ],
353+ user_id = args .user_id ,
354+ metadata = item ["metadata" ],
355+ categories = item ["categories" ],
356+ infer = False ,
357+ )
358+ except Exception as e2 :
359+ logger .warning ("Skipping session for question %s: %s" , question_id , e2 )
360+
361+ # Batch enrich after all sessions loaded (deferred mode)
362+ if use_deferred :
363+ try :
364+ memory .enrich_pending (user_id = args .user_id , batch_size = 10 , max_batches = 50 )
365+ except Exception as e :
366+ logger .warning ("Enrichment failed for question %s: %s" , question_id , e )
310367
311368 query = str (entry .get ("question" , "" )).strip ()
312369 search_payload = memory .search (
@@ -436,8 +493,10 @@ def parse_args() -> argparse.Namespace:
436493 parser .add_argument ("--embedding-dims" , type = int , default = 1536 , help = "Embedding dimensions for simple/memory configs." )
437494 parser .add_argument ("--vector-store-provider" , choices = ["memory" , "sqlite_vec" ], default = "memory" )
438495 parser .add_argument ("--history-db-path" , default = "/content/engram-longmemeval.db" , help = "SQLite db path." )
496+ parser .add_argument ("--defer-enrichment" , action = "store_true" , default = False , help = "Use deferred enrichment (0 LLM calls at ingestion, batch enrich after)." )
439497 args = parser .parse_args ()
440498 args .full_potential = not args .minimal
499+ args .defer_enrichment = args .defer_enrichment
441500 return args
442501
443502
0 commit comments