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