File size: 1,432 Bytes
83b28fa
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
29
30
31
32
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