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