JuloEco/opsiom-fr-tokenizer
1
1# ============================================================================2# Mini-LLM style LLaMA/Qwen — script complet prêt pour une cellule Colab3# RMSNorm + RoPE + SwiGLU + Attention SDPA (FlashAttention) + KV-Cache4# ============================================================================5 6# --- Dépendances (décommente si tu es sur Colab, sinon `pip install` en local) ---7# %pip install -q datasets tokenizers torch accelerate8 9import os10 11# ⚠️ Doit être positionné AVANT tout import de torch qui toucherait CUDA:12# suggéré directement par le message d'erreur OOM rencontré ("If reserved but13# unallocated memory is large try setting PYTORCH_ALLOC_CONF=expandable_segments:True").14# Réduit la fragmentation mémoire du cache d'allocateur CUDA de PyTorch — utile15# vu qu'on est déjà à la limite de VRAM du T4 avec Opsiom-Large.16os.environ.setdefault("PYTORCH_ALLOC_CONF", "expandable_segments:True")17 18# ⚠️ Sans ça, CHAQUE upload_file()/hf_hub_download() du système de sauvegarde19# externe (voir push_backup_to_hf / try_restore_backup_from_hf plus bas)20# affiche une barre de progression tqdm complète dans la sortie du notebook.21# Avec un push toutes les quelques dizaines de steps en tout début22# d'entraînement (val_loss qui s'améliore souvent) + un push "latest" toutes23# les 10 minutes, ça pollue rapidement la sortie de plusieurs milliers de24# lignes, ralentit l'affichage Kaggle/Colab, et gonfle inutilement la taille25# du commit final. On garde les print() explicites de push_backup_to_hf26# (concis, un par sauvegarde) comme seule trace utile.27os.environ.setdefault("HF_HUB_DISABLE_PROGRESS_BARS", "1")28import re29import json30import time31import sys32import math33import inspect34import random35import numpy as np36import torch37import torch.nn as nn38import torch.nn.functional as F39import torch.multiprocessing as mp40from accelerate import Accelerator41from dataclasses import dataclass42 43# ⚠️ Ampere (RTX A5000, compute capability 8.6) supporte le TF32 pour les44# matmuls/convs qui restent en fp32 pendant l'entraînement (ex: certaines45# réductions internes non couvertes par l'autocast bf16) — gain de vitesse46# quasi gratuit, sans perte de précision significative pour ce genre47# d'entraînement. cudnn.benchmark=True est sûr ici car block_size (donc la48# forme des tenseurs d'entrée) est constant tout au long du run: cuDNN peut49# mettre en cache le meilleur algorithme trouvé au 1er step au lieu de le50# redécouvrir. Sur un T4 (Turing, pas Ampere) le gain TF32 est nul mais sans51# risque non plus — ces lignes sont sûres sur toute génération de GPU NVIDIA.52if torch.cuda.is_available():53 torch.backends.cuda.matmul.allow_tf32 = True54 torch.backends.cudnn.allow_tf32 = True55 torch.backends.cudnn.benchmark = True56 57# ----------------------------------------------------------------------------58# Sortie non-bufferisée: dans un notebook Kaggle/Colab, stdout n'est pas un59# vrai terminal (c'est un pipe, ou même un objet `OutStream` custom d'ipykernel60# qui n'a pas de méthode `.reconfigure()`), donc Python bufferise entièrement61# au lieu d'écrire ligne par ligne. Résultat: les print() n'apparaissent qu'au62# bout d'un moment, voire seulement quand le processus se termine (ou est63# interrompu). On force ici `flush=True` sur CHAQUE print() via un monkeypatch64# — ça fonctionne quel que soit le type de flux (contrairement à65# `sys.stdout.reconfigure()`, absent sur l'OutStream de Kaggle), aussi bien66# dans le processus principal que dans chaque processus enfant spawné (ce67# bloc, hors de `if __name__ == "__main__"`, s'exécute dans les deux cas).68import builtins as _builtins69 70# ⚠️ Garde-fou anti-double-patch: si ce script s'exécute une 2e fois dans le71# MÊME kernel (ex: on relance la cellule après un crash, sans redémarrer),72# `_builtins.print` est déjà le `_flushing_print` du 1er passage. Sans ce73# garde-fou, `_original_print` capturerait cette version déjà patchée plutôt74# que le vrai print d'origine, et `_flushing_print` finirait par s'appeler75# lui-même indéfiniment (RecursionError observée en pratique) dès que le76# namespace du module est réutilisé d'une exécution à l'autre. Le marqueur77# `_is_flushing_wrapper` permet de détecter ce cas et de ne PAS re-patcher.78if not getattr(_builtins.print, "_is_flushing_wrapper", False):79 _original_print = _builtins.print80 81 def _flushing_print(*args, **kwargs):82 kwargs.setdefault("flush", True)83 _original_print(*args, **kwargs)84 85 _flushing_print._is_flushing_wrapper = True86 _builtins.print = _flushing_print87 88 89# ============================================================================90# Détection d'environnement — Colab (Drive), Kaggle, ou local91# ============================================================================92# Sur Colab, monte le Drive et fait pointer TOKENIZER_CACHE_PATH / CHECKPOINT_PATH93# / PRETRAIN_BIN_PATH dedans, pour que tokenizer, checkpoints et corpus survivent94# à la fermeture de la session. Sur Kaggle, pas d'équivalent Drive automatique:95# on utilise /kaggle/working (le dossier "Output" du kernel). ⚠️ Contrairement à96# Drive, /kaggle/working n'est PAS automatiquement persistant d'une session à97# l'autre en édition interactive — il ne survit que si tu fais "Save Version"98# (commit), et pour vraiment reprendre entre plusieurs sessions il faut ensuite99# ajouter cet Output comme Input Dataset de ta prochaine session. Hors Colab et100# hors Kaggle (exécution locale), on retombe sur le répertoire courant.101 102DRIVE_SAVE_DIR = "/content/drive/MyDrive/mini_llm_fr"103KAGGLE_SAVE_DIR = "/kaggle/working/mini_llm_fr"104 105 106# ----------------------------------------------------------------------------107# 🔑 Authentification Hugging Face sans conflit d'environnement108# ----------------------------------------------------------------------------109# ⚠️ Ce bloc est du code de niveau module, hors de `if __name__ == "__main__"`:110# avec torch.multiprocessing.spawn (méthode 'spawn'), CHAQUE processus enfant111# réimporte et réexécute intégralement ce script — donc ce bloc tournerait une112# fois par GPU en plus du processus parent. `huggingface_hub.login()` pose un113# verrou fichier (filelock) sur le cache HF pour écrire le token: plusieurs114# processus qui l'appellent en même temps peuvent entrer en contention sur ce115# verrou et rester bloqués indéfiniment (c'est très probablement ce qui a causé116# le blocage observé juste après le lancement des processus). On ne fait donc117# l'appel réseau `login()` que dans le tout premier processus (le parent);118# les processus spawnés héritent de toute façon de `HF_TOKEN` via119# l'environnement (les variables d'env sont copiées au processus enfant), donc120# il leur suffit de le relire sans refaire l'appel réseau/verrou.121_IS_SPAWNED_CHILD = mp.current_process().name != "MainProcess"122 123hf_token = os.environ.get("HF_TOKEN")124 125# Si non présent dans l'OS, on cherche dans les secrets Kaggle126if not hf_token:127 try:128 from kaggle_secrets import UserSecretsClient129 hf_token = UserSecretsClient().get_secret("HF_TOKEN")130 except Exception:131 hf_token = None132 133if hf_token:134 hf_token = hf_token.strip()135 if _IS_SPAWNED_CHILD:136 # Processus enfant (spawn): pas de login() réseau, juste s'assurer que137 # la variable d'environnement est bien présente pour ce processus.138 os.environ["HF_TOKEN"] = hf_token139 else:140 try:141 from huggingface_hub import login142 # On passe le token directement à login()143 login(token=hf_token, add_to_git_credential=False)144 # On met à jour l'environnement APRÈS la validation réussie145 os.environ["HF_TOKEN"] = hf_token146 print("🔑 Authentification Hugging Face réussie !")147 except Exception as e:148 print(f"⚠️ Échec de l'authentification : {e}")149 150def _is_kaggle() -> bool:151 return os.path.exists("/kaggle/working") or "KAGGLE_KERNEL_RUN_TYPE" in os.environ152 153 154def _detect_env_and_get_save_dir() -> str:155 """Détecte l'environnement d'exécution et renvoie le dossier de sauvegarde à156 utiliser pour le tokenizer, les checkpoints et le corpus de pré-entraînement."""157 try:158 from google.colab import drive # disponible uniquement sur Colab159 print("📎 Montage de Google Drive...")160 drive.mount("/content/drive")161 os.makedirs(DRIVE_SAVE_DIR, exist_ok=True)162 print(f"✅ Google Drive monté — sauvegardes dans {DRIVE_SAVE_DIR}")163 return DRIVE_SAVE_DIR164 except:165 pass166 167 if _is_kaggle():168 os.makedirs(KAGGLE_SAVE_DIR, exist_ok=True)169 print(f"ℹ️ Environnement Kaggle détecté — sauvegardes dans {KAGGLE_SAVE_DIR}.")170 print(" ↳ Pense à faire 'Save Version' (commit) pour conserver ce dossier, "171 "puis à le réimporter comme Input Dataset la prochaine fois pour reprendre.")172 return KAGGLE_SAVE_DIR173 174 print("ℹ️ Ni Colab ni Kaggle détecté — sauvegarde en local dans le répertoire courant.")175 return "."176 177 178def _restore_kaggle_dataset_cache(save_dir: str) -> None:179 """Si une session Kaggle précédente a été sauvegardée (bouton 'Save180 Version') puis réimportée comme Input Dataset de cette nouvelle session,181 restaure automatiquement le tokenizer / corpus de pré-entraînement /182 checkpoint déjà construits — pour éviter de refaire les ~2h30 de183 téléchargement + tokenisation à chaque nouvelle session Kaggle (rappel:184 /kaggle/working est vidé entre deux sessions, contrairement à Drive sur185 Colab). Ne copie que les fichiers absents localement: n'écrase jamais un186 fichier déjà présent dans `save_dir` (au cas où la session en cours ait187 déjà progressé plus loin que l'ancien commit importé)."""188 if not _is_kaggle():189 return190 import glob191 import shutil192 193 candidates = glob.glob("/kaggle/input/*/mini_llm_fr") + glob.glob("/kaggle/input/*/*/mini_llm_fr")194 if not candidates:195 return196 src_dir = candidates[0]197 restored = []198 for fname in ("pretrain_corpus.bin", "pretrain_corpus_meta.json", "fr_bpe_tokenizer.json", "best_model.pt"):199 src = os.path.join(src_dir, fname)200 dst = os.path.join(save_dir, fname)201 if os.path.exists(src) and not os.path.exists(dst):202 shutil.copy2(src, dst)203 restored.append(fname)204 if restored:205 print(f"♻️ Fichiers restaurés depuis le dataset Kaggle importé ({src_dir}): {', '.join(restored)}")206 else:207 print(f"ℹ️ Dataset Kaggle importé détecté ({src_dir}) mais rien à restaurer "208 f"(fichiers déjà présents localement ou absents du dataset).")209 210 211_SAVE_DIR = _detect_env_and_get_save_dir()212_restore_kaggle_dataset_cache(_SAVE_DIR)213 214# --- Multi-GPU (utile sur Kaggle: accélérateur "GPU T4 x2") ---215# WORLD_SIZE > 1 déclenche un entraînement distribué (DistributedDataParallel):216# un processus par GPU, un identique modèle sur chacun, gradients synchronisés217# à chaque step. Voir main_worker() et le lancement dans `if __name__ == "__main__"`.218# ⚠️ BUG CORRIGÉ (était: FORCE_SINGLE_GPU = True en dur, donc WORLD_SIZE=1 même219# avec 2 GPU physiquement disponibles sur Kaggle T4x2). La cause du problème220# historique (logs dupliqués, signe que chaque GPU s'entraînait indépendamment221# sans réel échange de gradients) n'était PAS le double GPU en lui-même, mais222# le fait qu'Accelerator() ne détectait pas correctement le contexte distribué223# quand les processus sont lancés via torch.multiprocessing.spawn (au lieu de224# `accelerate launch`/`notebook_launcher`). Ça a été corrigé dans main_worker()225# en initialisant explicitement le process group PyTorch226# (torch.distributed.init_process_group) AVANT de construire l'Accelerator —227# voir plus bas. Le double GPU peut donc être réactivé en toute sécurité.228FORCE_SINGLE_GPU = False229WORLD_SIZE = 1 if FORCE_SINGLE_GPU else (torch.cuda.device_count() if torch.cuda.is_available() else 1)230 231 232# ============================================================================233# 🛟 Sauvegarde externe (Hugging Face Hub) — survit à un crash/coupure Kaggle234# ============================================================================235# Le vrai problème avec /kaggle/working: il n'est PAS persistant tant qu'un236# "Save Version" (commit) n'a pas abouti. Un crash, un OOM-kill, une coupure237# réseau ou électrique AVANT ce commit final efface tout, même si238# best_model.pt vient d'être écrit sur disque une seconde plus tôt. On pousse239# donc chaque nouveau meilleur checkpoint (+ le tokenizer) vers un repo HF Hub240# PRIVÉ dès qu'il est sauvegardé localement: le fichier quitte Kaggle241# immédiatement, indépendamment de ce qui arrive à la session ensuite.242#243# Prérequis: HF_TOKEN doit avoir les droits d'écriture (un token "write", pas244# "read"), et HF_BACKUP_REPO_ID doit pointer vers un repo que ce token peut245# créer/modifier (il est créé automatiquement s'il n'existe pas encore).246HF_BACKUP_ENABLED = True247HF_BACKUP_REPO_ID = "JuloEco/opsiom-fr-checkpoints" # ⚠️ adapte à ton propre namespace HF248HF_BACKUP_PRIVATE = True249# Filet de sécurité supplémentaire: même sans nouvelle amélioration de250# val_loss (donc sans nouveau "best_model.pt"), on repousse l'état ACTUEL du251# modèle toutes les N secondes, sous un nom distinct ("latest_model.pt"). Sans252# ça, une longue période sans amélioration + un crash = retour au dernier best253# parfois vieux de plusieurs heures, alors qu'un état plus récent (même non254# "meilleur") existait juste avant le crash.255HF_BACKUP_LATEST_INTERVAL_SECONDS = 600 # 10 min256 257_hf_backup_api = None258 259 260def _get_hf_backup_api():261 """Instancie paresseusement le client HfApi (une seule fois), et262 crée le repo de sauvegarde s'il n'existe pas déjà."""263 global _hf_backup_api264 if _hf_backup_api is not None:265 return _hf_backup_api266 from huggingface_hub import HfApi267 api = HfApi()268 try:269 api.create_repo(repo_id=HF_BACKUP_REPO_ID, private=HF_BACKUP_PRIVATE, repo_type="model", exist_ok=True)270 except Exception as e:271 print(f"⚠️ Impossible de créer/vérifier le repo de sauvegarde HF '{HF_BACKUP_REPO_ID}' ({e}). "272 f"La sauvegarde externe est désactivée pour cette session.")273 return None274 _hf_backup_api = api275 return api276 277 278def push_backup_to_hf(local_path: str, path_in_repo: str) -> None:279 """Pousse un fichier vers le repo de sauvegarde HF Hub. Ne lève JAMAIS280 d'exception: un problème réseau ponctuel ne doit pas interrompre281 l'entraînement — on log juste un avertissement et on continue."""282 if not HF_BACKUP_ENABLED or not os.path.exists(local_path):283 return284 api = _get_hf_backup_api()285 if api is None:286 return287 try:288 api.upload_file(289 path_or_fileobj=local_path,290 path_in_repo=path_in_repo,291 repo_id=HF_BACKUP_REPO_ID,292 repo_type="model",293 )294 size_mb = os.path.getsize(local_path) / 1e6295 print(f" ↳ 🛟 Sauvegarde externe: '{path_in_repo}' poussé vers {HF_BACKUP_REPO_ID} ({size_mb:.1f} Mo)")296 except Exception as e:297 print(f" ↳ ⚠️ Échec de la sauvegarde externe de '{path_in_repo}' ({e}) — entraînement non interrompu.")298 299 300def try_restore_backup_from_hf(local_path: str, path_in_repo: str) -> bool:301 """Tente de restaurer un fichier depuis le repo de sauvegarde HF Hub si le302 fichier local est absent — utile après une session Kaggle perdue sans303 'Save Version', pour reprendre automatiquement sans réimport manuel.304 Retourne True si un fichier a été restauré."""305 if not HF_BACKUP_ENABLED or os.path.exists(local_path):306 return False307 try:308 from huggingface_hub import hf_hub_download309 downloaded = hf_hub_download(repo_id=HF_BACKUP_REPO_ID, filename=path_in_repo, repo_type="model")310 import shutil311 os.makedirs(os.path.dirname(local_path) or ".", exist_ok=True)312 shutil.copy2(downloaded, local_path)313 print(f"♻️ Restauré depuis la sauvegarde externe HF Hub: '{path_in_repo}' -> '{local_path}'")314 return True315 except Exception as e:316 print(f"ℹ️ Pas de sauvegarde externe utilisable pour '{path_in_repo}' ({e}).")317 return False318 319 320# ============================================================================321# Configuration322# ============================================================================323 324@dataclass325class ModelArgs:326 """Configuration hyperparamétrique pour architecture Transformer type LLaMA/Qwen."""327 vocab_size: int = 16000 # Valeur par défaut — écrasée dynamiquement par la taille328 # réelle du vocabulaire du tokenizer français entraîné plus bas329 dim: int = 1024 # Dimension des embeddings — Opsiom-Large330 n_layers: int = 16 # Nombre de blocs Transformer331 n_heads: int = 16 # Nombre de têtes d'attention (Query)332 n_kv_heads: int | None = 4 # Grouped-Query Attention (GQA): 4 têtes K/V pour 16 têtes Q333 max_seq_len: int = 512 # Fenêtre de contexte334 dropout: float = 0.1335 rope_theta: float = 10000.0336 norm_eps: float = 1e-6337 device: str = "cuda" if torch.cuda.is_available() else "cpu"338 339 340# Hyperparamètres d'entraînement — modifie-les librement341N_STORIES = 20000 # Plafond d'histoires TinyStories-French utilisées (le dataset n'en342 # contient qu'environ 1000 au total, donc en pratique tout est utilisé)343TOKENIZER_VOCAB_SIZE = 16000 # Taille cible du vocabulaire du tokenizer BPE français344TOKENIZER_CACHE_PATH = os.path.join(_SAVE_DIR, "fr_bpe_tokenizer.json")345WIKI_CONFIG = "wikitext-72" # Plus grande des deux configs d'asi/wikitext_fr (quality + good articles)346WIKI_MAX_CHARS = 25_000_000 # Plafond de caractères Wikipedia chargés (tokenizer + corpus LM)347# MAX_STEPS: 1000 steps à batch=32/block=256 ne couvre qu'~1 epoch sur le corpus348# (TinyStories-FR + Wikipedia ≈ 5-6M tokens). Pour un modèle de ~26M paramètres,349# c'est trop peu pour stabiliser les statistiques de sous-mots (cause principale350# des artefacts type "dçant", "garès êtreux"). On vise ~15-20 epochs sur les351# données disponibles plutôt qu'un budget de compute abstrait — ajustez à la352# hausse si votre val_loss continue de baisser à la fin de l'entraînement.353MAX_STEPS = 8000 # ~8-10 epochs sur le corpus combiné (était 1000)354WARMUP_STEPS = 400 # gardé à 5% de MAX_STEPS, comme avant355# ⚠️ BATCH_SIZE remonté 8 -> 24 pour la RTX A5000 (24 Go VRAM, contre 14,56 Go356# utilisables sur le T4 qui avait motivé la valeur de 8). C'est un point de357# départ raisonnable pour Opsiom-Large (196,77M params, max_seq_len=512, bf16358# natif sur Ampere) — surveille `nvidia-smi` pendant les premiers steps: s'il359# reste beaucoup de VRAM libre, remonte encore (32, 40...) par paliers ; en360# cas d'OOM, redescends. Le TF32/bf16 + le batch plus gros font l'essentiel du361# gain de throughput par rapport au réglage T4.362BATCH_SIZE = 24363# GRAD_ACCUM_STEPS abaissé en conséquence: avec un batch micro déjà bien plus364# gros (24 au lieu de 8), moins d'accumulation suffit pour un batch EFFECTIF365# confortable (24 x 2 = 48 en mono-GPU, ou 24 x n_GPU x 2 en multi-GPU) tout366# en gardant plus de vraies mises à jour de poids par unité de temps (moins de367# micro-steps "gaspillés" avant chaque step réel) — l'A5000 a largement la368# VRAM pour se permettre moins d'accumulation. Ajuste à la hausse si tu veux369# un batch effectif plus grand sans remonter BATCH_SIZE.370GRAD_ACCUM_STEPS = 4 if FORCE_SINGLE_GPU else 2371MAX_LR = 3e-4372MIN_LR = 3e-5373EVAL_INTERVAL = 200 # Évaluation + sauvegarde du meilleur modèle tous les N steps374GEN_INTERVAL = 400 # Génération d'un échantillon de contrôle tous les N steps375CHECKPOINT_PATH = os.path.join(_SAVE_DIR, "best_model.pt")376LATEST_CHECKPOINT_PATH = os.path.join(_SAVE_DIR, "latest_model.pt")377SEED = 1337378 379# Cible de tokens pour le pré-entraînement à grande échelle (modèle "Large").380# Règle Chinchilla (~20 tokens / paramètre): 196.77M params x 20 ≈ 3.9355e9.381# Ajuste ce chiffre au nombre de paramètres réel de ton modèle si tu changes382# ModelArgs (dim/n_layers/n_heads) — ce fichier ne fixe QUE la taille du corpus,383# pas l'architecture.384TARGET_PRETRAIN_TOKENS = 3_935_400_000385PRETRAIN_BIN_PATH = os.path.join(_SAVE_DIR, "pretrain_corpus.bin")386PRETRAIN_META_PATH = os.path.join(_SAVE_DIR, "pretrain_corpus_meta.json")387# Nombre de tokens réservés à la validation, plafonné: à 3,9 milliards de tokens,388# 5% donnerait ~195M tokens de val — bien plus que nécessaire pour une estimation389# stable de la val_loss, et ça gaspillerait du budget de tokens d'entraînement390# chèrement acquis (téléchargement + tokenisation).391VAL_TOKENS_CAP = 20_000_000392 393# Reprend l'entraînement depuis CHECKPOINT_PATH s'il existe, au lieu de394# repartir de poids aléatoires. Pratique pour étendre un entraînement déjà395# fait (ex: vous aviez tourné 1000 steps, vous voulez continuer) sans perdre396# ce qui a déjà été appris. Le tokenizer/vocab doit être identique.397RESUME_FROM_CHECKPOINT = True398 399 400# ============================================================================401# Normalization402# ============================================================================403 404class RMSNorm(nn.Module):405 """Root Mean Square Layer Normalization (plus rapide et stable que LayerNorm)."""406 407 def __init__(self, dim: int, eps: float = 1e-6):408 super().__init__()409 self.eps = eps410 self.weight = nn.Parameter(torch.ones(dim)) # (dim,) gain appris411 412 def _norm(self, x: torch.Tensor) -> torch.Tensor:413 return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps) # (B, T, C)414 415 def forward(self, x: torch.Tensor) -> torch.Tensor:416 # Calcul en float32 pour la stabilité numérique, recast dans le dtype d'origine417 out = self._norm(x.float()).type_as(x) # (B, T, C)418 return out * self.weight # (B, T, C) * (C,) broadcast419 420 421# ============================================================================422# Rotary Position Embeddings (RoPE)423# ============================================================================424 425class RotaryEmbedding(nn.Module):426 """Rotary Position Embeddings (RoPE) — précalcule cos/sin pour toutes les positions."""427 428 def __init__(self, dim: int, max_seq_len: int = 2048, theta: float = 10000.0):429 super().__init__()430 assert dim % 2 == 0, "head_dim doit être pair pour RoPE"431 inv_freq = 1.0 / (theta ** (torch.arange(0, dim, 2).float() / dim)) # (dim/2,)432 self.register_buffer("inv_freq", inv_freq, persistent=False)433 434 t = torch.arange(max_seq_len).float() # (max_seq_len,)435 freqs = torch.outer(t, inv_freq) # (max_seq_len, dim/2)436 emb = torch.cat((freqs, freqs), dim=-1) # (max_seq_len, dim)437 self.register_buffer("cos_cached", emb.cos(), persistent=False) # (max_seq_len, dim)438 self.register_buffer("sin_cached", emb.sin(), persistent=False) # (max_seq_len, dim)439 440 def forward(self, x: torch.Tensor, seq_len: int, start_pos: int = 0):441 # Positions [start_pos, start_pos + seq_len) — essentiel pour le décodage avec KV-Cache442 cos = self.cos_cached[start_pos:start_pos + seq_len].to(dtype=x.dtype, device=x.device) # (T, head_dim)443 sin = self.sin_cached[start_pos:start_pos + seq_len].to(dtype=x.dtype, device=x.device) # (T, head_dim)444 return cos, sin445 446 447def rotate_half(x: torch.Tensor) -> torch.Tensor:448 """Fait pivoter la moitié des dimensions: [-x2, x1] où x = [x1, x2]."""449 half = x.shape[-1] // 2450 x1 = x[..., :half]451 x2 = x[..., half:]452 return torch.cat((-x2, x1), dim=-1)453 454 455def apply_rotary_emb(xq: torch.Tensor, xk: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor):456 """Applique la rotation RoPE sur les Query et Key tensors.457 458 xq, xk: (B, n_heads, T, head_dim) ; cos, sin: (T, head_dim)459 """460 cos = cos.unsqueeze(0).unsqueeze(0) # (1, 1, T, head_dim)461 sin = sin.unsqueeze(0).unsqueeze(0) # (1, 1, T, head_dim)462 xq_rotated = (xq * cos) + (rotate_half(xq) * sin) # (B, n_heads, T, head_dim)463 xk_rotated = (xk * cos) + (rotate_half(xk) * sin) # (B, n_kv_heads, T, head_dim)464 return xq_rotated, xk_rotated465 466 467# ============================================================================468# SwiGLU FeedForward469# ============================================================================470 471class SwiGLUFeedForward(nn.Module):472 """Couche MLP SwiGLU (Gated Linear Unit avec SiLU) utilisée dans LLaMA/Qwen."""473 474 def __init__(self, args: ModelArgs):475 super().__init__()476 hidden_dim = int(8 * args.dim / 3)477 multiple_of = 256478 hidden_dim = multiple_of * ((hidden_dim + multiple_of - 1) // multiple_of)479 480 self.w1 = nn.Linear(args.dim, hidden_dim, bias=False) # gate481 self.w2 = nn.Linear(hidden_dim, args.dim, bias=False) # down482 self.w3 = nn.Linear(args.dim, hidden_dim, bias=False) # up483 self.dropout = nn.Dropout(args.dropout)484 485 def forward(self, x: torch.Tensor) -> torch.Tensor:486 gate = F.silu(self.w1(x)) # (B, T, hidden_dim)487 up = self.w3(x) # (B, T, hidden_dim)488 out = self.w2(gate * up) # (B, T, dim)489 return self.dropout(out)490 491 492# ============================================================================493# Attention (SDPA / FlashAttention) avec support KV-Cache494# ============================================================================495 496class ModernCausalAttention(nn.Module):497 """Attention causale moderne avec PyTorch SDPA (FlashAttention / Scaled Dot-Product)."""498 499 def __init__(self, args: ModelArgs):500 super().__init__()501 self.n_heads = args.n_heads502 self.n_kv_heads = args.n_kv_heads if args.n_kv_heads is not None else args.n_heads503 assert args.n_heads % self.n_kv_heads == 0, "n_heads doit être divisible par n_kv_heads (GQA)"504 self.n_rep = self.n_heads // self.n_kv_heads505 self.head_dim = args.dim // args.n_heads506 self.dropout_p = args.dropout507 508 self.wq = nn.Linear(args.dim, self.n_heads * self.head_dim, bias=False)509 self.wk = nn.Linear(args.dim, self.n_kv_heads * self.head_dim, bias=False)510 self.wv = nn.Linear(args.dim, self.n_kv_heads * self.head_dim, bias=False)511 self.wo = nn.Linear(self.n_heads * self.head_dim, args.dim, bias=False)512 self.resid_dropout = nn.Dropout(args.dropout)513 514 def forward(self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, kv_cache=None):515 B, T, C = x.shape516 517 q = self.wq(x).view(B, T, self.n_heads, self.head_dim).transpose(1, 2) # (B, n_heads, T, head_dim)518 k = self.wk(x).view(B, T, self.n_kv_heads, self.head_dim).transpose(1, 2) # (B, n_kv_heads, T, head_dim)519 v = self.wv(x).view(B, T, self.n_kv_heads, self.head_dim).transpose(1, 2) # (B, n_kv_heads, T, head_dim)520 521 q, k = apply_rotary_emb(q, k, cos, sin)522 523 if kv_cache is not None:524 past_k, past_v = kv_cache525 if past_k is not None:526 k = torch.cat((past_k, k), dim=2) # (B, n_kv_heads, T_past+T, head_dim)527 v = torch.cat((past_v, v), dim=2)528 new_kv_cache = (k, v)529 else:530 new_kv_cache = None531 532 if self.n_rep > 1:533 k = k.repeat_interleave(self.n_rep, dim=1) # (B, n_heads, T_kv, head_dim)534 v = v.repeat_interleave(self.n_rep, dim=1)535 536 is_causal = kv_cache is None or k.shape[2] == q.shape[2]537 y = F.scaled_dot_product_attention(538 q, k, v,539 attn_mask=None,540 dropout_p=self.dropout_p if self.training else 0.0,541 is_causal=is_causal,542 ) # (B, n_heads, T, head_dim)543 544 y = y.transpose(1, 2).contiguous().view(B, T, self.n_heads * self.head_dim) # (B, T, dim)545 y = self.resid_dropout(self.wo(y))546 547 if kv_cache is not None:548 return y, new_kv_cache549 return y550 551 552class TransformerBlock(nn.Module):553 def __init__(self, args: ModelArgs):554 super().__init__()555 self.attn = ModernCausalAttention(args)556 self.ffn = SwiGLUFeedForward(args)557 self.norm1 = RMSNorm(args.dim, eps=args.norm_eps)558 self.norm2 = RMSNorm(args.dim, eps=args.norm_eps)559 560 def forward(self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, kv_cache=None):561 if kv_cache is not None:562 attn_out, new_kv_cache = self.attn(self.norm1(x), cos, sin, kv_cache=kv_cache)563 x = x + attn_out564 x = x + self.ffn(self.norm2(x))565 return x, new_kv_cache566 else:567 x = x + self.attn(self.norm1(x), cos, sin)568 x = x + self.ffn(self.norm2(x))569 return x570 571 572# ============================================================================573# Modèle complet574# ============================================================================575 576class ModernLLM(nn.Module):577 """Architecture complète GPT/LLaMA auto-régressive."""578 579 def __init__(self, args: ModelArgs):580 super().__init__()581 self.args = args582 583 self.tok_embeddings = nn.Embedding(args.vocab_size, args.dim)584 self.dropout = nn.Dropout(args.dropout)585 586 head_dim = args.dim // args.n_heads587 self.rope = RotaryEmbedding(head_dim, max_seq_len=args.max_seq_len, theta=args.rope_theta)588 589 self.layers = nn.ModuleList([TransformerBlock(args) for _ in range(args.n_layers)])590 self.norm_f = RMSNorm(args.dim, eps=args.norm_eps)591 592 self.lm_head = nn.Linear(args.dim, args.vocab_size, bias=False)593 self.tok_embeddings.weight = self.lm_head.weight # Weight Tying594 595 self.apply(self._init_weights)596 for name, p in self.named_parameters():597 if name.endswith("w2.weight") or name.endswith("wo.weight"):598 nn.init.normal_(p, mean=0.0, std=0.02 / math.sqrt(2 * args.n_layers))599 600 def _init_weights(self, module: nn.Module):601 if isinstance(module, nn.Linear):602 nn.init.normal_(module.weight, mean=0.0, std=0.02)603 if module.bias is not None:604 nn.init.zeros_(module.bias)605 elif isinstance(module, nn.Embedding):606 nn.init.normal_(module.weight, mean=0.0, std=0.02)607 608 def num_params(self) -> int:609 # Les poids liés (tok_embeddings == lm_head) ne sont comptés qu'une fois:610 # on déduplique par id() du tenseur avant de sommer.611 unique_params = {id(p): p for p in self.parameters()}612 return sum(p.numel() for p in unique_params.values())613 614 def forward(615 self,616 tokens: torch.Tensor,617 targets: torch.Tensor | None = None,618 kv_caches: list | None = None,619 start_pos: int = 0,620 ):621 B, T = tokens.shape622 assert start_pos + T <= self.args.max_seq_len, (623 f"Position {start_pos + T} > max_seq_len {self.args.max_seq_len}"624 )625 626 x = self.tok_embeddings(tokens) # (B, T, dim)627 x = self.dropout(x)628 629 cos, sin = self.rope(x, seq_len=T, start_pos=start_pos)630 631 if kv_caches is not None:632 new_kv_caches = []633 for i, layer in enumerate(self.layers):634 x, layer_cache = layer(x, cos, sin, kv_cache=kv_caches[i])635 new_kv_caches.append(layer_cache)636 else:637 for layer in self.layers:638 x = layer(x, cos, sin)639 new_kv_caches = None640 641 x = self.norm_f(x)642 logits = self.lm_head(x) # (B, T, vocab_size)643 644 loss = None645 if targets is not None:646 loss = F.cross_entropy(647 logits.view(-1, logits.size(-1)),648 targets.view(-1),649 ignore_index=-1,650 )651 652 if kv_caches is not None:653 return logits, loss, new_kv_caches654 return logits, loss655 656 def _sample_next_token(657 self,658 logits: torch.Tensor,659 temperature: float,660 top_k: int | None,661 top_p: float | None,662 generated_ids: torch.Tensor | None = None,663 repetition_penalty: float = 1.0,664 ) -> torch.Tensor:665 """Échantillonne le prochain token. logits: (1, vocab_size)."""666 667 # --- Pénalité de répétition (style HuggingFace) ---668 # Pour chaque token déjà généré, on divise son logit par la pénalité s'il est669 # positif (on le rend moins probable) ou on le multiplie s'il est négatif670 # (même effet: on pousse le score vers -inf). Casse les boucles de mots répétés.671 if repetition_penalty != 1.0 and generated_ids is not None and generated_ids.numel() > 0:672 unique_ids = torch.unique(generated_ids)673 prev_logits = logits[0, unique_ids] # (n_unique,)674 penalized = torch.where(675 prev_logits > 0,676 prev_logits / repetition_penalty,677 prev_logits * repetition_penalty,678 )679 logits[0, unique_ids] = penalized680 681 if temperature <= 0.0:682 return torch.argmax(logits, dim=-1, keepdim=True) # (1, 1)683 684 logits = logits / temperature685 686 if top_k is not None:687 top_k_clamped = min(top_k, logits.size(-1))688 v, _ = torch.topk(logits, top_k_clamped)689 threshold = v[:, [-1]]690 logits = torch.where(logits < threshold, torch.full_like(logits, float("-inf")), logits)691 692 probs = F.softmax(logits, dim=-1)693 694 if top_p is not None and top_p < 1.0:695 sorted_probs, sorted_indices = torch.sort(probs, descending=True, dim=-1)696 cumulative_probs = torch.cumsum(sorted_probs, dim=-1)697 sorted_mask = cumulative_probs - sorted_probs > top_p698 sorted_probs[sorted_mask] = 0.0699 sorted_probs = sorted_probs / sorted_probs.sum(dim=-1, keepdim=True)700 probs = torch.zeros_like(probs).scatter_(-1, sorted_indices, sorted_probs)701 702 return torch.multinomial(probs, num_samples=1) # (1, 1)703 704 @torch.no_grad()705 def generate(706 self,707 prompt: str,708 tokenizer: "FrenchTokenizerWrapper",709 max_new_tokens: int = 50,710 temperature: float = 0.8,711 top_p: float = 0.9,712 top_k: int | None = 40,713 repetition_penalty: float = 1.3,714 ):715 """Génération auto-régressive avec KV-Cache, temperature, top-k, top-p et716 pénalité de répétition. Le prompt est traité en une seule passe (prefill),717 puis chaque nouveau token n'est calculé qu'une seule fois (decode step)."""718 self.eval()719 device = next(self.parameters()).device720 721 token_ids = tokenizer.encode(prompt, allowed_special="all")722 tokens = torch.tensor([token_ids], dtype=torch.long, device=device) # (1, T0)723 724 max_prompt_len = self.args.max_seq_len - 1725 tokens = tokens[:, -max_prompt_len:]726 727 kv_caches = [(None, None) for _ in range(self.args.n_layers)]728 729 logits, _, kv_caches = self.forward(tokens, kv_caches=kv_caches, start_pos=0)730 logits = logits[:, -1, :]731 cur_pos = tokens.shape[1]732 733 for _ in range(max_new_tokens):734 next_token = self._sample_next_token(735 logits, temperature, top_k, top_p,736 generated_ids=tokens, repetition_penalty=repetition_penalty,737 )738 tokens = torch.cat((tokens, next_token), dim=1)739 740 if next_token.item() == tokenizer.eot_token:741 break742 if cur_pos >= self.args.max_seq_len:743 break744 745 logits, _, kv_caches = self.forward(next_token, kv_caches=kv_caches, start_pos=cur_pos)746 logits = logits[:, -1, :]747 cur_pos += 1748 749 self.train()750 return tokenizer.decode(tokens[0].tolist())751 752 753# ============================================================================754# Dataset755# ============================================================================756 757class TextDataset(torch.utils.data.Dataset):758 """Découpe un long corpus tokenisé en fenêtres (input, target) décalées d'un token."""759 760 def __init__(self, token_ids: list[int], block_size: int):761 self.data = torch.tensor(token_ids, dtype=torch.long)762 self.block_size = block_size763 764 def __len__(self):765 return max(0, len(self.data) - self.block_size)766 767 def __getitem__(self, idx: int):768 x = self.data[idx: idx + self.block_size]769 y = self.data[idx + 1: idx + 1 + self.block_size]770 return x, y771 772 773class TextDatasetMemmap(torch.utils.data.Dataset):774 """Équivalent de TextDataset, mais lit les tokens depuis un fichier binaire sur775 disque via np.memmap au lieu de tout garder en RAM. Indispensable au-delà de776 quelques centaines de millions de tokens: un tenseur de 3,9 milliards de777 tokens en int64 pèserait à lui seul ~31 Go, hors de portée d'une VM Colab."""778 779 def __init__(self, bin_path: str, dtype, start: int, end: int, block_size: int):780 self.bin_path = bin_path781 self.dtype = dtype782 self.start = start783 self.end = end # exclusif784 self.block_size = block_size785 self._mm = None # ouvert paresseusement, voir _ensure_mmap786 787 def _ensure_mmap(self):788 # np.memmap ne se transmet pas proprement aux processus worker d'un789 # DataLoader multi-process — on l'ouvre à la demande dans chaque worker790 # (le premier __getitem__ appelé dans ce processus l'initialise).791 if self._mm is None:792 self._mm = np.memmap(self.bin_path, dtype=self.dtype, mode="r")793 794 def __len__(self):795 return max(0, (self.end - self.start) - self.block_size)796 797 def __getitem__(self, idx: int):798 self._ensure_mmap()799 i = self.start + idx800 x = torch.from_numpy(self._mm[i: i + self.block_size].astype(np.int64))801 y = torch.from_numpy(self._mm[i + 1: i + 1 + self.block_size].astype(np.int64))802 return x, y803 804 805class RandomWindowIterableDataset(torch.utils.data.IterableDataset):806 """Dataset infini pour un corpus memmap gigantesque (potentiellement des807 milliards de fenêtres): pioche des positions de départ aléatoires une par808 une (`random.randint`), au lieu de matérialiser une permutation complète809 des indices comme le font `shuffle=True` / `DistributedSampler` par défaut.810 811 ⚠️ `DistributedSampler(..., shuffle=True)` et le `shuffle=True` standard812 d'un DataLoader appellent en interne `torch.randperm(len(dataset))` pour813 mélanger les indices. Sur un dataset de ~3,9 milliards de fenêtres, ça814 alloue un tenseur de ~31 Go rien que pour les indices (int64), PUIS le815 convertit en liste Python de 3,9 milliards d'objets `int` (des dizaines de816 Go supplémentaires, chaque int Python pesant bien plus que 8 octets). Sur817 Kaggle, cette tentative d'allocation dépasse la RAM système disponible et818 le noyau Linux tue le processus (SIGKILL) — c'est précisément ce qui s'est819 produit au tout premier step d'entraînement.820 821 Un tirage aléatoire AVEC remise (bootstrap), fenêtre par fenêtre, est822 statistiquement équivalent à un vrai mélange pour du pré-entraînement à823 cette échelle (on ne complète de toute façon jamais une epoch entière sur824 un corpus de cette taille en quelques sessions Kaggle), et ne matérialise825 jamais plus d'un seul index à la fois."""826 827 def __init__(self, bin_path: str, dtype, start: int, end: int, block_size: int, seed: int = 0):828 super().__init__()829 self.bin_path = bin_path830 self.dtype = dtype831 self.start = start832 self.end = end # exclusif833 self.block_size = block_size834 self.seed = seed835 self._mm = None836 837 def _ensure_mmap(self):838 if self._mm is None:839 self._mm = np.memmap(self.bin_path, dtype=self.dtype, mode="r")840 841 def __iter__(self):842 self._ensure_mmap()843 # Graine différente par worker DataLoader (si num_workers > 0 un jour)844 # pour ne pas tirer exactement la même séquence dans chaque worker.845 worker_info = torch.utils.data.get_worker_info()846 worker_id = worker_info.id if worker_info is not None else 0847 rng = random.Random(self.seed + worker_id)848 hi = self.end - self.block_size - 1849 while True:850 i = rng.randint(self.start, hi)851 x = torch.from_numpy(self._mm[i: i + self.block_size].astype(np.int64))852 y = torch.from_numpy(self._mm[i + 1: i + 1 + self.block_size].astype(np.int64))853 yield x, y854 855 856# ============================================================================857# Tokenizer BPE français — vocabulaire entraîné sur asi/wikitext_fr (Hugging Face)858# ============================================================================859 860class FrenchTokenizerWrapper:861 """Adapte un `tokenizers.Tokenizer` (BPE byte-level) à l'interface utilisée862 dans le reste du script: `.encode()`, `.decode()`, `.eot_token`."""863 864 def __init__(self, tokenizer):865 self._tok = tokenizer866 eot_id = tokenizer.token_to_id("<|endoftext|>")867 assert eot_id is not None, "Le tokenizer doit contenir le token spécial <|endoftext|>"868 self.eot_token = eot_id869 self.vocab_size = tokenizer.get_vocab_size()870 871 def encode(self, text: str, allowed_special: str = "all") -> list[int]:872 return self._tok.encode(text).ids873 874 def decode(self, ids: list[int]) -> str:875 return self._tok.decode(ids, skip_special_tokens=False)876 877 878def _import_datasets_module():879 try:880 return __import__("datasets")881 except ImportError as e:882 raise ImportError(883 "Le package 'datasets' n'est pas installé. Installez-le avec `pip install datasets`."884 ) from e885 886 887# ============================================================================888# Filtre anti-LaTeX (corpus Wikipedia)889# ============================================================================890_LATEX_COMMAND_RE = re.compile(r"\\(?:[a-zA-Z]+|[^a-zA-Z\s])") # \frac, \alpha, \{, \\, ...891_LATEX_SCRIPT_RE = re.compile(r"[_^]\{[^{}]{0,80}\}") # T_{ij}, x^{2}892_LATEX_INLINE_MATH_RE = re.compile(r"\${1,2}[^$\n]{1,200}\${1,2}") # $...$ ou $$...$$893_LATEX_ENV_RE = re.compile(r"\\(?:begin|end)\{[a-zA-Z*]+\}")894 895LATEX_DROP_THRESHOLD = 0.08896 897 898def _latex_pollution_ratio(text: str) -> float:899 if not text:900 return 0.0901 matches = (902 len(_LATEX_COMMAND_RE.findall(text))903 + len(_LATEX_SCRIPT_RE.findall(text))904 + len(_LATEX_INLINE_MATH_RE.findall(text))905 + len(_LATEX_ENV_RE.findall(text))906 )907 approx_chars = matches * 6908 return approx_chars / max(1, len(text))909 910 911def _strip_latex_noise(text: str) -> str:912 text = _LATEX_ENV_RE.sub(" ", text)913 text = _LATEX_INLINE_MATH_RE.sub(" ", text)914 text = _LATEX_SCRIPT_RE.sub(" ", text)915 text = _LATEX_COMMAND_RE.sub(" ", text)916 return re.sub(r"\s{2,}", " ", text).strip()917 918 919def filter_latex_pollution(paragraphs: list[str]) -> list[str]:920 cleaned = []921 dropped = 0922 for p in paragraphs:923 ratio = _latex_pollution_ratio(p)924 if ratio > LATEX_DROP_THRESHOLD:925 dropped += 1926 continue927 if ratio > 0:928 p = _strip_latex_noise(p)929 if len(p) < 20:930 dropped += 1931 continue932 cleaned.append(p)933 if dropped:934 print(f"🧹 Filtre anti-LaTeX: {dropped:,} paragraphe(s) pollué(s) retiré(s)/nettoyé(s) "935 f"sur {len(paragraphs):,} ({dropped / max(1, len(paragraphs)):.1%}).")936 return cleaned937 938 939def _load_wikitext_fr_via_datasets(max_chars: int) -> list[str]:940 datasets = _import_datasets_module()941 print(f"📥 Téléchargement de asi/wikitext_fr (config '{WIKI_CONFIG}') via load_dataset...")942 ds = datasets.load_dataset("asi/wikitext_fr", WIKI_CONFIG, split="train")943 paragraphs, total_chars = [], 0944 for row in ds:945 p = row["paragraph"].strip()946 if not p:947 continue948 paragraphs.append(p)949 total_chars += len(p)950 if total_chars >= max_chars:951 break952 return paragraphs953 954 955def _load_wikitext_fr_via_zip(max_chars: int) -> list[str]:956 from huggingface_hub import hf_hub_download957 import zipfile958 959 folder = "wikitext_72" if WIKI_CONFIG == "wikitext-72" else "wikitext_35"960 print(f"📥 Téléchargement direct de {folder}/wiki.zip depuis asi/wikitext_fr...")961 zip_path = hf_hub_download(repo_id="asi/wikitext_fr", repo_type="dataset", filename=f"{folder}/wiki.zip")962 extract_dir = zip_path + "_extracted"963 with zipfile.ZipFile(zip_path, "r") as zf:964 zf.extractall(extract_dir)965 966 train_file = None967 for root, _, files in os.walk(extract_dir):968 for fname in files:969 if "train" in fname.lower():970 train_file = os.path.join(root, fname)971 break972 if train_file is None:973 raise FileNotFoundError("Fichier d'entraînement introuvable dans l'archive extraite.")974 975 with open(train_file, "r", encoding="utf-8", errors="ignore") as f:976 raw = f.read(max_chars)977 return [p.strip() for p in raw.split("\n") if len(p.strip()) > 20]978 979 980def _load_wikimedia_wikipedia_fr(max_chars: int) -> list[str]:981 try:982 from datasets import load_dataset983 except Exception: # pragma: no cover - graceful fallback when `datasets` is not installed984 load_dataset = None985 import warnings986 987 warnings.warn(988 "Optional dependency 'datasets' is not available. Functions that rely on it will raise an error if used.\n"989 "Install it with: pip install datasets"990 )991 print("📥 Téléchargement (streaming) de wikimedia/wikipedia (fr) en repli...")992 ds = load_dataset("wikimedia/wikipedia", "20231101.fr", split="train", streaming=True)993 paragraphs, total_chars = [], 0994 for row in ds:995 text = row["text"].strip()996 for p in text.split("\n\n"):997 p = p.strip()998 if len(p) < 50: # ignore titres/fragments trop courts999 continue1000 paragraphs.append(p)1001 total_chars += len(p)1002 if total_chars >= max_chars:1003 break1004 return paragraphs1005 1006 1007def load_wikipedia_paragraphs(max_chars: int = WIKI_MAX_CHARS) -> list[str]:1008 # ⚠️ asi/wikitext_fr désactivé: son chargement via `load_dataset` échoue1009 # systématiquement (script de chargement obsolète), ET son repli par1010 # téléchargement/extraction de zip a provoqué un SIGKILL/SIGTERM externe1011 # (probable OOM) en tournant en parallèle de 2 process GPU. On va1012 # directement sur wikimedia/wikipedia, fiable dans toutes nos runs1013 # précédentes et streamé (pas d'extraction de zip en mémoire).1014 attempts = [1015 (lambda: _load_wikimedia_wikipedia_fr(max_chars), "wikimedia/wikipedia (fr)"),1016 ]1017 for loader, label in attempts:1018 try:1019 paragraphs = loader()1020 print(f"✅ {len(paragraphs):,} paragraphes chargés depuis {label}.")1021 paragraphs = filter_latex_pollution(paragraphs)1022 return paragraphs1023 except Exception as e:1024 print(f"⚠️ Échec avec {label} ({e}).")1025 print("↪️ Tous les téléchargements ont échoué — utilisation du texte de secours en français embarqué.")1026 return [FALLBACK_TEXT_FR]1027 1028 1029def build_or_load_french_tokenizer(1030 paragraphs: list[str],1031 vocab_size: int = TOKENIZER_VOCAB_SIZE,1032 cache_path: str = TOKENIZER_CACHE_PATH,1033) -> FrenchTokenizerWrapper:1034 """Entraîne un tokenizer BPE byte-level (façon GPT-2, donc pas d'OOV possible)1035 sur des paragraphes Wikipedia en français — vocabulaire authentiquement1036 français, beaucoup plus compact que le vocab anglais de gpt2 (50257 tokens).1037 Si un tokenizer entraîné est déjà en cache sur disque, il est rechargé direct.1038 Avant ça, tente de le restaurer depuis la sauvegarde externe HF Hub si absent1039 localement (ex: nouvelle session Kaggle sans 'Save Version' précédent)."""1040 from tokenizers import Tokenizer1041 1042 try_restore_backup_from_hf(cache_path, "fr_bpe_tokenizer.json")1043 1044 if os.path.exists(cache_path):1045 print(f"📂 Tokenizer français rechargé depuis {cache_path}.")1046 tok = Tokenizer.from_file(cache_path)1047 return FrenchTokenizerWrapper(tok)1048 1049 from tokenizers.models import BPE1050 from tokenizers.trainers import BpeTrainer1051 from tokenizers.pre_tokenizers import ByteLevel as ByteLevelPreTokenizer1052 from tokenizers.decoders import ByteLevel as ByteLevelDecoder1053 1054 print(f"🛠️ Entraînement d'un tokenizer BPE français (vocab_size={vocab_size}) "1055 f"sur {len(paragraphs):,} paragraphes...")1056 tokenizer = Tokenizer(BPE(unk_token=None))1057 tokenizer.pre_tokenizer = ByteLevelPreTokenizer(add_prefix_space=False)1058 tokenizer.decoder = ByteLevelDecoder()1059 trainer = BpeTrainer(vocab_size=vocab_size, min_frequency=2, special_tokens=["<|endoftext|>"])1060 tokenizer.train_from_iterator(paragraphs, trainer=trainer)1061 # Enregistre <|endoftext|> comme token spécial "atomique": sans ça, encode()1062 # découperait la chaîne littérale en octets au lieu de la reconnaître d'un bloc.1063 tokenizer.add_special_tokens(["<|endoftext|>"])1064 tokenizer.save(cache_path)1065 print(f"✅ Tokenizer entraîné (vocabulaire réel: {tokenizer.get_vocab_size()} tokens) "1066 f"et sauvegardé dans {cache_path}.")1067 # Sauvegarde externe immédiate: le tokenizer ne change plus jamais ensuite,1068 # donc un seul push suffit (pas besoin de le répéter à chaque restauration).1069 push_backup_to_hf(cache_path, "fr_bpe_tokenizer.json")1070 return FrenchTokenizerWrapper(tokenizer)1071 1072 1073# ============================================================================1074# Chargement du dataset — TinyStories-French, avec repli sur un texte français1075# embarqué si le téléchargement échoue (pas de réseau, dataset gated, etc.)1076# ============================================================================1077 1078FALLBACK_TEXT_FR = """1079Il était une fois un petit renard curieux qui vivait à l'orée d'une forêt tranquille.1080Chaque matin, il sortait de son terrier pour explorer les sentiers couverts de mousse.1081Un jour, il rencontra une chouette sage perchée sur une branche basse.1082La chouette lui dit: si tu veux comprendre la forêt, il faut d'abord apprendre à écouter le silence.1083Le petit renard s'assit et ferma les yeux. Il entendit le vent dans les feuilles, le ruisseau au loin, et les oiseaux qui chantaient.1084Depuis ce jour, il revenait souvent voir la chouette pour apprendre de nouvelles histoires.1085Un lapin nommé Noisette vivait aussi dans cette forêt. Il aimait collectionner de petits cailloux ronds.1086Chaque cailloux avait une couleur différente, et Noisette les rangeait soigneusement dans un panier tressé.1087Un matin de printemps, la rivière déborda légèrement à cause de la fonte des neiges.1088Le renard et le lapin décidèrent de construire un petit pont avec des branches pour aider leurs amis à traverser.1089Ensemble, ils travaillèrent toute la journée, portant des bâtons et les attachant avec des lianes solides.1090Quand le pont fut terminé, tous les animaux de la forêt vinrent le remercier avec des fleurs et des fruits.1091La chouette sage regarda la scène depuis son arbre et sourit: la coopération est la plus belle des forces.1092Le soir venu, les étoiles apparurent une à une dans le ciel violet, et la forêt s'endormit doucement.1093Le lendemain, le petit renard raconta cette aventure à tous ses amis, encore et encore, avec des étoiles dans les yeux.1094""" * 40 # répété pour donner un corpus d'entraînement de taille suffisante1095 1096 1097def _fix_mojibake(text: str) -> str:1098 if "é" in text or "è" in text or "’" in text:1099 try:1100 return text.encode("latin1").decode("utf8")1101 except (UnicodeEncodeError, UnicodeDecodeError):1102 return text1103 return text1104 1105 1106def load_training_text(n_stories: int = N_STORIES) -> str:1107 try:1108 from datasets import load_dataset1109 print(f"📥 Téléchargement de TinyStories-French...")1110 ds = load_dataset("iproskurina/TinyStories-French", split="train")1111 column = "french-tinystories" if "french-tinystories" in ds.column_names else ds.column_names[0]1112 texts = [t for t in ds[column] if t and t.strip()]1113 texts = texts[:n_stories] if n_stories < len(texts) else texts1114 texts = [_fix_mojibake(t) for t in texts]1115 dataset_text = "\n<|endoftext|>\n".join(texts)1116 print(f"✅ Dataset chargé: {len(texts)} histoires en français, {len(dataset_text):,} caractères.")1117 if len(texts) < 1500:1118 print("ℹ️ Corpus restreint (~1000 histoires) — le modèle reverra plusieurs fois "1119 "les mêmes textes sur 1000 steps, ce qui reste adapté à un modèle de cette taille.")1120 return dataset_text1121 except Exception as e:1122 print(f"⚠️ Impossible de charger TinyStories-French ({e}).")1123 print("↪️ Utilisation du texte de secours en français (corpus embarqué).")1124 return FALLBACK_TEXT_FR1125 1126 1127# ============================================================================1128# Sources de pré-entraînement à grande échelle (plusieurs milliards de tokens)1129# ============================================================================1130 1131def _stream_wikipedia_fr_docs():1132 from datasets import load_dataset1133 ds = load_dataset("wikimedia/wikipedia", "20231101.fr", split="train", streaming=True)1134 for row in ds:1135 text = row["text"].strip()1136 if text:1137 yield text1138 1139 1140def _stream_fineweb2_fr_docs():1141 from datasets import load_dataset1142 ds = load_dataset("HuggingFaceFW/fineweb-2", "fra_Latn", split="train", streaming=True)1143 for row in ds:1144 text = row["text"].strip()1145 if text:1146 yield text1147 1148 1149def _stream_oscar_fr_docs():1150 token = os.environ.get("HF_TOKEN")1151 if not token:1152 raise RuntimeError(1153 "HF_TOKEN absent — OSCAR-2301 nécessite d'accepter les conditions d'accès "1154 "sur huggingface.co/datasets/oscar-corpus/OSCAR-2301 puis d'exporter un token."1155 )1156 from datasets import load_dataset1157 ds = load_dataset("oscar-corpus/OSCAR-2301", language="fr", split="train", streaming=True, token=token)1158 for row in ds:1159 text = row["text"].strip()1160 if text:1161 yield text1162 1163 1164PRETRAIN_SOURCES = [1165 {"name": "Wikipedia FR (wikimedia/wikipedia)", "weight": 0.15, "make_gen": _stream_wikipedia_fr_docs},1166 {"name": "FineWeb-2 FR (HuggingFaceFW/fineweb-2, fra_Latn)", "weight": 0.70, "make_gen": _stream_fineweb2_fr_docs},1167 {"name": "OSCAR-2301 FR (opportuniste, requiert HF_TOKEN)", "weight": 0.15, "make_gen": _stream_oscar_fr_docs},1168]1169 1170 1171def build_pretraining_corpus_bin(1172 tokenizer: FrenchTokenizerWrapper,1173 target_tokens: int = TARGET_PRETRAIN_TOKENS,1174 bin_path: str = PRETRAIN_BIN_PATH,1175 meta_path: str = PRETRAIN_META_PATH,1176 extra_text_once: str = "",1177) -> tuple[str, int]:1178 dtype = np.uint16 if tokenizer.vocab_size <= 65535 else np.uint321179 print("cwd:", os.getcwd())1180 print("bin existe:", os.path.exists(PRETRAIN_BIN_PATH), PRETRAIN_BIN_PATH)1181 print("meta existe:", os.path.exists(PRETRAIN_META_PATH), PRETRAIN_META_PATH)1182 if os.path.exists(bin_path) and os.path.exists(meta_path):1183 with open(meta_path, "r", encoding="utf-8") as f:1184 meta = json.load(f)1185 if meta.get("target_tokens") == target_tokens and meta.get("vocab_size") == tokenizer.vocab_size:1186 total = meta["total_tokens"]1187 print(f"📂 Corpus de pré-entraînement déjà construit sur Drive ({total:,} tokens) — réutilisation.")1188 return bin_path, total1189 print("⚠️ Corpus existant sur Drive mais cible ou vocabulaire différents — reconstruction complète.")1190 1191 print(f"🏗️ Construction du corpus de pré-entraînement (~{target_tokens / 1e9:.2f}G tokens visés) "1192 f"depuis {len(PRETRAIN_SOURCES)} sources...")1193 1194 stats = {src["name"]: 0 for src in PRETRAIN_SOURCES}1195 active = []1196 for src in PRETRAIN_SOURCES:1197 try:1198 gen = src["make_gen"]()1199 active.append({"name": src["name"], "weight": src["weight"], "gen": gen})1200 print(f" ✅ Source active: {src['name']} (poids {src['weight']:.0%})")