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