File size: 1,478 Bytes
3d11298
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
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