Module 09 · Leçon 5

L'entraîner : apprendre au modèle à écrire des poèmes Tang

Entraîner le GPT de la leçon précédente sur trente-quatre mille poèmes Tang, sept minutes sur le CPU d'un portable, et le voir passer du charabia aux vers de cinq et sept caractères. Puis ne lui donner que trois cents poèmes, et voir comment il les apprend par cœur au lieu d'apprendre à écrire des poèmes.

  • Environ 50 minutes
  • Niveau : Approfondi
  • Testé : 2026-09-15 torch 2.14, CPU Apple M4, graine aléatoire fixée

Le code et les sorties des programmes sont reproduits tels qu’ils ont tourné : commentaires et sorties sont donc en chinois.

Le modèle est construit ; entraînons-le maintenant. Le code de cette leçon, train.py, ne fait qu'un peu plus de cent lignes, avec la même structure que la boucle d'entraînement du module 08 : prendre un lot de données, calculer la perte, rétropropager, mettre à jour les paramètres.

python train.py

Les données : décalées d'une position

Les données d'entraînement d'un modèle de langue n'ont pas besoin d'annotation humaine. Un texte est à lui seul question et réponse :

def get_batch(ids, generator=None):
    """随机截取 BATCH 段长度为 block_size 的文字。目标 y 就是 x 往后错一位:每个位置都要预测下一个字。"""
    starts = torch.randint(len(ids) - cfg.block_size - 1, (BATCH,), generator=generator)
    x = torch.stack([ids[s:s + cfg.block_size] for s in starts])
    y = torch.stack([ids[s + 1:s + cfg.block_size + 1] for s in starts])
    return x, y

Tous les poèmes sont reliés par des retours à la ligne en une longue chaîne, dont on extrait au hasard 32 segments de 128 caractères comme entrée x ; chaque segment décalé d'une position donne la réponse y. Si l'entrée est « 白日依山尽, », la réponse est « 日依山尽,黄 » : la 1re position voit « 白 » et doit deviner « 日 », la 2e voit « 白日 » et doit deviner « 依 ». La leçon précédente l'a dit : le masque causal garantit qu'aucune position ne peut regarder la réponse.

Le retour à la ligne a ici un sens particulier : il marque la fin d'un poème et le début du suivant. Le modèle apprend que « le point est suivi d'un retour à la ligne », et à la génération, il peut s'arrêter à un retour à la ligne : le poème est terminé.

Les 1000 derniers poèmes sont gardés comme jeu de validation, jamais vus à l'entraînement. On peut ainsi, avec la méthode de la leçon 6 du module 08, voir si le modèle apprend ou s'il récite.

训练集 34135 首诗,1547255 个词元;验证集 1000 首;词表 6289 个字符
模型:4 层,4 个头,向量维度 128,共 1,614,720 个参数

Quelques détails de la boucle d'entraînement

optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=0.1)
for step in range(1, args.steps + 1):
    for group in optimizer.param_groups:
        group["lr"] = lr_at(step - 1)
    x, y = get_batch(train_ids)
    _, loss = model(x, y)
    optimizer.zero_grad()
    loss.backward()
    torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)  # 梯度裁剪:防止偶尔一步梯度太大把训练带飞
    optimizer.step()

Par rapport au module 08, deux choses s'ajoutent.

Le taux d'apprentissage varie. Sur les 100 premiers pas, il monte doucement de 0 à 0,001 : c'est l'échauffement (warm-up). Au début de l'entraînement, les paramètres sont aléatoires et la direction des gradients peu fiable ; des pas trop grands dévieraient facilement. Ensuite, le taux descend lentement jusqu'à 0,0001 selon une courbe cosinus : en fin d'entraînement, des pas plus petits permettent de trouver plus finement les zones de faible perte.

def lr_at(step, peak=1e-3, warmup=100):
    """学习率:前 100 步从 0 慢慢升上去(预热),然后按余弦曲线降到峰值的十分之一。"""
    if step < warmup:
        return peak * (step + 1) / warmup
    progress = (step - warmup) / max(1, args.steps - warmup)
    return peak * (0.1 + 0.9 * 0.5 * (1 + math.cos(math.pi * progress)))

L'écrêtage du gradient. Si le gradient calculé à un pas est particulièrement grand (en tombant sur un lot de données inhabituel), on le réduit proportionnellement pour que sa norme ne dépasse pas 1. La leçon 2 du module 08 l'a montré : un pas trop grand peut faire directement diverger l'entraînement.

« Échauffement + décroissance en cosinus + écrêtage du gradient + AdamW » est la combinaison la plus courante pour entraîner un Transformer ; les grands modèles s'entraînent essentiellement ainsi, avec d'autres valeurs.

Le voir apprendre à écrire des poèmes

3000 pas d'entraînement ; de temps en temps, on regarde la perte et on lui fait écrire deux poèmes :

