CoolFace
Modelpublic

JuloEco/opsiom-fr-tokenizer

sourceHugging Faceunknownupdated 2d agoView on Hugging Face
1likes
main_GPU.py1739 linesDownload Raw Back to root
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%})")

Showing the first 1,200 of 1739 lines. Download the file for the rest.