diff --git a/MyIA.AI.Notebooks/GenAI/Texte/TransformerVariants/TV-03-Internalisation-CoT.ipynb b/MyIA.AI.Notebooks/GenAI/Texte/TransformerVariants/TV-03-Internalisation-CoT.ipynb new file mode 100644 index 0000000000..f772567c43 --- /dev/null +++ b/MyIA.AI.Notebooks/GenAI/Texte/TransformerVariants/TV-03-Internalisation-CoT.ipynb @@ -0,0 +1,657 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "id": "626e10c7", + "metadata": { + "papermill": { + "duration": 0.00345, + "end_time": "2026-09-29T10:41:41.369604+00:00", + "exception": false, + "start_time": "2026-09-29T10:41:41.366154+00:00", + "status": "completed" + }, + "tags": [] + }, + "source": [ + "# TV-03 -- Internalisation du raisonnement (CoT -> calcul interne) -- v3 CoT supervisee + comparaison answer-only\n", + "\n", + "Quatrieme tranche d'execution de l'Epic #17540 (Russell & Norvig arc B, raisonnement internalise). Cette v3 livre :\n", + "\n", + "1. Le module tv/task.py etendu : nouvelles fonctions lot_multi_hop_cot(), evaluer_multi_hop_cot(), entrainer_cot(), entrainer_multi_seed_cot() qui implementent la supervision CoT (chain-of-thought) sur la tache multi-sauts.\n", + "2. Comparaison multi-seed (4 graines) answer-only vs CoT supervisee sur multi-sauts 3 sauts : mesure discriminante de l'internalisation du raisonnement (cf Huang et al. 2026).\n", + "3. Cycle de mesure complet : entrainement CoT avec perte supervisee sur la chaine entiere (chaque token PAS/RECAP_j/cible doit etre predit correctement).\n", + "\n", + "## Pourquoi cette v3\n", + "\n", + "v1 (c.808) posait le discriminant H.2 sur 1 graine. v2 (c.811) etendait a 4 graines et validait le protocole PR review-discipline §C (edge >= 2sigma cross-seed, 17.2sigma). Cette v3 livre la COMPARAISON discriminante : reponse only vs CoT supervisee, les deux sur la meme architecture MHA 64-dim.\n", + "\n", + "## Ce qui n'est PAS dans cette tranche\n", + "\n", + "- Lecture SAE des representations internes (autre lane po-2027:CoursIA).\n", + "- Modeles plus grands (la discrimination MHA 64-dim est suffisante pour la preuve de concept).\n", + "- Optimisation du vocabulaire CoT (les jetons PAS/RECAP_j sont des placeholders pedagogiques ; un tokenization appris est un autre grain).\n" + ] + }, + { + "cell_type": "code", + "execution_count": 1, + "id": "c8c8ecd3", + "metadata": { + "execution": { + "iopub.execute_input": "2026-09-29T10:41:41.378926Z", + "iopub.status.busy": "2026-09-29T10:41:41.378656Z", + "iopub.status.idle": "2026-09-29T10:41:42.896883Z", + "shell.execute_reply": "2026-09-29T10:41:42.895929Z" + }, + "papermill": { + "duration": 1.523613, + "end_time": "2026-09-29T10:41:42.898599+00:00", + "exception": false, + "start_time": "2026-09-29T10:41:41.374986+00:00", + "status": "completed" + }, + "tags": [] + }, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "torch 2.13.0+cu126 | python 3.13\n", + "package tv charge OK : single-hop + multi-hop + multi-hop-CoT + multi-seed\n" + ] + } + ], + "source": [ + "import math\n", + "import sys\n", + "import time\n", + "\n", + "import torch\n", + "import torch.nn.functional as F\n", + "\n", + "sys.path.insert(0, '.')\n", + "from tv import ( # noqa: E402\n", + " PetitLM, Vocab,\n", + " evaluer_single_hop, evaluer_multi_hop, evaluer_multi_hop_cot,\n", + " entrainer, entrainer_multi_seed,\n", + " entrainer_cot, entrainer_multi_seed_cot,\n", + ")\n", + "\n", + "print(f'torch {torch.__version__} | python {sys.version_info.major}.{sys.version_info.minor}')\n", + "print('package tv charge OK : single-hop + multi-hop + multi-hop-CoT + multi-seed')\n" + ] + }, + { + "cell_type": "markdown", + "id": "0c3392b5", + "metadata": { + "papermill": { + "duration": 0.00292, + "end_time": "2026-09-29T10:41:42.908178+00:00", + "exception": false, + "start_time": "2026-09-29T10:41:42.905258+00:00", + "status": "completed" + }, + "tags": [] + }, + "source": [ + "## 1. Vocabulaires\n", + "\n", + "Deux vocabulaires de même structure :\n", + "\n", + "- **Single-hop** : `N_MARQUEURS=8`, `N_REMPLISSAGE=10`, `N_QUESTIONS=1` (le `QUESTION(0)` trivial). Vocab total 20 tokens.\n", + "- **Multi-sauts** : `N_MARQUEURS=8`, `N_REMPLISSAGE=10`, `N_QUESTIONS=3` (la cible est le `q`-ième marqueur, avec `q` choisi uniformément sur [0, 3)). Vocab total 22 tokens.\n", + "\n", + "Hasard exactitude : 1/8 dans les deux cas (la cible est l'un des 8 marqueurs ; le mécanisme à apprendre est la sélection conditionnelle, pas la mémorisation)." + ] + }, + { + "cell_type": "code", + "execution_count": 2, + "id": "1d2423b0", + "metadata": { + "execution": { + "iopub.execute_input": "2026-09-29T10:41:42.915763Z", + "iopub.status.busy": "2026-09-29T10:41:42.915345Z", + "iopub.status.idle": "2026-09-29T10:41:42.922056Z", + "shell.execute_reply": "2026-09-29T10:41:42.921136Z" + }, + "papermill": { + "duration": 0.011776, + "end_time": "2026-09-29T10:41:42.922862+00:00", + "exception": false, + "start_time": "2026-09-29T10:41:42.911086+00:00", + "status": "completed" + }, + "tags": [] + }, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "Vocab single-hop : VOCAB=20, N_QUESTIONS=1\n", + "Vocab multi-sauts : VOCAB=22, N_QUESTIONS=3\n", + "VOCAB_COT (multi-sauts + CoT) : 26 (= 22 + 1 JETON_PAS + 3 jetons RECAP)\n", + "Hasard exactitude : 1 / 8 = 0.1250\n" + ] + } + ], + "source": [ + "T = 64 # longueur de sequence commune aux deux taches\n", + "T_COT = T + 7 # 7 = 2*max_q+1 = 2*3+1 (chaine CoT)\n", + "v_single = Vocab(N_MARQUEURS=8, N_REMPLISSAGE=10, N_QUESTIONS=1)\n", + "v_multi = Vocab(N_MARQUEURS=8, N_REMPLISSAGE=10, N_QUESTIONS=3)\n", + "\n", + "MAX_Q = 3 # q_idx in [0, 3) -- borne sup pour le vocabulaire CoT\n", + "VOCAB_COT = v_multi.VOCAB + 1 + MAX_Q # VOCAB + JETON_PAS + MAX_Q jetons de recap\n", + "\n", + "def fabrique(vocab):\n", + " \"\"\"Fabrique un MHA vierge pour taches sans CoT (VOCAB standard).\"\"\"\n", + " return PetitLM(\n", + " vocab=vocab.VOCAB,\n", + " d_model=64,\n", + " n_heads=4,\n", + " n_kv_heads=4,\n", + " window=None,\n", + " n_couches=2,\n", + " )\n", + "\n", + "def fabrique_cot(vocab):\n", + " \"\"\"Fabrique un MHA vierge pour tache CoT (VOCAB etendu avec PAS + RECAP).\"\"\"\n", + " return PetitLM(\n", + " vocab=VOCAB_COT,\n", + " d_model=64,\n", + " n_heads=4,\n", + " n_kv_heads=4,\n", + " window=None,\n", + " n_couches=2,\n", + " )\n", + "\n", + "print(f'Vocab single-hop : VOCAB={v_single.VOCAB}, N_QUESTIONS={v_single.N_QUESTIONS}')\n", + "print(f'Vocab multi-sauts : VOCAB={v_multi.VOCAB}, N_QUESTIONS={v_multi.N_QUESTIONS}')\n", + "print(f'VOCAB_COT (multi-sauts + CoT) : {VOCAB_COT} (= {v_multi.VOCAB} + 1 JETON_PAS + {MAX_Q} jetons RECAP)')\n", + "print(f'Hasard exactitude : 1 / {v_single.N_MARQUEURS} = {1/v_single.N_MARQUEURS:.4f}')\n" + ] + }, + { + "cell_type": "markdown", + "id": "588afde7", + "metadata": { + "papermill": { + "duration": 0.002165, + "end_time": "2026-09-29T10:41:42.926871+00:00", + "exception": false, + "start_time": "2026-09-29T10:41:42.924706+00:00", + "status": "completed" + }, + "tags": [] + }, + "source": [ + "## 2. Mesure multi-seed single-hop (4 graines, 300 pas)\n", + "\n", + "Single-hop est trivialement resolue par MHA 102K params : on s'attend a 1.0000 +/- ~0 sur 4 graines.\n" + ] + }, + { + "cell_type": "code", + "execution_count": 3, + "id": "e6bc6b97", + "metadata": { + "execution": { + "iopub.execute_input": "2026-09-29T10:41:42.931695Z", + "iopub.status.busy": "2026-09-29T10:41:42.931433Z", + "iopub.status.idle": "2026-09-29T10:42:18.050409Z", + "shell.execute_reply": "2026-09-29T10:42:18.049594Z" + }, + "papermill": { + "duration": 35.123175, + "end_time": "2026-09-29T10:42:18.051840+00:00", + "exception": false, + "start_time": "2026-09-29T10:41:42.928665+00:00", + "status": "completed" + }, + "tags": [] + }, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "SINGLE-HOP multi-seed (4 graines, 300 pas) :\n", + " EXACTITUDE = 1.0000 +/- 0.0000 (hasard = 0.1250)\n", + " PERPLEXITE = 1.0016 +/- 0.0001 (hasard = 8.0)\n", + " Secondes total = 35.11 s\n", + " graine 0 : acc=1.0000 ppl=1.0017 sec=8.87\n", + " graine 1 : acc=1.0000 ppl=1.0015 sec=7.90\n", + " graine 7 : acc=1.0000 ppl=1.0015 sec=7.82\n", + " graine 42 : acc=1.0000 ppl=1.0015 sec=7.88\n" + ] + } + ], + "source": [ + "GRAINES = [0, 1, 7, 42]\n", + "PAS = 300\n", + "\n", + "result_s = entrainer_multi_seed(\n", + " fabrique, v_single, T=T, multi_hop=False, graines=GRAINES, pas=PAS\n", + ")\n", + "print(f'SINGLE-HOP multi-seed ({len(GRAINES)} graines, {PAS} pas) :')\n", + "print(f' EXACTITUDE = {result_s[\"acc_moy\"]:.4f} +/- {result_s[\"acc_std\"]:.4f} (hasard = 0.1250)')\n", + "print(f' PERPLEXITE = {result_s[\"ppl_moy\"]:.4f} +/- {result_s[\"ppl_std\"]:.4f} (hasard = 8.0)')\n", + "print(f' Secondes total = {result_s[\"secondes\"]:.2f} s')\n", + "for graine, acc, ppl, sec in result_s[\"brut\"]:\n", + " print(f' graine {graine} : acc={acc:.4f} ppl={ppl:.4f} sec={sec:.2f}')\n" + ] + }, + { + "cell_type": "markdown", + "id": "8f88d873", + "metadata": { + "papermill": { + "duration": 0.001726, + "end_time": "2026-09-29T10:42:18.055442+00:00", + "exception": false, + "start_time": "2026-09-29T10:42:18.053716+00:00", + "status": "completed" + }, + "tags": [] + }, + "source": [ + "## 3. Mesure multi-seed multi-sauts (4 graines, 300 pas)\n", + "\n", + "Multi-sauts 3 questions : le discriminant H.2 doit tenir sur 4 graines (moyenne significativement au-dessus du hasard, sans atteindre 1.0).\n" + ] + }, + { + "cell_type": "code", + "execution_count": 4, + "id": "c204ef74", + "metadata": { + "execution": { + "iopub.execute_input": "2026-09-29T10:42:18.060417Z", + "iopub.status.busy": "2026-09-29T10:42:18.059966Z", + "iopub.status.idle": "2026-09-29T10:42:50.573598Z", + "shell.execute_reply": "2026-09-29T10:42:50.572670Z" + }, + "papermill": { + "duration": 32.517539, + "end_time": "2026-09-29T10:42:50.574799+00:00", + "exception": false, + "start_time": "2026-09-29T10:42:18.057260+00:00", + "status": "completed" + }, + "tags": [] + }, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "MULTI-SAUTS multi-seed (4 graines, 300 pas, 3 sauts) :\n", + " EXACTITUDE = 0.4019 +/- 0.0161 (hasard = 0.1250)\n", + " PERPLEXITE = 2.9156 +/- 0.0943 (hasard = 8.0)\n", + " Secondes total = 32.51 s\n", + " graine 0 : acc=0.3945 ppl=2.9109 sec=8.01\n", + " graine 1 : acc=0.4180 ppl=3.0707 sec=8.00\n", + " graine 7 : acc=0.4160 ppl=2.8310 sec=8.06\n", + " graine 42 : acc=0.3789 ppl=2.8499 sec=8.03\n" + ] + } + ], + "source": [ + "result_m = entrainer_multi_seed(\n", + " fabrique, v_multi, T=T, multi_hop=True, graines=GRAINES, pas=PAS\n", + ")\n", + "print(f'MULTI-SAUTS multi-seed ({len(GRAINES)} graines, {PAS} pas, 3 sauts) :')\n", + "print(f' EXACTITUDE = {result_m[\"acc_moy\"]:.4f} +/- {result_m[\"acc_std\"]:.4f} (hasard = 0.1250)')\n", + "print(f' PERPLEXITE = {result_m[\"ppl_moy\"]:.4f} +/- {result_m[\"ppl_std\"]:.4f} (hasard = 8.0)')\n", + "print(f' Secondes total = {result_m[\"secondes\"]:.2f} s')\n", + "for graine, acc, ppl, sec in result_m[\"brut\"]:\n", + " print(f' graine {graine} : acc={acc:.4f} ppl={ppl:.4f} sec={sec:.2f}')\n" + ] + }, + { + "cell_type": "markdown", + "id": "f2dc297d", + "metadata": { + "papermill": { + "duration": 0.001819, + "end_time": "2026-09-29T10:42:50.578713+00:00", + "exception": false, + "start_time": "2026-09-29T10:42:50.576894+00:00", + "status": "completed" + }, + "tags": [] + }, + "source": [ + "## 4. Bilan tranche v2\n", + "\n", + "Multi-seed >=4 sur single-hop vs multi-sauts -- discriminant H.2 mesure sur 4 graines.\n" + ] + }, + { + "cell_type": "code", + "execution_count": 5, + "id": "bcaee334", + "metadata": { + "execution": { + "iopub.execute_input": "2026-09-29T10:42:50.583987Z", + "iopub.status.busy": "2026-09-29T10:42:50.583639Z", + "iopub.status.idle": "2026-09-29T10:42:50.589453Z", + "shell.execute_reply": "2026-09-29T10:42:50.588632Z" + }, + "papermill": { + "duration": 0.009575, + "end_time": "2026-09-29T10:42:50.590097+00:00", + "exception": false, + "start_time": "2026-09-29T10:42:50.580522+00:00", + "status": "completed" + }, + "tags": [] + }, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "BILAN multi-seed :\n", + " single-hop : 1.0000 +/- 0.0000\n", + " multi-sauts : 0.4019 +/- 0.0161\n", + " rapport multi/single = 0.402\n", + " secondes total (4 graines single + 4 graines multi) = 67.62 s\n", + "\n", + "Conclusion : Single-hop trivialement resolu (>=0.99) ; multi-sauts plus bas avec >=4 graines.\n", + "Discriminant H.2 multi-seed : OK (single >> multi, multi > hasard).\n" + ] + } + ], + "source": [ + "print('BILAN multi-seed :')\n", + "print(f' single-hop : {result_s[\"acc_moy\"]:.4f} +/- {result_s[\"acc_std\"]:.4f}')\n", + "print(f' multi-sauts : {result_m[\"acc_moy\"]:.4f} +/- {result_m[\"acc_std\"]:.4f}')\n", + "rapport = result_m[\"acc_moy\"] / result_s[\"acc_moy\"] if result_s[\"acc_moy\"] > 0 else float(\"inf\")\n", + "print(f' rapport multi/single = {rapport:.3f}')\n", + "print(f' secondes total (4 graines single + 4 graines multi) = {result_s[\"secondes\"] + result_m[\"secondes\"]:.2f} s')\n", + "print()\n", + "if result_s['acc_moy'] >= 0.99 and result_m['acc_moy'] < result_s['acc_moy']:\n", + " print('Conclusion : Single-hop trivialement resolu (>=0.99) ; multi-sauts plus bas avec >=4 graines.')\n", + " print('Discriminant H.2 multi-seed : OK (single >> multi, multi > hasard).')\n", + "else:\n", + " print(f'ATTENTION : single-hop = {result_s[\"acc_moy\"]:.4f}, multi-sauts = {result_m[\"acc_moy\"]:.4f}.')\n", + " print('Le discriminant H.2 n est PAS clairement tenu. Investiguer la graine fautive.')\n" + ] + }, + { + "cell_type": "markdown", + "id": "6e5efbfd", + "metadata": { + "papermill": { + "duration": 0.001829, + "end_time": "2026-09-29T10:42:50.593768+00:00", + "exception": false, + "start_time": "2026-09-29T10:42:50.591939+00:00", + "status": "completed" + }, + "tags": [] + }, + "source": [ + "## 5. Comparaison CoT supervisee vs answer-only (4 graines, 300 pas)\n", + "\n", + "Le discriminant central du programme #17540 : la supervision CoT ameliore-t-elle significativement la memorisation multi-sauts ?\n", + "\n", + "- answer-only (v2 cellule 7) : `entrainer_multi_seed` avec `multi_hop=True` (cible = marqueur direct)\n", + "- CoT supervise (v3) : `entrainer_multi_seed_cot` avec `max_q=3` (cible finale + chaine PAS/RECAP)\n", + "\n", + "Tell c.1493 strict fondateur nuance : la mesure CoT supervise la chaine entiere via une perte massee par item (les jetons au-dela du q_idx reel sont ignores ; voir entrainer_cot dans tv/task.py)." + ] + }, + { + "cell_type": "code", + "execution_count": 6, + "id": "6c184943", + "metadata": { + "execution": { + "iopub.execute_input": "2026-09-29T10:42:50.598752Z", + "iopub.status.busy": "2026-09-29T10:42:50.598437Z", + "iopub.status.idle": "2026-09-29T10:43:25.823094Z", + "shell.execute_reply": "2026-09-29T10:43:25.822022Z" + }, + "papermill": { + "duration": 35.228925, + "end_time": "2026-09-29T10:43:25.824583+00:00", + "exception": false, + "start_time": "2026-09-29T10:42:50.595658+00:00", + "status": "completed" + }, + "tags": [] + }, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "CoT multi-seed (4 graines, 300 pas, 3 sauts) :\n", + " EXACTITUDE = 0.4146 +/- 0.0257 (hasard = 0.1250)\n", + " PERPLEXITE = 3.0279 +/- 0.0852 (hasard = 8.0)\n", + " Secondes total = 35.22 s\n", + " graine 0 : acc=0.4453 ppl=2.9318 sec=8.79\n", + " graine 1 : acc=0.4121 ppl=3.1603 sec=8.68\n", + " graine 7 : acc=0.4258 ppl=3.0383 sec=8.54\n", + " graine 42 : acc=0.3750 ppl=2.9810 sec=8.78\n" + ] + } + ], + "source": [ + "result_cot = entrainer_multi_seed_cot(\n", + " fabrique_cot, v_multi, T_cot=T_COT, graines=GRAINES, pas=PAS, batch=32, max_q=MAX_Q\n", + ")\n", + "print(f'CoT multi-seed ({len(GRAINES)} graines, {PAS} pas, 3 sauts) :')\n", + "print(f' EXACTITUDE = {result_cot[\"acc_moy\"]:.4f} +/- {result_cot[\"acc_std\"]:.4f} (hasard = 0.1250)')\n", + "print(f' PERPLEXITE = {result_cot[\"ppl_moy\"]:.4f} +/- {result_cot[\"ppl_std\"]:.4f} (hasard = 8.0)')\n", + "print(f' Secondes total = {result_cot[\"secondes\"]:.2f} s')\n", + "for graine, acc, ppl, sec in result_cot[\"brut\"]:\n", + " print(f' graine {graine} : acc={acc:.4f} ppl={ppl:.4f} sec={sec:.2f}')\n" + ] + }, + { + "cell_type": "markdown", + "id": "281b8902", + "metadata": { + "papermill": { + "duration": 0.001997, + "end_time": "2026-09-29T10:43:25.828588+00:00", + "exception": false, + "start_time": "2026-09-29T10:43:25.826591+00:00", + "status": "completed" + }, + "tags": [] + }, + "source": [ + "## 6. Bilan tranche v3 -- comparaison discriminante\n", + "\n", + "Compare answer-only (multi-hop) vs CoT supervisee (multi-hop + chain-of-thought) sur la meme architecture MHA 64-dim et 4 graines communes.\n", + "\n", + "## Corrections de mesure appliquees en v3\n", + "\n", + "L'exactitude CoT mesuree en cellule 11 etait `0.0000 +/- 0.0000` sur 4 graines (`ppl = 17 994`). Deux defauts de `tv/task.py` en sont la cause, corriges dans cette tranche :\n", + "\n", + "1. **`evaluer_multi_hop_cot` (ligne 372)** : lisait `modele(lot.x)[:, -1]` -- les logits a la derniere position predisent au-dela de la sequence. La derniere position predictive entrainee par `entrainer_cot` est `T_seq - 2` (couvre `[start_pred - 1, T - 1)`). Corrige en `[:, -2]`.\n", + "2. **`entrainer_cot` (lignes 421-436)** : le masque `(pas_positions <= q)` n'incluait jamais le slot cible. La chaine etant de **taille fixe** (cf `lot_multi_hop_cot`), la cible est toujours en fin de chaine (`j = 2*max_q`), quel que soit `q_idx` -- le slot `j = 2*q + 1` ne porte que `RECAP_q`. Le modele n'etait donc jamais entraine a produire la cible. Corrige en `(pas_positions <= q) | (j == 2*max_q)`.\n", + "\n", + "Mesure de controle (env `coursia-sae`, torch 2.13, CPU, 300 pas, batch 32, cf commentaire ai-01 c.5331913672) sur la tete b71ef3782b + les deux corrections :\n", + "\n", + "| Correctif | Graine 0 | Graine 1 |\n", + "|---|---|---|\n", + "| Lecture `-1` (defaut 1) | 0.0000 | 0.0000 |\n", + "| Lecture `-2` seule (defaut 2 reste) | 0.0000 | 0.0000 |\n", + "| Lecture `-2` + cible dans la perte | **0.4336** (ppl 3.02) | **0.4219** (ppl 3.07) |\n", + "\n", + "Les deux corrections ensemble ramènent le CoT supervise au niveau d'answer-only (0.4019 +/- 0.0161) -- la conclusion \"CoT < answer-only\" disparait.\n", + "\n", + "## Caveat pedagogique (porte par ai-01)\n", + "\n", + "La chaine generee `PAS, RECAP_0, PAS, RECAP_1, PAS, RECAP_2` (cf `tv/task.py` ligne 224 et suivantes) est **constante pour tous les items** : elle ne porte aucune information tirée de l'entree. Meme corrigee, la comparaison oppose answer-only a answer-only precede d'un prefixe constant. Pour tester l'internalisation du raisonnement, les jetons intermediaires doivent dependre de la sequence (par exemple, les marqueurs visites a chaque saut). C'est un autre grain.\n", + "\n", + "Cette v3 delivre donc un **protocole fonctionnel** (CoT supervise evalue sur la position correcte + cible dans la perte), pas une **mesure d'internalisation** au sens fort de Huang et al. 2026. La mesure elle-meme (egalite CoT ≈ answer-only) reflete le caractere trivial de la chaine, pas une absence de capacite du CoT supervise.\n" + ] + }, + { + "cell_type": "code", + "execution_count": 7, + "id": "67e6b80e", + "metadata": { + "execution": { + "iopub.execute_input": "2026-09-29T10:43:25.833627Z", + "iopub.status.busy": "2026-09-29T10:43:25.833377Z", + "iopub.status.idle": "2026-09-29T10:43:25.839436Z", + "shell.execute_reply": "2026-09-29T10:43:25.838582Z" + }, + "papermill": { + "duration": 0.009627, + "end_time": "2026-09-29T10:43:25.840213+00:00", + "exception": false, + "start_time": "2026-09-29T10:43:25.830586+00:00", + "status": "completed" + }, + "tags": [] + }, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + " Rapport CoT / answer-only = 1.032 (ecart +0.0127, std combine 0.0303)\n", + "\n", + "Conclusion : egalite answer-only / CoT supervisee a l'interieur du bruit.\n", + "Cohérent avec le caveat de la cellule précédente : la chaîne supervisée est constante\n", + "(préfixe fixe, aucune information d'entrée), donc la comparaison oppose answer-only à\n", + "answer-only précédé d'un préfixe constant -- le protocole v3 valide la boucle CoT,\n", + "pas une internalisation du raisonnement.\n" + ] + } + ], + "source": [ + "rapport_cot = result_cot['acc_moy'] / max(result_m['acc_moy'], 1e-9)\n", + "ecart_cot = result_cot['acc_moy'] - result_m['acc_moy']\n", + "# Std combine des deux moyennes (variances independantes, 4 graines chacune)\n", + "std_combine = (result_m['acc_std']**2 + result_cot['acc_std']**2) ** 0.5\n", + "print(f' Rapport CoT / answer-only = {rapport_cot:.3f} (ecart {ecart_cot:+.4f}, std combine {std_combine:.4f})')\n", + "print()\n", + "if ecart_cot > 2 * std_combine:\n", + " print('Conclusion : CoT supervisee > answer-only au-dela de 2*std combine.')\n", + " print('La supervision de la chaine aide le modele au-dela du bruit de mesure.')\n", + "elif ecart_cot < -2 * std_combine:\n", + " print('Conclusion : CoT supervisee < answer-only au-dela de 2*std combine.')\n", + " print('Surprenant -- generer la chaine entiere en 300 pas est difficile. Investiguer.')\n", + "else:\n", + " print(\"Conclusion : egalite answer-only / CoT supervisee a l'interieur du bruit.\")\n", + " print(\"Cohérent avec le caveat de la cellule précédente : la chaîne supervisée est constante\")\n", + " print(\"(préfixe fixe, aucune information d'entrée), donc la comparaison oppose answer-only à\")\n", + " print(\"answer-only précédé d'un préfixe constant -- le protocole v3 valide la boucle CoT,\")\n", + " print(\"pas une internalisation du raisonnement.\")" + ] + }, + { + "cell_type": "markdown", + "id": "73e2bf66", + "metadata": { + "papermill": { + "duration": 0.002139, + "end_time": "2026-09-29T10:43:25.844505+00:00", + "exception": false, + "start_time": "2026-09-29T10:43:25.842366+00:00", + "status": "completed" + }, + "tags": [] + }, + "source": [ + "## 7. Selfcheck invariant attn_banded ≡ attn_masked\n", + "\n", + "Ancre first-hand (Tell c.1493 strict fondateur nuance strict) de l'invariant documente dans la docstring de tv/model.py. TV-00b cellules 16 et 31 mesurent ~1.2e-07 ; le selfcheck ici produit **max_abs_diff = 1.788e-07** (global, W=8) avec [W=8: 1.788e-07, W=16: 1.490e-07, W=32: 1.192e-07] -- une mesure du **meme ordre de grandeur** que TV-00b (1.2e-07), l'ecart W=8 vs W=32 tenant a la discretisation fp32 cumulee dans le produit QK^T. Seed 42, T=64, dh=64, W in (8, 16, 32). Verdict : `passed = max_abs_diff <= 1e-5`." + ] + }, + { + "cell_type": "code", + "execution_count": 8, + "id": "4436a3b7", + "metadata": { + "execution": { + "iopub.execute_input": "2026-09-29T10:43:25.849964Z", + "iopub.status.busy": "2026-09-29T10:43:25.849700Z", + "iopub.status.idle": "2026-09-29T10:43:25.877166Z", + "shell.execute_reply": "2026-09-29T10:43:25.876244Z" + }, + "papermill": { + "duration": 0.03139, + "end_time": "2026-09-29T10:43:25.877969+00:00", + "exception": false, + "start_time": "2026-09-29T10:43:25.846579+00:00", + "status": "completed" + }, + "tags": [] + }, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "max_abs_diff global : 1.788e-07\n", + "passed : True\n", + " W= 8 : max_abs_diff = 1.788e-07\n", + " W= 16 : max_abs_diff = 1.490e-07\n", + " W= 32 : max_abs_diff = 1.192e-07\n" + ] + } + ], + "source": [ + "import sys\n", + "sys.path.insert(0, '.')\n", + "from tv import selfcheck_attn_equivalence\n", + "\n", + "r = selfcheck_attn_equivalence()\n", + "print(f\"max_abs_diff global : {r['max_abs_diff']:.3e}\")\n", + "print(f\"passed : {r['passed']}\")\n", + "for s in r['samples']:\n", + " print(f\" W={s['window']:>3d} : max_abs_diff = {s['max_abs_diff']:.3e}\")\n", + "\n", + "assert r['passed'], f\"selfcheck rate : {r['max_abs_diff']:.3e} > tol {r['tol']}\"\n" + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": "Python 3", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.13.15" + }, + "papermill": { + "default_parameters": {}, + "duration": 107.282966, + "end_time": "2026-09-29T10:43:27.091514+00:00", + "environment_variables": {}, + "exception": null, + "input_path": "TV-03-Internalisation-CoT.ipynb", + "output_path": "TV-03-Internalisation-CoT_output.ipynb", + "parameters": {}, + "start_time": "2026-09-29T10:41:39.808548+00:00", + "version": "2.7.0" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} \ No newline at end of file diff --git a/MyIA.AI.Notebooks/GenAI/Texte/TransformerVariants/tv/__init__.py b/MyIA.AI.Notebooks/GenAI/Texte/TransformerVariants/tv/__init__.py new file mode 100644 index 0000000000..1b4cd359cc --- /dev/null +++ b/MyIA.AI.Notebooks/GenAI/Texte/TransformerVariants/tv/__init__.py @@ -0,0 +1,86 @@ +"""tv -- package réutilisable issu de TV-00b (Attention Variants from scratch). + +Ce package extrait les briques canoniques du notebook TV-00b en modules Python +importables, sans dépendance sur le reste du notebook. Trois classes principales : + +- :class:`VariantAttn` : attention unifiée MHA / MQA / GQA / SWA (leviers : n_kv_heads, window, banded). +- :class:`Bloc` : bloc pré-norm résiduel (LN -> attn -> résiduel, LN -> MLP -> résiduel). +- :class:`PetitLM` : transformer minimal (embedding + N blocs + LayerNorm + tête linéaire). + +Helpers publics : :func:`build_kv_heads`, :func:`attn_masked`, :func:`attn_banded`, +:func:`causal_window_mask`. + +Origine +------- + +Issue **#17540** (Russell & Norvig arc B, raisonnement internalisé). Vérification +organ-first c.807 : TV-00b contient bien ces classes (cellules 12, 26). Avant +l'extraction, les classes n'étaient **pas** importables : elles vivaient dans les +cellules du notebook. Le grain d'exécution est cette extraction. + +Témoin négatif +-------------- + +Le témoin négatif du grain #17540 (Huang 2026) est l'entraînement avec vs sans +supervision de chaîne de pensée (CoT) sur une tâche synthétique multi-sauts. +L'extraction du modèle canonique n'est pas le témoin négatif elle-même ; elle +rend le témoin **faisable** dans un grain ultérieur, et c'est précisément le +geste attendu par l'organ-first (question 3 : exporter / refactorer dans la +série source). + +Voir aussi +---------- + +- TV-00b : ``Attention-Variants-from-scratch.ipynb`` (origine). +- TV-03 : ``TV-03-Internalisation-CoT.ipynb`` (à venir — grain d'exécution). +- Issue : **#17540**. +""" +from __future__ import annotations + +from .model import ( + VariantAttn, + Bloc, + PetitLM, + build_kv_heads, + attn_masked, + attn_banded, + causal_window_mask, + selfcheck_attn_equivalence, +) +from .task import ( + Vocab, + Lot, + lot_single_hop, + lot_multi_hop, + lot_multi_hop_cot, + evaluer_single_hop, + evaluer_multi_hop, + evaluer_multi_hop_cot, + entrainer, + entrainer_multi_seed, + entrainer_cot, + entrainer_multi_seed_cot, +) + +__all__ = [ + "VariantAttn", + "Bloc", + "PetitLM", + "build_kv_heads", + "attn_masked", + "attn_banded", + "causal_window_mask", + "selfcheck_attn_equivalence", + "Vocab", + "Lot", + "lot_single_hop", + "lot_multi_hop", + "lot_multi_hop_cot", + "evaluer_single_hop", + "evaluer_multi_hop", + "evaluer_multi_hop_cot", + "entrainer", + "entrainer_multi_seed", + "entrainer_cot", + "entrainer_multi_seed_cot", +] diff --git a/MyIA.AI.Notebooks/GenAI/Texte/TransformerVariants/tv/model.py b/MyIA.AI.Notebooks/GenAI/Texte/TransformerVariants/tv/model.py new file mode 100644 index 0000000000..014efd186f --- /dev/null +++ b/MyIA.AI.Notebooks/GenAI/Texte/TransformerVariants/tv/model.py @@ -0,0 +1,212 @@ +"""Briques canoniques d'un mini-transformer, extraites de TV-00b. + +Ce module isole les pièces qui vivent dans les cellules du notebook +``TV-00b-Attention-Variants-from-scratch.ipynb`` pour les rendre importables +par un autre notebook ou un test, sans dépendance au reste du notebook. + +Trois classes principales : + +- :class:`VariantAttn` : attention unifiée (MHA / MQA / GQA / SWA). +- :class:`Bloc` : bloc pré-norm résiduel. +- :class:`PetitLM` : transformer canonique (embedding + N blocs + LN + tête). + +Quatre helpers : + +- :func:`build_kv_heads` : étale les têtes KV sur les têtes Q (GQA / MQA). +- :func:`attn_masked` : attention masquée (causale + fenêtre) O(T^2). +- :func:`attn_banded` : attention en bande SWA O(T * W). +- :func:`causal_window_mask` : masque booléen (T, T) partagé par les deux fonctions. + +Origine : cellules 12 (VariantAttn + build_kv_heads), 4 (causal_window_mask + +attn_masked), 15 (attn_banded), 26 (Bloc + PetitLM) de TV-00b. + +Invariants vérifiés (cf TV-00b cellules 12-15) : + +- L'étalement KV utilise ``repeat_interleave`` (vues partagées, pas de copie). +- Le coût en paramètres des projections K/V décroît avec ``n_kv_heads``. +- ``attn_banded`` et ``attn_masked`` rendent le même tenseur pour les mêmes (q, k, v, W). + Ancré en continu par :func:`selfcheck_attn_equivalence` (B=2, H=4, T=64, dh=64, + W in (8, 16, 32), seed 42, tol float32 1e-5). Mesure typique ~1.2e-07 (cf + TV-00b cellules 16 et 31). +""" +from __future__ import annotations + +import math + +import torch +import torch.nn as nn +import torch.nn.functional as F + + +def causal_window_mask(T: int, window: int | None = None, device: torch.device | str = "cpu") -> torch.Tensor: + """Masque booléen (T, T) : True si la requête i a le droit de lire la clé j. + + window=None : causal pur (j <= i). + window=W : causal ET bande locale (i - j < W), au plus W clés par requête. + """ + i = torch.arange(T, device=device).view(T, 1) + j = torch.arange(T, device=device).view(1, T) + mask = j <= i + if window is not None: + mask = mask & ((i - j) < window) + return mask + + +def attn_masked(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, window: int | None = None) -> torch.Tensor: + """Attention sur (B, H, T, dh). Forme lisible ; reste O(T^2) même en fenêtre.""" + scores = q @ k.transpose(-2, -1) / math.sqrt(q.shape[-1]) + scores = scores.masked_fill(~causal_window_mask(q.shape[-2], window, q.device), float("-inf")) + return torch.softmax(scores, dim=-1) @ v + + +def attn_banded(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, window: int) -> torch.Tensor: + """SWA en bande : ne matérialise que les W colonnes utiles → O(T*W). + + On décale les clés de W-1 vers la droite (remplissage à gauche), puis on + découpe une fenêtre glissante de largeur W le long de l'axe des positions. + La requête i voit alors les positions décalées relatives 0..W-1, dont la + dernière (offset W-1) est la position i. + """ + B, H, T, dh = q.shape + W = min(window, T) + k_win = F.pad(k, (0, 0, W - 1, 0)).unfold(2, W, 1) + v_win = F.pad(v, (0, 0, W - 1, 0)).unfold(2, W, 1) + scores = torch.einsum("bhtd,bhtdw->bhtw", q, k_win) / math.sqrt(dh) + decalage = torch.arange(W, device=q.device) + positions = torch.arange(T, device=q.device) + keep = decalage.view(1, 1, 1, W) >= (W - 1 - positions).view(1, 1, T, 1) + scores = scores.masked_fill(~keep, float("-inf")) + return torch.einsum("bhtw,bhtdw->bhtd", torch.softmax(scores, dim=-1), v_win) + + +def build_kv_heads(kv: torch.Tensor, n_heads: int, n_kv_heads: int) -> torch.Tensor: + """Étale les têtes KV sur les têtes Q. + + MHA (n_kv_heads == n_heads) : identité, on ne touche pas au tenseur. + MQA (n_kv_heads == 1) : l'unique tête KV est répétée n_heads fois. + GQA : la tête Q h reçoit la tête KV h // (n_heads / n_kv_heads). + + Retourne des **vues** (``repeat_interleave``) : les poids restent partagés. + """ + if n_kv_heads == n_heads: + return kv + return kv.repeat_interleave(n_heads // n_kv_heads, dim=1) + + +class VariantAttn(nn.Module): + """Les quatre variantes d'un seul tenant : les leviers sont (n_kv_heads, window). + + n_kv_heads == n_heads -> MHA | window is not None -> SWA + n_kv_heads == 1 -> MQA | banded=True -> SWA en bande O(T*W) + 1 < n_kv_heads < n_heads -> GQA | sinon -> masque O(T^2) + """ + + def __init__( + self, + d_model: int, + n_heads: int, + n_kv_heads: int | None = None, + window: int | None = None, + banded: bool = False, + ): + super().__init__() + n_kv_heads = n_heads if n_kv_heads is None else n_kv_heads + assert n_heads % n_kv_heads == 0, "les têtes Q doivent se répartir en groupes égaux" + self.h, self.g, self.window, self.banded = n_heads, n_kv_heads, window, banded + self.dh = d_model // n_heads + self.q_proj = nn.Linear(d_model, n_heads * self.dh, bias=False) + self.k_proj = nn.Linear(d_model, n_kv_heads * self.dh, bias=False) + self.v_proj = nn.Linear(d_model, n_kv_heads * self.dh, bias=False) + self.o_proj = nn.Linear(n_heads * self.dh, d_model, bias=False) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + B, T, _ = x.shape + q = self.q_proj(x).view(B, T, self.h, self.dh).transpose(1, 2) + k = self.k_proj(x).view(B, T, self.g, self.dh).transpose(1, 2) + v = self.v_proj(x).view(B, T, self.g, self.dh).transpose(1, 2) + k = build_kv_heads(k, self.h, self.g) + v = build_kv_heads(v, self.h, self.g) + if self.banded and self.window is not None: + o = attn_banded(q, k, v, self.window) + else: + o = attn_masked(q, k, v, self.window) + return self.o_proj(o.transpose(1, 2).reshape(B, T, self.h * self.dh)) + + +class Bloc(nn.Module): + """Bloc pré-norm résiduel : x = x + attn(LN(x)); x = x + MLP(LN(x)).""" + + def __init__(self, d_model: int, n_heads: int, n_kv_heads: int, window: int | None): + super().__init__() + self.ln1, self.ln2 = nn.LayerNorm(d_model), nn.LayerNorm(d_model) + self.attn = VariantAttn(d_model, n_heads, n_kv_heads=n_kv_heads, window=window) + self.mlp = nn.Sequential( + nn.Linear(d_model, 4 * d_model), + nn.GELU(), + nn.Linear(4 * d_model, d_model), + ) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + x = x + self.attn(self.ln1(x)) + return x + self.mlp(self.ln2(x)) + + +class PetitLM(nn.Module): + """Transformer minimal : embedding + N blocs pré-norm + LN final + tête linéaire.""" + + def __init__( + self, + vocab: int, + d_model: int, + n_heads: int, + n_kv_heads: int, + window: int | None, + n_couches: int, + ): + super().__init__() + self.emb = nn.Embedding(vocab, d_model) + self.blocs = nn.ModuleList( + [Bloc(d_model, n_heads, n_kv_heads, window) for _ in range(n_couches)] + ) + self.ln_f = nn.LayerNorm(d_model) + self.tete = nn.Linear(d_model, vocab, bias=False) + + def forward(self, idx: torch.Tensor) -> torch.Tensor: + x = self.emb(idx) + for b in self.blocs: + x = b(x) + return self.tete(self.ln_f(x)) + + +def selfcheck_attn_equivalence(B: int = 2, H: int = 4, T: int = 64, dh: int = 64, + windows: tuple[int, ...] = (8, 16, 32), + seed: int = 42, tol: float = 1e-5) -> dict: + """Ancre l'invariant documente dans la docstring du module. + + Pour chaque fenetre W dans ``windows``, mesure ``max|attn_banded - attn_masked|`` + sur des tenseurs (q, k, v) tires aleatoirement avec la graine ``seed``. Le + resultat est borne par ``tol`` (float32, juste en deca de la limite pratique + mesuree dans TV-00b cellules 16 et 31, ~1.2e-07). + + Retourne ``{"max_abs_diff": float, "passed": bool, "samples": list[dict]}`` + ou chaque sample porte ``{"window": int, "max_abs_diff": float}``. + + Le test est strictement CPU (le device suit ``q.device``), pas de gradient + necessaire. Si un echantillon echoue, le dictionnaire le signale dans + ``samples`` et ``passed = False``. + """ + g = torch.Generator(device="cpu").manual_seed(seed) + samples: list[dict] = [] + worst = 0.0 + for W in windows: + q = torch.randn(B, H, T, dh, generator=g) * 0.5 + k = torch.randn(B, H, T, dh, generator=g) * 0.5 + v = torch.randn(B, H, T, dh, generator=g) * 0.5 + with torch.no_grad(): + o_band = attn_banded(q, k, v, W) + o_mask = attn_masked(q, k, v, W) + diff = float((o_band - o_mask).abs().max().item()) + samples.append({"window": W, "max_abs_diff": diff}) + worst = max(worst, diff) + return {"max_abs_diff": worst, "passed": worst <= tol, "samples": samples, + "tol": tol, "B": B, "H": H, "T": T, "dh": dh, "seed": seed} diff --git a/MyIA.AI.Notebooks/GenAI/Texte/TransformerVariants/tv/task.py b/MyIA.AI.Notebooks/GenAI/Texte/TransformerVariants/tv/task.py new file mode 100644 index 0000000000..89f39bc81f --- /dev/null +++ b/MyIA.AI.Notebooks/GenAI/Texte/TransformerVariants/tv/task.py @@ -0,0 +1,497 @@ +"""Tâches synthétiques pour l'étude de l'internalisation du raisonnement. + +Ce module pose les **deux** tâches utilisées par TV-03 : + +- :class:`MarqueurSingleHop` : un marqueur en position 0, du remplissage, un jeton de requête + en dernière position. Le modèle doit retrouver le marqueur en sortie. Tâche **single-hop** + (résolue trivialement par attention directe, voir TV-00b cellule 26). + +- :class:`MarqueurMultiHop` : **n** marqueurs en début de séquence, puis un **jeton de + question** indiquant lequel des marqueurs est demandé. Tâche **multi-sauts** au sens de + Huang et al. 2026 : le modèle doit (a) reconnaître la valeur du jeton QUESTION, puis + (b) récupérer le marqueur correspondant. C'est le discriminant H.2 du dispatch ai-01 + sur #17540 — une tâche qu'un modèle sans chaîne de pensée résout aussi bien ne démontre + rien. + +Hasard exactitude : + +- :class:`MarqueurSingleHop` : 1 / ``N_MARQUEURS`` (1/8 par défaut). +- :class:`MarqueurMultiHop` : 1 / ``N_MARQUEURS`` (la cible est l'un des marqueurs, + indépendamment de la question ; le mécanisme à apprendre est la sélection conditionnelle, + pas la simple mémorisation). + +Tell c.1493 strict fondateur nuance — **origine** : ces deux tâches sont dérivées de TV-00b +cellule 26 (la single-hop), avec extension multi-sauts. La structure est volontairement +minimale pour qu'un transformer jouet (~100K params) puisse la résoudre avec marge en +CPU, et qu'un entraînement CoT-vs-answer-only y soit discriminable. +""" +from __future__ import annotations + +import math +from dataclasses import dataclass + +import torch +import torch.nn.functional as F + + +@dataclass +class Vocab: + """Vocabulaire d'une tâche marqueur. + + Le vocabulaire a trois couches : + - ``N_MARQUEURS`` tokens distincts (les "objets" à récupérer). + - ``N_REMPLISSAGE`` tokens de remplissage (les "distracteurs"). + - ``N_QUESTIONS`` tokens QUESTION(1..N_QUESTIONS) (le "selecteur"). + - 1 jeton REQUETE final (le "déclencheur de réponse"). + """ + + N_MARQUEURS: int = 8 + N_REMPLISSAGE: int = 10 + N_QUESTIONS: int = 1 # 1 = single-hop (Q=0 trivial), >1 = multi-sauts + + def __post_init__(self): + assert self.N_MARQUEURS >= 2 + assert self.N_REMPLISSAGE >= 2 + assert self.N_QUESTIONS >= 1 + self.taille_marqueurs = self.N_MARQUEURS + self.taille_remplissage = self.N_MARQUEURS + self.N_REMPLISSAGE + self.taille_questions = self.taille_remplissage + self.N_QUESTIONS + self.JETON_REQUETE = self.taille_questions + self.VOCAB = self.JETON_REQUETE + 1 + + def jeton_question(self, k: int) -> int: + """Indice du token QUESTION(k). k dans [0, N_QUESTIONS).""" + assert 0 <= k < self.N_QUESTIONS + return self.taille_remplissage + k + + +@dataclass +class Lot: + """Un lot (batch) de séquences avec leur cible. + + - ``x`` : (B, T) tenseur d'indices. + - ``y`` : (B,) cible = l'indice du marqueur que le modèle doit prédire à la position + de requête. + - ``q`` : (B,) indice de la question (quel marqueur est demandé). 0 en single-hop. + """ + + x: torch.Tensor + y: torch.Tensor + q: torch.Tensor + + +def lot_single_hop(n: int, T: int, gen: torch.Generator, vocab: Vocab) -> Lot: + """Génère un lot single-hop : un marqueur en position 0, requête en T-1, cible = marqueur. + + Reproduit TV-00b cellule 26. + """ + assert vocab.N_QUESTIONS == 1, "lot_single_hop exige N_QUESTIONS == 1" + marqueur = torch.randint(0, vocab.N_MARQUEURS, (n, 1), generator=gen) + remplissage = torch.randint( + vocab.N_MARQUEURS, + vocab.taille_remplissage, + (n, T - 2), + generator=gen, + ) + requete = torch.full((n, 1), vocab.JETON_REQUETE) + x = torch.cat([marqueur, remplissage, requete], dim=1) + q = torch.zeros(n, dtype=torch.long) + return Lot(x=x, y=x[:, 0].clone(), q=q) + + +def lot_multi_hop( + n: int, + T: int, + gen: torch.Generator, + vocab: Vocab, +) -> Lot: + """Génère un lot multi-sauts : N_QUESTIONS marqueurs + jeton QUESTION(k) + requête. + + Séquence : + [M1, M2, ..., M_Q, ..., remplissage, QUESTION(k), REQUETE] + + où ``M_k`` est la cible que le modèle doit prédire à la position de requête. + Le discriminateur : ``QUESTION(k)`` force le modèle à ignorer les ``Q-1`` autres + marqueurs ; un modèle qui se contente d'extraire le dernier marqueur est en + échec quand ``k != Q-1``. + """ + Q = vocab.N_QUESTIONS + assert Q >= 2, "lot_multi_hop exige N_QUESTIONS >= 2" + + # Marqueurs en début (Q positions) + marqueurs = torch.randint(0, vocab.N_MARQUEURS, (n, Q), generator=gen) + + # Question choisie par item (uniforme sur [0, Q)) + q_idx = torch.randint(0, Q, (n,), generator=gen) + + # Remplissage (entre les marqueurs et la zone question/requete) + n_remplissage = T - Q - 2 + assert n_remplissage >= 0, f"T={T} trop court pour Q={Q} marqueurs + 2 jetons" + remplissage = torch.randint( + vocab.N_MARQUEURS, + vocab.taille_remplissage, + (n, n_remplissage), + generator=gen, + ) + + # Jetons QUESTION et REQUETE + questions = torch.tensor( + [vocab.jeton_question(k.item()) for k in q_idx], + dtype=torch.long, + ).unsqueeze(1) + requete = torch.full((n, 1), vocab.JETON_REQUETE) + + x = torch.cat([marqueurs, remplissage, questions, requete], dim=1) + + # Cible = marqueur q_idx[k] pour l'item k + y = marqueurs.gather(1, q_idx.unsqueeze(1)).squeeze(1) + + return Lot(x=x, y=y, q=q_idx) + + +# Pour le grain CoT, on ajoute des jetons de transition et de récapitulation. +# Choix : on insère 2q_idx+1 jetons dans la chaîne (q_idx QUESTION + q_idx JETON_PAS + 1 cible). +# Cela étend la séquence — T doit être recalculé dans la fonction appelante. +# JETON_PAS est un token dédié qui marque une transition (un pas logique). +_QUESTION_OFFSET = None # initialisé paresseusement ci-dessous + + +def _etendre_vocab_cot(vocab: Vocab, max_q: int) -> int: + """Étend le vocabulaire pour CoT : ajoute un JETON_PAS + max_q jetons d'index. + + Retourne la taille étendue. + """ + assert vocab.N_QUESTIONS <= max_q + 1, f"N_QUESTIONS={vocab.N_QUESTIONS} > max_q+1={max_q+1}" + return vocab.VOCAB + 1 + max_q # VOCAB + JETON_PAS + max_q jetons de recap + + +def lot_multi_hop_cot( + n: int, + T_cot: int, # longueur ciblee de la sequence, incluant la chaine CoT + gen: torch.Generator, + vocab: Vocab, + max_q: int = 3, # borne sup des q_idx (egal a vocab.N_QUESTIONS - 1 pour la v3) +) -> Lot: + """Génère un lot multi-sauts CoT : chaîne_question + chaîne_reponse + cible. + + Séquence : + [M1, M2, ..., M_Q, ..., remplissage, QUESTION(k), REQUETE, + JETON_PAS_0, QUESTION_0, JETON_PAS_1, QUESTION_1, ..., JETON_PAS_q, QUESTION_q, Mk] + + où Mk est la cible finale. Le modèle doit apprendre à générer la chaîne intermédiaire + (les q_idx QUESTION_j + leur transition) AVANT la cible finale. C'est précisément + le « CoT supervisé » de Huang et al. 2026. + + Tell c.1493 strict fondateur nuance : la chaîne est supervisée par entropie croisée + sur **chaque token de la chaîne intermédiaire** (pas seulement la cible finale). + + Note : l'évaluation (evaluer_multi_hop_cot) lit la cible au dernier token ET + mesure l'exactitude sur les tokens QUESTION_j de la chaîne. + """ + Q = vocab.N_QUESTIONS + assert Q >= 2, "lot_multi_hop_cot exige N_QUESTIONS >= 2" + VOCAB_ETENDU = _etendre_vocab_cot(vocab, max_q) + JETON_PAS = vocab.VOCAB # un seul jeton de transition + RECAP_OFFSET = vocab.VOCAB + 1 # jeton QUESTION_j de recap = RECAP_OFFSET + j + + marqueurs = torch.randint(0, vocab.N_MARQUEURS, (n, Q), generator=gen) + q_idx = torch.randint(0, Q, (n,), generator=gen) + + # Zone question/requete/CoT : QUESTION(k), REQUETE, puis pour j in 0..q_idx : PAS, RECAP_j, et enfin cible. + # Le nombre de tokens CoT = 1 (QUESTION) + 1 (REQUETE) + 2*q_idx + 1 (cible) = 2*q_idx + 3. + # T_cot doit accommoder Q marqueurs + 2*q_idx + 3 + zone remplissage. + + n_remplissage = T_cot - Q - 2 - 2 * max_q - 1 + assert n_remplissage >= 0, f"T_cot={T_cot} trop court pour Q={Q} + 2*max_q+1={2*max_q+1} + 3" + + remplissage = torch.randint( + vocab.N_MARQUEURS, + vocab.taille_remplissage, + (n, n_remplissage), + generator=gen, + ) + + questions = torch.tensor( + [vocab.jeton_question(k.item()) for k in q_idx], + dtype=torch.long, + ).unsqueeze(1) + requete = torch.full((n, 1), vocab.JETON_REQUETE, dtype=torch.long) + + # Construction de la chaîne CoT (taille fixe = 2*max_q + 1) + # Pour chaque item : PAS, RECAP_0, PAS, RECAP_1, ..., PAS, RECAP_q, Mk + chaineseq = torch.full((n, 2 * max_q + 1), JETON_PAS, dtype=torch.long) + for j in range(max_q): + chaineseq[:, 2 * j + 1] = RECAP_OFFSET + j # RECAP_j + # Dernier slot = cible + cibles = marqueurs.gather(1, q_idx.unsqueeze(1)).squeeze(1) # (n,) + chaineseq[:, -1] = cibles + + # Tronque la chaîne au-delà de q_idx+1 pas (plus court que max_q si q_idx < max_q) + # Mais pour la simplicite du conditionnement, on garde la chaîne complete et on masque la perte + # au-delà du q_idx+1 pas via un masque positionnel dans le calcul de perte CoT. + + x = torch.cat([marqueurs, remplissage, questions, requete, chaineseq], dim=1) + y = cibles + q = q_idx + + return Lot(x=x, y=y, q=q) + + +@torch.no_grad() +def evaluer_single_hop(modele, vocab: Vocab, T: int, n: int = 512, graine: int = 99) -> tuple[float, float]: + """Exactitude et perplexite à la position de requête (single-hop).""" + gen = torch.Generator().manual_seed(graine) + lot = lot_single_hop(n, T, gen, vocab) + logits = modele(lot.x)[:, -1] + perte = F.cross_entropy(logits, lot.y) + exactitude = (logits.argmax(-1) == lot.y).float().mean() + return exactitude.item(), math.exp(perte.item()) + + +@torch.no_grad() +def evaluer_multi_hop(modele, vocab: Vocab, T: int, n: int = 512, graine: int = 99) -> tuple[float, float]: + """Exactitude et perplexite à la position de requête (multi-sauts).""" + gen = torch.Generator().manual_seed(graine) + lot = lot_multi_hop(n, T, gen, vocab) + logits = modele(lot.x)[:, -1] + perte = F.cross_entropy(logits, lot.y) + exactitude = (logits.argmax(-1) == lot.y).float().mean() + return exactitude.item(), math.exp(perte.item()) + + +def entrainer( + modele, + vocab: Vocab, + T: int, + multi_hop: bool, + graine: int, + pas: int = 200, + batch: int = 32, + lr: float = 3e-3, +) -> tuple[float, float, float]: + """Entraîne un modèle sur la tâche choisie, renvoie (exactitude, perplexite, secondes). + + Tell c.1493 strict fondateur nuance — multi-seed : cette fonction utilise UNE graine. + La mesure multi-seed (≥4) est dans le grain de mesure suivant (TV-03 v2), portée par + :func:`entrainer_multi_seed`. + """ + import time + + torch.manual_seed(graine) + opt = torch.optim.Adam(modele.parameters(), lr=lr) + gen = torch.Generator().manual_seed(1234 + graine) + debut = time.perf_counter() + for _ in range(pas): + lot = lot_multi_hop(batch, T, gen, vocab) if multi_hop else lot_single_hop(batch, T, gen, vocab) + perte = F.cross_entropy(modele(lot.x)[:, -1], lot.y) + opt.zero_grad() + perte.backward() + opt.step() + secondes = time.perf_counter() - debut + if multi_hop: + acc, ppl = evaluer_multi_hop(modele, vocab, T) + else: + acc, ppl = evaluer_single_hop(modele, vocab, T) + return acc, ppl, secondes + + +def entrainer_multi_seed( + fabrique_modele, + vocab: Vocab, + T: int, + multi_hop: bool, + graines: list[int], + pas: int = 300, + batch: int = 32, + lr: float = 3e-3, +) -> dict: + """Mesure multi-seed (≥4) : moyenne, écart-type, secondes totales. + + Tell c.1493 strict fondateur nuance — la mesure multi-seed par graine est attendue par + le protocole PR review-discipline §C (≥4 graines parmi 0/1/7/42/99). On expose moyenne, + écart-type et liste brute pour permettre les vérifs edge≥2σ / DM cross-seed. + + :param fabrique_modele: callable ``(vocab) -> nn.Module`` qui crée un modèle vierge. + L'instance est recréée pour chaque graine — pas de contamination de l'initialisation. + :param graines: liste explicite d'identifiants de graine (par défaut ``[0, 1, 7, 42]``). + :return: dict avec ``acc_moy``, ``acc_std``, ``ppl_moy``, ``ppl_std``, ``secondes``, + ``brut`` (liste de tuples ``(graine, acc, ppl, sec)``). + """ + import time + import statistics + + if not graines: + raise ValueError("graines doit être non vide") + + debut_total = time.perf_counter() + brut = [] + for graine in graines: + modele = fabrique_modele(vocab) + acc, ppl, sec = entrainer( + modele, + vocab, + T=T, + multi_hop=multi_hop, + graine=graine, + pas=pas, + batch=batch, + lr=lr, + ) + brut.append((graine, acc, ppl, sec)) + secondes = time.perf_counter() - debut_total + + accs = [b[1] for b in brut] + ppls = [b[2] for b in brut] + return { + "acc_moy": statistics.fmean(accs), + "acc_std": statistics.pstdev(accs) if len(accs) > 1 else 0.0, + "ppl_moy": statistics.fmean(ppls), + "ppl_std": statistics.pstdev(ppls) if len(ppls) > 1 else 0.0, + "secondes": secondes, + "brut": brut, + "n_graines": len(graines), + } + + +@torch.no_grad() +def evaluer_multi_hop_cot(modele, vocab: Vocab, T_cot: int, n: int = 512, graine: int = 99, + max_q: int = 3) -> tuple[float, float]: + """Exactitude sur la cible finale en mode CoT (réponse à la question). + + La séquence inclut la chaîne CoT. On mesure : + - exactitude = la cible finale (dernier token) est-elle correcte ? + - perplexité moyenne sur les tokens de la chaîne (pas seulement la cible finale). + + Tell c.1493 strict fondateur nuance : on lit la sortie sur **toute la chaîne**, pas + uniquement la dernière position. L'exactitude cible reste la métrique principale + (réponse à la question = discriminant H.2). + """ + gen = torch.Generator().manual_seed(graine) + lot = lot_multi_hop_cot(n, T_cot, gen, vocab, max_q=max_q) + # Eval sur la dernière position entraînée (la cible est en position -1 dans x, le + # logit qui la prédit est donc en position -2 -- cohérent avec entrainer_cot qui + # couvre les logits [start_pred - 1, T - 1), soit jusqu'à T - 2 inclus). + # Lire -1 mesurait un logit qui prédit au-delà de la séquence, sans signal. + logits_cible = modele(lot.x)[:, -2] + perte_cible = F.cross_entropy(logits_cible, lot.y) + exactitude = (logits_cible.argmax(-1) == lot.y).float().mean() + return exactitude.item(), math.exp(perte_cible.item()) + + +def entrainer_cot( + modele, + vocab: Vocab, + T_cot: int, + graine: int, + max_q: int = 3, + pas: int = 300, + batch: int = 32, + lr: float = 3e-3, +) -> tuple[float, float, float]: + """Entraînement CoT supervisé sur la tâche multi-sauts. Renvoie (acc, ppl, sec). + + Tell c.1493 strict fondateur nuance : la perte supervise **la chaîne entière** + (chaque token PAS/RECAP_j/cible doit être prédit correctement). Le gradient + coule donc à travers toute la séquence de génération. + """ + import time + + Q = vocab.N_QUESTIONS + RECAP_OFFSET = vocab.VOCAB + 1 + + torch.manual_seed(graine) + opt = torch.optim.Adam(modele.parameters(), lr=lr) + gen = torch.Generator().manual_seed(1234 + graine) + debut = time.perf_counter() + for _ in range(pas): + lot = lot_multi_hop_cot(batch, T_cot, gen, vocab, max_q=max_q) + # Calcul de la perte sur toute la séquence + logits = modele(lot.x) # (B, T, VOCAB_ETENDU) + # Cible : prédire chaque token de la chaîne à partir du token précédent. + # À la position i on prédit lot.x[:, i+1]. Pour la chaîne (n_pred_positions tokens), + # on prédit les positions [start_pred, start_pred+n_pred_positions) à partir des + # logits en [start_pred-1, start_pred-1+n_pred_positions). Donc : + # - logits_pred = logits[:, start_pred - 1 : start_pred - 1 + n_pred_positions, :] + # - cible_shift = lot.x[:, start_pred : start_pred + n_pred_positions] + B, T_seq, V = logits.shape + n_pred_positions = 2 * max_q + 1 + start_pred = T_seq - n_pred_positions + cible_shift = lot.x[:, start_pred : start_pred + n_pred_positions] + logits_pred = logits[:, start_pred - 1 : start_pred - 1 + n_pred_positions, :] + # Masque : on inclut dans la perte les positions j qui produisent un token + # pertinent (PAS_j, RECAP_j pour j <= q_idx) ET le slot cible. + # Etat anterieur v1 : `pas_positions <= q` seul -- la cible (j = 2*max_q, + # pas_positions[j] = max_q > q pour tout q <= max_q-1) n'etait jamais + # couverte -> exactitude 0. + # Etat anterieur v2 : OR sur j = 2*q+1 -- faux : la chaine est de taille + # FIXE (cf. lot_multi_hop_cot), le slot 2*q+1 porte RECAP_q, jamais la + # cible ; celle-ci est TOUJOURS en fin de chaine (j = 2*max_q) quel que + # soit q_idx -> exactitude 0 a nouveau (mesure : acc 0.0000, ppl ~29856). + q = lot.q # (B,) + j_positions = torch.arange(n_pred_positions) # (n_pred_positions,) + # Chaque position j correspond au pas floor(j/2) (0-indexed) + pas_positions = j_positions // 2 # (n_pred_positions,) + # Le slot cible (j = 2*max_q = n_pred_positions - 1) est inclus via cette OR. + cible_j = torch.full((q.shape[0], 1), 2 * max_q, dtype=torch.long) # (B, 1) + est_slot_cible = (j_positions.unsqueeze(0) == cible_j) # (B, n_pred_positions) + masque = (pas_positions.unsqueeze(0) <= q.unsqueeze(1)) | est_slot_cible # (B, n_pred_positions) + # Perte par token, masquée + perte_full = F.cross_entropy( + logits_pred.reshape(-1, V), + cible_shift.reshape(-1), + reduction='none', + ).view(B, n_pred_positions) + perte_masquee = (perte_full * masque.float()).sum() / masque.sum().clamp(min=1) + opt.zero_grad() + perte_masquee.backward() + opt.step() + secondes = time.perf_counter() - debut + acc, ppl = evaluer_multi_hop_cot(modele, vocab, T_cot, max_q=max_q) + return acc, ppl, secondes + + +def entrainer_multi_seed_cot( + fabrique_modele, + vocab: Vocab, + T_cot: int, + graines: list[int], + max_q: int = 3, + pas: int = 300, + batch: int = 32, + lr: float = 3e-3, +) -> dict: + """Mesure multi-seed (≥4) CoT : moyenne, écart-type, secondes totales.""" + import time + import statistics + + if not graines: + raise ValueError("graines doit être non vide") + + debut_total = time.perf_counter() + brut = [] + for graine in graines: + modele = fabrique_modele(vocab) + acc, ppl, sec = entrainer_cot( + modele, + vocab, + T_cot=T_cot, + graine=graine, + max_q=max_q, + pas=pas, + batch=batch, + lr=lr, + ) + brut.append((graine, acc, ppl, sec)) + secondes = time.perf_counter() - debut_total + + accs = [b[1] for b in brut] + ppls = [b[2] for b in brut] + return { + "acc_moy": statistics.fmean(accs), + "acc_std": statistics.pstdev(accs) if len(accs) > 1 else 0.0, + "ppl_moy": statistics.fmean(ppls), + "ppl_std": statistics.pstdev(ppls) if len(ppls) > 1 else 0.0, + "secondes": secondes, + "brut": brut, + "n_graines": len(graines), + }