M31Tesla / m31_runtime.py
eshanized's picture
add M31 v5 standalone inference runtime
83b28fa verified
Raw
History Blame Contribute Delete
1.43 kB
import os, json, torch
from huggingface_hub import hf_hub_download
from safetensors.torch import load_file
from tokenizers import Tokenizer
from modeling_m31 import M31Model
from generation_m31 import generate
REPO_ID = os.environ.get('M31_REPO_ID', 'eshanized/M31Tesla')
CHECKPOINTS = {
'pretrain': 'experiments/M31-Python-Agent-220M-v5/checkpoints/pretrain/step-00000733',
'sft': 'experiments/M31-Python-Agent-220M-v5/checkpoints/sft/step-00000123',
'agent': 'experiments/M31-Python-Agent-220M-v5/checkpoints/agent/step-00000123',
'repair': 'experiments/M31-Python-Agent-220M-v5/checkpoints/repair/step-00000123',
}
def load_stage(stage='repair', device=None):
if stage not in CHECKPOINTS:
raise ValueError(f'Unknown stage {stage!r}: {sorted(CHECKPOINTS)}')
device = device or ('cuda' if torch.cuda.is_available() else 'cpu')
prefix = CHECKPOINTS[stage]
config_path = hf_hub_download(REPO_ID, f'{prefix}/config.json')
weights_path = hf_hub_download(REPO_ID, f'{prefix}/model.safetensors')
tokenizer_path = hf_hub_download(REPO_ID, 'tokenizer.json')
with open(config_path, 'r', encoding='utf-8') as f:
config = json.load(f)
model = M31Model()
state = load_file(weights_path, device='cpu')
model.load_state_dict(state, strict=True)
model.to(device).eval()
tokenizer = Tokenizer.from_file(tokenizer_path)
return model, tokenizer, config