memfold / load_components.py
Johnny221B's picture
Release MemFold main-result adapters and matching memory components
0ecc5ec verified
Raw History Blame Contribute Delete
2.47 kB
"""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.'}))