| 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 |
|
|