"""Load MemFold's learned memory components with the public implementation.""" from pathlib import Path import argparse,json,sys import torch def load_memory_components(bundle, code): bundle=Path(bundle).resolve();code=Path(code).resolve() sys.path[:0]=[str(code),str(code/'src')] config=json.loads((bundle/'bundle.json').read_text()) if 'bridge' in config: from locomo_pipeline.compressor_core import build saved=torch.load(bundle/config['bridge'],map_location='cpu',weights_only=True) bridge=build(code,saved['config']);bridge.load_state_dict(saved['bridge'],strict=True) return {'bridge':bridge.eval(),'config':config} from locomo_pipeline.session_memory_trial import build_compressor mp=torch.load(bundle/config['mapper'],map_location='cpu',weights_only=True) cp=torch.load(bundle/config['compressor'],map_location='cpu',weights_only=True) import hashlib with (bundle/config['mapper']).open('rb') as f: assert hashlib.file_digest(f,'sha256').hexdigest()==cp['mapper_hash'] h=mp['hidden'] mapper=torch.nn.Sequential(torch.nn.LayerNorm(h),torch.nn.Linear(h,1024),torch.nn.GELU(),torch.nn.Linear(1024,h)) mapper.load_state_dict(mp['state_dict'],strict=True) compressor=build_compressor(h);compressor.load_state_dict(cp['state_dict'],strict=True) return {'mapper':mapper.eval(),'compressor':compressor.eval(),'config':config} def load_reader(bundle, **model_kwargs): """Load the reader LoRA; generation still requires the learned memory prefix.""" from transformers import AutoTokenizer,AutoModelForCausalLM from peft import PeftModel bundle=Path(bundle);cfg=json.loads((bundle/'bundle.json').read_text()) tokenizer=AutoTokenizer.from_pretrained(cfg['backbone'],revision=cfg['backbone_revision']) base=AutoModelForCausalLM.from_pretrained(cfg['backbone'],revision=cfg['backbone_revision'],**model_kwargs) return PeftModel.from_pretrained(base,str(bundle/cfg['reader']),is_trainable=False).eval(),tokenizer if __name__=='__main__': p=argparse.ArgumentParser(description=__doc__);p.add_argument('--bundle',type=Path,required=True);p.add_argument('--code',type=Path,required=True);args=p.parse_args() loaded=load_memory_components(args.bundle,args.code) print(json.dumps({'loaded':list(k for k in loaded if k!='config'),'backbone':loaded['config']['backbone'],'note':'Memory component loading only; no full-model inference performed.'}))