vtava's picture
Publish accepted Memory Fusion prefix [0]
3d11298 verified
Raw History Blame Contribute Delete
1.48 kB
from __future__ import annotations
import json
import torch
from huggingface_hub import hf_hub_download
from safetensors.torch import load_file
from transformers import AutoModelForCausalLM, AutoTokenizer
from tinycenn_lm.smollm2_memory_fusion import SmolMemoryFusionConfig, replace_attention_layers
def load_model(repo_id: str, token=None, device=None):
device=torch.device(device or ('cuda' if torch.cuda.is_available() else 'cpu'))
dtype=torch.bfloat16 if device.type=='cuda' and torch.cuda.is_bf16_supported() else (torch.float16 if device.type=='cuda' else torch.float32)
meta_path=hf_hub_download(repo_id,'tinycenn_model.json',token=token)
with open(meta_path,encoding='utf-8') as f:
meta=json.load(f)
tokenizer=AutoTokenizer.from_pretrained(repo_id,token=token,use_fast=True)
if tokenizer.pad_token_id is None:
tokenizer.pad_token=tokenizer.eos_token
model=AutoModelForCausalLM.from_pretrained(meta['base_model'],dtype=dtype)
cfg=SmolMemoryFusionConfig.from_dict(meta['memory_fusion'])
replace_attention_layers(model,cfg,meta['accepted_layers'])
state_path=hf_hub_download(repo_id,'model.safetensors',token=token)
state=load_file(state_path,device='cpu')
model.load_state_dict(state,strict=True)
model.to(device).eval(); model.requires_grad_(False); model.config.use_cache=False
if hasattr(model,'generation_config'):
model.generation_config.use_cache=False
return model,tokenizer,meta