Koda, un LLM entraîné de zéro
Un modèle de langage de 1,27 milliard de paramètres entraîné from scratch, juste pour comprendre comment ça marche à l'intérieur. Decoder-only façon LLaMA (24 couches, GQA, SwiGLU, RoPE), entraîné en JAX/Flax NNX sur 2 GPU L40S. Checkpoints publiés sur Hugging Face, exports HF, GGUF et MLX.
1,27 Mdde paramètres, entraînés de zéro pour comprendre
Le projet
KodaLite-1.3B est un modèle de langage que j'ai entraîné de zéro, non pas pour la taille, mais pour comprendre les internals : architecture decoder-only façon LLaMA (24 couches, hidden 2048, GQA 32/8, SwiGLU, RMSNorm pre-norm, RoPE), tokenizer GPT-2 BPE. Le pipeline complet tourne en JAX + Flax NNX sur 2 GPU NVIDIA L40S en bf16 : pré-entraînement sur SlimPajama (~1,6 milliard de tokens, ~25 heures) avec un orchestrateur à reprise sur crash, SFT LoRA sur Dolly et OASST, extension de contexte NTK-aware de 1024 à 2048 tokens. Les checkpoints sont publiés sur Hugging Face avec exports vers Transformers, GGUF (llama.cpp, Ollama, LM Studio) et MLX, plus un benchmark maison de 8 tâches zero-shot.
Architecture façon LLaMA
Un decoder-only de 1,27 milliard de paramètres : 24 couches, hidden 2048, GQA 32/8, SwiGLU, RMSNorm pre-norm, RoPE, tokenizer GPT-2 BPE. Le contexte natif est de 1024 tokens, étendu à 2048 en phase 3.
Pré-entraînement sur 2 L40S
Le pipeline tourne en JAX + Flax NNX en bf16 sur 2 GPU NVIDIA L40S (96 GB de VRAM). Environ 1,6 milliard de tokens SlimPajama en 25 heures, avec un orchestrateur à reprise sur crash pour ne pas perdre de progression.
Benchmark honnête
Un benchmark maison de 8 tâches zero-shot (HellaSwag, ARC, WinoGrande, PIQA, BoolQ, OpenBookQA, LAMBADA) compare KodaLite à 8 modèles d'environ 1 milliard de paramètres. Il arrive dernier, et la fiche le dit.
Loi de Chinchilla
La fiche explique pourquoi un modèle 10 fois plus gros que GPT-2-124M score en dessous : 1,64 milliard de tokens vus, soit 6,5 % de la cible Chinchilla d'environ 25 milliards. Les tokens comptent plus que les paramètres à ce budget.
Exports HF, GGUF, MLX
Les checkpoints sont publiés sur Hugging Face avec des exports vers Transformers, GGUF (llama.cpp, Ollama, LM Studio) et MLX, en fp16 et 8 bits, pour rendre le modèle utilisable en dehors de JAX.
Défis
- Entraîner un modèle de 1,27 Md de paramètres sur un budget GPU limité (2x L40S, 96 GB VRAM)
- Tenir un run de pré-entraînement d'environ 25 heures sans perdre de progression
- Étendre le contexte de 1024 à 2048 tokens après le pré-entraînement
- Rendre le modèle utilisable en dehors de JAX (Transformers, GGUF, MLX)
Solutions
- Implémentation JAX + Flax NNX en bf16 avec orchestrateur de reprise sur crash
- Pré-entraînement SlimPajama puis SFT LoRA (Dolly, OASST) et fix du token EOS
- Extension de contexte NTK-aware sans ré-entraînement complet
- Pipeline d'export vers Hugging Face Transformers, GGUF et MLX (fp16 et 8 bits)
Résultats
- Modèle KodaLite-1.3B publié sur Hugging Face (YoAbriel/KodaLite-1.3B, variantes GGUF et MLX)
- Pré-entraînement complet : ~1,6 Md de tokens SlimPajama en ~25 h sur 2x L40S
- Benchmark maison de 8 tâches zero-shot pour mesurer ce que le modèle sait vraiment faire
- Code public sur GitHub (Koda-v0.1)