File size: 1,024 Bytes
507f954
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""
Model registry for PyTorch model architectures.

Each model class must implement:
    __init__(self, vocab_size, block_size, n_layer, n_head, n_embd, dropout)
    forward(self, idx, targets=None) -> (logits, loss)
    generate(self, idx, max_new_tokens, temperature=1.0, top_k=40)
"""

from .micro_gpt import MicroGPT
from .micro_gpt_tf import MicroGPT_TF
from .micro_glm import MicroGLM
from .micro_deepseek import MicroDeepSeek

MODEL_REGISTRY = {
    'micro_gpt': MicroGPT,
    'micro_gpt_tf': MicroGPT_TF,
    'micro_glm': MicroGLM,
    'micro_deepseek': MicroDeepSeek,
}

def get_model(arch_name: str):
    """Return the model class for a given architecture name."""
    if arch_name not in MODEL_REGISTRY:
        available = ', '.join(sorted(MODEL_REGISTRY.keys()))
        raise ValueError(f"Unknown architecture '{arch_name}'. Available: {available}")
    return MODEL_REGISTRY[arch_name]


def list_models():
    """Return list of available model architecture names."""
    return sorted(MODEL_REGISTRY.keys())