Download m31_runtime.py from eshanized/M31Tesla: direct link, hf CLI and curl.
- Browser
- Download file 1.43 kB
-
https://huggingface.co/eshanized/M31Tesla/resolve/main/m31_runtime.py
- Command line
-
hf download hf://eshanized/M31Tesla/m31_runtime.py
-
curl -L -o m31_runtime.py https://huggingface.co/eshanized/M31Tesla/resolve/main/m31_runtime.py
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 | |