训练前:训练损失 8.779,验证损失 8.778(随便猜的话是 ln(6289) = 8.747)
  训练前随便写的: 耦呌同捐睢寮办猬囊斵樬劫鞑顷饼遁瞩赴冈譀嗾雅踯牧写禧氛迹樯毁尤平螵蔻郓祐菼灺嵊匦

第 300 步(27 秒):训练损失 5.778,验证损失 5.818
  将南。
  门不独。到李斜,功觅青。谁将风花不可碧,不和山。

第 1000 步(97 秒):训练损失 4.777,验证损失 4.879
  山里怅望两间道,千里长亭寺断肠。曾梳白兰草,试向玉关鱼。
  此日醉前年,相思高至兹。云如汉陵子,暮水九重宫。

第 2000 步(249 秒):训练损失 4.410,验证损失 4.583
  秋江野鸟过高楼,野树猿声怨见人。处处秋风满孤照,沧海无端行处闻。
  春山度滟月,多少漫为秋。药罢已开葬,松阴不道开。

第 3000 步(407 秒):训练损失 4.280,验证损失 4.488
  江山古馆响幽幽,山鸟无人见白头。此时猩猩争得语,又将杯酒醉参差。
  分明人在泪,明月更经过。小谷闲烟树,红潭古石床。归来扶白首,立向卧青山。独有安行处,如何却得还。

训练用了 408 秒,模型存到 .cache/gpt.pt

Sur mon ordinateur (Apple M4, CPU seulement), 3000 pas ont pris moins de 7 minutes. Un pas traite 32×128 = 4096 caractères ; 3000 pas voient environ 12 millions de caractères, soit à peu près 8 passages sur le jeu d'entraînement.

