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