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