Ce qu'il apprend apparaît par étapes :

  • Avant l'entraînement : une suite de caractères aléatoires, perte 8,78, le pur hasard.
  • 300 pas : il a appris la ponctuation, n'écrit que des caractères courants, mais les vers sont de longueur variable.
  • 1000 pas : la plupart des vers ont cinq ou sept caractères, virgules et points alternent. Mais dans le premier poème, les deux premiers vers ont sept caractères et les deux derniers cinq : il ne sait pas encore qu'un poème doit être homogène.
  • 2000 pas : le format est à peu près juste, et apparaissent des vers aux images cohérentes comme « 秋江野鸟过高楼,野树猿声怨见人 » (sur le fleuve d'automne, les oiseaux sauvages passent la haute tour ; dans les arbres sauvages, les cris des singes se plaignent des hommes).
  • 3000 pas : le premier est un jueju complet à sept caractères, « 江山古馆响幽幽,山鸟无人见白头 » ; le second un lüshi complet à cinq caractères, huit vers.

Ce qu'il n'a pas appris est tout aussi net : le sens se suit souvent mal (« 此时猩猩争得语 », en ce moment les orangs-outans se disputent la parole), et il ne se soucie ni des tons ni de la rime. Pour 1,61 million de paramètres et 7 minutes d'entraînement, c'est déjà bien.

Pendant tout l'entraînement, pertes d'entraînement et de validation baissent et restent très proches (4,28 et 4,49 à la fin). Selon la leçon 6 du module 08, cela signifie pas de surapprentissage marqué : il s'en sort presque aussi bien sur des poèmes jamais vus.

Écrit-il des poèmes, ou les récite-t-il ?

Les modules précédents de ce cours l'ont répété : qu'un texte généré ait l'air bon ne prouve pas que le modèle a vraiment appris ; il a peut-être simplement appris par cœur les données d'entraînement. À la fin de l'entraînement, le script fait donc une vérification : générer 100 poèmes et compter combien de vers sont identiques à un vers du jeu d'entraînement.

train_sentences = {s for p in train_poems for s in p.replace("。", ",").split(",") if s}
torch.manual_seed(1)
generated = sample(100)
sentences = [s for p in generated for s in p.replace("。", ",").split(",") if s]
copied = sum(s in train_sentences for s in sentences)
生成 100 首诗,共 576 句,其中 0 句(0%)和训练集里的某一句一模一样
整首和训练集里某一首一模一样的:0 首

576 vers, pas un seul recopié. Ce qu'il a appris, c'est la « manière d'écrire » des poèmes Tang : le format, les caractères et mots courants, quels caractères vont souvent ensemble, qu'il combine en vers nouveaux.

Seulement 300 poèmes

La leçon 6 du module 08 l'a dit : avec trop peu de données, le modèle récite. Essayons maintenant le même modèle avec seulement 300 poèmes :

python train.py --poems 300 --steps 1500 --out small.pt
训练集 300 首诗,13612 个词元;验证集 1000 首;词表 6289 个字符

第 150 步(22 秒):训练损失 5.079,验证损失 6.584
第 500 步(70 秒):训练损失 0.421,验证损失 8.864
第 1000 步(145 秒):训练损失 0.103,验证损失 9.817
  雕鹗途程在碧天,彩衣东去复何言。二千宾客旧知己,十二山河新故园。吟看桂生溪月上,醉听鲲化海涛翻。好期圣代重相见,莫学袁生老竹轩。
第 1500 步(230 秒):训练损失 0.071,验证损失 9.984

生成 100 首诗,共 593 句,其中 454 句(77%)和训练集里的某一句一模一样
整首和训练集里某一首一模一样的:40 首

La perte d'entraînement descend à 0,07, bien plus bas que les 4,28 avec toutes les données. Mais la perte de validation monte jusqu'à 9,98, plus haut encore que les 8,78 du hasard d'avant l'entraînement.

Le lüshi à sept caractères écrit au pas 1000 est d'une belle régularité, parce qu'il est récité mot pour mot depuis le jeu d'entraînement. La vérification finale le confirme : 77 % des vers sont recopiés, et 40 poèmes sur 100 le sont en entier.

1,61 million de paramètres pour seulement 13 000 caractères à apprendre : le modèle est parfaitement capable de tout mémoriser. Mémoriser est le moyen le plus commode de faire baisser la perte d'entraînement. Sur des poèmes jamais vus, il fait même pire que s'il n'avait rien appris : trop « sûr de lui » sur ce qu'il a appris par cœur, il donne à la bonne réponse une probabilité dérisoire face à un autre poème.

Même modèle, même code, seules les données passent de 34 000 à 300 poèmes, et le résultat passe de « sait écrire des poèmes » à « récite des poèmes ». C'est la conclusion de la leçon 6 du module 08 rejouée sur un modèle de langue, et la raison pour laquelle les grands modèles ont besoin de données massives.

Que signifie une perte de 4,49

La perte de validation finale vaut 4,49. L'entropie croisée vaut -log(probabilité du bon caractère) : en moyenne, le modèle donne donc au bon caractère suivant une probabilité d'environ e puissance -4,49, soit environ 1,1 %.

Cela paraît faible, mais il faut savoir qu'en poésie, le caractère suivant admet de toute façon beaucoup de choix raisonnables. Après « 白日依山尽,黄河入海 », le modèle répartit un peu de probabilité sur « 无 », « 间 », « 多 », « 深 » (la leçon suivante montre les chiffres exacts). Le poème original dit « 流 » (coule), mais un autre caractère n'est pas forcément faux. La perte ne peut pas descendre à 0, pour la même raison qu'à la leçon 1 du module 08 : « les données fluctuent, la perte ne descend jamais à 0 ».

Un repère est la comparaison avec le hasard : deviner au hasard, c'est choisir uniformément parmi 6289 caractères, perte 8,75. 4,49 revient à réduire le choix de 6289 caractères à environ e puissance 4,49, soit environ 89 caractères.

Exercices

  1. Passez --steps à 6000. La perte continue-t-elle de baisser ? Les poèmes générés s'améliorent-ils ?
  2. Fixez le nombre de poèmes d'entraînement à 1000, 3000 et 10 000 (--poems), entraînez chacun 1500 pas, notez la perte de validation et la proportion de « vers recopiés », et dressez-en un tableau.
  3. Retirez l'échauffement et la décroissance du taux d'apprentissage (lr_at renvoie directement 0,001), entraînez aussi 3000 pas, et comparez la perte de validation finale.

Auto-test

1. Pour entraîner un modèle de langue, que sont l'entrée et la réponse ?

On extrait un segment de texte comme entrée, et ce même segment décalé d'une position sert de réponse. Chaque position doit prédire le caractère suivant d'après elle-même et les caractères précédents. Un texte de 128 caractères fournit à la fois 128 questions de prédiction.

2. Quels problèmes résolvent respectivement l'échauffement du taux d'apprentissage et l'écrêtage du gradient ?

L'échauffement : au début de l'entraînement, les paramètres sont aléatoires et la direction des gradients peu fiable ; un grand taux d'emblée dévie facilement, on commence donc petit et on monte peu à peu. L'écrêtage : parfois le gradient d'un pas est particulièrement grand ; mettre à jour directement avec lui ferait sauter les paramètres trop loin, et l'entraînement risquerait de diverger ; on réduit donc proportionnellement les gradients trop grands.

3. Avec seulement 300 poèmes, pourquoi la perte d'entraînement est-elle très basse, et la perte de validation plus haute que le hasard ?

Le modèle a assez de paramètres pour apprendre par cœur les 300 poèmes, d'où une perte d'entraînement très basse. Mais ce qu'il a appris, ce sont ces 300 poèmes eux-mêmes, et non les régularités de l'écriture poétique. Très sûr de lui sur ce qu'il a mémorisé, il donne à la bonne réponse une probabilité très faible face à un poème jamais vu, et l'entropie croisée dépasse celle du hasard. Que 77 % des vers générés soient recopiés le confirme.

Questions et discussion

Bloqué sur cette leçon ? Posez votre question ici. Et si vous pouvez répondre à quelqu'un, n'hésitez pas.

Une question rapporte 3 points, une réponse 6. Les messages paraissent après vérification.

Chargement de la discussion…