singhamAstram's picture
Upload complete MicroGPT model (hf_model_micro_gpt_tf_ckpt_step_500)
9eb0412 verified
Raw
History Blame Contribute Delete
1.02 kB
"""
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())