round5: 全前向训练: head / head+lora / head+block + 方向①②开关
Browse files- train_ponder_head.py +775 -0
train_ponder_head.py
ADDED
|
@@ -0,0 +1,775 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# -*- coding: utf-8 -*-
|
| 2 |
+
"""
|
| 3 |
+
train_ponder_head.py — 微调 PonderLlamaForCausalLM 的 halting 头 (真 PonderNet 训练)
|
| 4 |
+
|
| 5 |
+
三种训练模式 (--train-mode):
|
| 6 |
+
head 只训 ponder_head (2049 参数), 主干全部冻结。最便宜, 单卡 8GB 也能跑;
|
| 7 |
+
但"多想几步"带来的增益有限 (思考块没学过被重复执行)。
|
| 8 |
+
head+lora ponder_head + LoRA (--lora-scope ponder 只注入思考块 / all 注入全部层)。
|
| 9 |
+
推荐: 让思考块学会"被重复执行", 完整 PonderNet 语义, 显存友好的代价。
|
| 10 |
+
head+block ponder_head + 思考块全参微调。效果上限最高, 显存/算力开销最大。
|
| 11 |
+
|
| 12 |
+
数据格式 (jsonl, 每行一条):
|
| 13 |
+
{"messages": [{"role": "user", "content": "..."}, {"role": "assistant", "content": "..."}]}
|
| 14 |
+
或 {"text": "一段完整文本"} (用 --text-field text)
|
| 15 |
+
可用 --make-demo-data demo.jsonl 生成 16 条中/英混合演示数据 (简单问答 + 推理题,
|
| 16 |
+
让 halting 头有机会学到"难题多想、简单题少想")。
|
| 17 |
+
|
| 18 |
+
日志中的关键指标是 avg_ponder_steps: 训练初期接近 max_steps, 随 CE+KL 联合作用,
|
| 19 |
+
简单题的步数应逐步下降 —— 这就是模型在学"该想几步"。
|
| 20 |
+
|
| 21 |
+
示例:
|
| 22 |
+
python train_ponder_head.py --model /path/to/MiniCPM5-2B-cpu \
|
| 23 |
+
--train-mode head+lora --lora-scope ponder --data data.jsonl \
|
| 24 |
+
--epochs 2 --output out/ponder-lora
|
| 25 |
+
"""
|
| 26 |
+
from __future__ import annotations
|
| 27 |
+
|
| 28 |
+
import argparse
|
| 29 |
+
import json
|
| 30 |
+
import math
|
| 31 |
+
import os
|
| 32 |
+
import random
|
| 33 |
+
import time
|
| 34 |
+
|
| 35 |
+
import torch
|
| 36 |
+
import torch.nn.functional as F
|
| 37 |
+
from torch.utils.data import DataLoader, Dataset
|
| 38 |
+
|
| 39 |
+
from transformers import AutoTokenizer, get_cosine_schedule_with_warmup
|
| 40 |
+
|
| 41 |
+
from ponder_llama import PonderLlamaConfig, PonderLlamaForCausalLM, LlamaConfig
|
| 42 |
+
|
| 43 |
+
|
| 44 |
+
# ========================================================================== #
|
| 45 |
+
# 数据
|
| 46 |
+
# ========================================================================== #
|
| 47 |
+
DEMO_DATA = [
|
| 48 |
+
{"messages": [{"role": "user", "content": "你好"}, {"role": "assistant", "content": "你好!很高兴见到你。"}]},
|
| 49 |
+
{"messages": [{"role": "user", "content": "1+1等于几?"}, {"role": "assistant", "content": "1+1=2。"}]},
|
| 50 |
+
{"messages": [{"role": "user", "content": "今天星期几"}, {"role": "assistant", "content": "抱歉,我无法获取实时日期。"}]},
|
| 51 |
+
{"messages": [{"role": "user", "content": "把'苹果'翻译成英文"}, {"role": "assistant", "content": "苹果的英文是 apple。"}]},
|
| 52 |
+
{"messages": [{"role": "user", "content": "天空是什么颜色"}, {"role": "assistant", "content": "晴天时天空通常是蓝色的。"}]},
|
| 53 |
+
{"messages": [{"role": "user", "content": "水在多少度结冰?"}, {"role": "assistant", "content": "标准大气压下,水在 0 摄氏度结冰。"}]},
|
| 54 |
+
{"messages": [{"role": "user", "content": "9.11 和 9.9 哪个大?"}, {"role": "assistant", "content": "9.9 更大。比较小数:9.11 = 9 + 0.11,9.9 = 9 + 0.90,因为 0.90 > 0.11,所以 9.9 > 9.11。"}]},
|
| 55 |
+
{"messages": [{"role": "user", "content": "一个水池甲管8小时注满,乙管12小时注满,两管同开几小时注满?"}, {"role": "assistant", "content": "设总量为1。甲每小时注 1/8,乙每小时注 1/12,合计 1/8+1/12 = 3/24+2/24 = 5/24。注满需要 24/5 = 4.8 小时。"}]},
|
| 56 |
+
{"messages": [{"role": "user", "content": "鸡兔同笼,头36,脚100,各几只?"}, {"role": "assistant", "content": "设鸡 x 只,兔 36-x 只。脚数 2x + 4(36-x) = 100,得 144-2x = 100,x = 22。鸡 22 只,兔 14 只。"}]},
|
| 57 |
+
{"messages": [{"role": "user", "content": "证明 sqrt(2) 是无理数"}, {"role": "assistant", "content": "反证法:若 sqrt(2)=p/q(最简分数),则 p^2=2q^2,故 p 为偶数,设 p=2k,则 4k^2=2q^2,q^2=2k^2,q 也为偶数,与 p/q 最简矛盾。故 sqrt(2) 无理。"}]},
|
| 58 |
+
{"messages": [{"role": "user", "content": "写一个Python函数判断质数"}, {"role": "assistant", "content": "def is_prime(n):\n if n < 2: return False\n for i in range(2, int(n**0.5)+1):\n if n % i == 0: return False\n return True"}]},
|
| 59 |
+
{"messages": [{"role": "user", "content": "如果所有的A都是B,所有的B都是C,那么所有的A都是C吗?"}, {"role": "assistant", "content": "是的。这是三段论的传递性:A⊆B 且 B⊆C,则 A⊆C。"}]},
|
| 60 |
+
{"messages": [{"role": "user", "content": "给我讲个笑话"}, {"role": "assistant", "content": "程序员最讨厌的两件事:别人不写注释,和别人让自己写注释。"}]},
|
| 61 |
+
{"messages": [{"role": "user", "content": "一张纸对折10次后有多少层?"}, {"role": "assistant", "content": "每对折一次层数翻倍,对折 10 次后是 2^10 = 1024 层。"}]},
|
| 62 |
+
{"messages": [{"role": "user", "content": "用'但是'造句"}, {"role": "assistant", "content": "他很想去看电影,但是作业还没写完。"}]},
|
| 63 |
+
{"messages": [{"role": "user", "content": "5个人互相握手一次,共握手几次?"}, {"role": "assistant", "content": "这是组合问题:C(5,2) = 5×4/2 = 10 次。"}]},
|
| 64 |
+
]
|
| 65 |
+
|
| 66 |
+
|
| 67 |
+
def make_demo_data(path: str):
|
| 68 |
+
with open(path, "w", encoding="utf-8") as f:
|
| 69 |
+
for item in DEMO_DATA:
|
| 70 |
+
f.write(json.dumps(item, ensure_ascii=False) + "\n")
|
| 71 |
+
print(f"演示数据已写入 {path} ({len(DEMO_DATA)} 条)")
|
| 72 |
+
|
| 73 |
+
|
| 74 |
+
# ========================================================================== #
|
| 75 |
+
# 难度分级数据生成 (easy / medium / hard)
|
| 76 |
+
# 设计目的: PonderNet 的核心卖点是"难题多想、简单题少想"。只有当数据存在清晰
|
| 77 |
+
# 的难度梯度时, halting 头才有信号可学 —— 训练后各难度桶的平均思考步数应当分化。
|
| 78 |
+
# 全部题目反向构造 (先选解再出题), 保证答案正确; "difficulty" 字段供 eval 分层统计。
|
| 79 |
+
# ========================================================================== #
|
| 80 |
+
EASY_GREET = [
|
| 81 |
+
("你好", "你好!很高兴见到你,有什么可以帮你的吗?"),
|
| 82 |
+
("你好呀", "你好呀!今天过得怎么样?"),
|
| 83 |
+
("早上好", "早上好!祝你今天心情愉快。"),
|
| 84 |
+
("晚安", "晚安,祝你做个好梦。"),
|
| 85 |
+
("谢谢", "不客气,能帮到你我很开心。"),
|
| 86 |
+
("你是谁", "我是一个 AI 助手,可以回答问题、陪你聊天。"),
|
| 87 |
+
("你会什么", "我可以聊天、答题、算数,也可以帮你翻译。"),
|
| 88 |
+
("再见", "再见!欢迎随时来找我聊天。"),
|
| 89 |
+
("你叫什么名字", "你可以叫我小助手,很高兴认识你。"),
|
| 90 |
+
("今天天气不错", "是啊,这样的天气很适合出门走走。"),
|
| 91 |
+
("讲个笑话", "程序员最讨厌两件事:别人不写注释,和让自己写注释。"),
|
| 92 |
+
("你吃饭了吗", "我是 AI 不用吃饭,不过谢谢你关心!"),
|
| 93 |
+
("中午好", "中午好!记得吃午饭休息一下哦。"),
|
| 94 |
+
("在吗", "在的,请说。"),
|
| 95 |
+
("谢谢你", "不用谢,这是我应该做的。"),
|
| 96 |
+
("你会说英文吗", "会的,你可以用英文和我聊天。"),
|
| 97 |
+
("新年快乐", "新年快乐!祝万事如意。"),
|
| 98 |
+
("生日快乐", "谢谢!也祝你天天开心。"),
|
| 99 |
+
("周末愉快", "周末愉快!好好放松一下吧。"),
|
| 100 |
+
("下午好", "下午好!喝杯茶休息一下吧。"),
|
| 101 |
+
]
|
| 102 |
+
COMMON_SENSE = [
|
| 103 |
+
("天空是什么颜色", "晴天时天空通常是蓝色的。"),
|
| 104 |
+
("一年有几个月", "一年有 12 个月。"),
|
| 105 |
+
("一个星期有几天", "一个星期有 7 天。"),
|
| 106 |
+
("冰遇到热会怎么样", "冰遇到热会融化成水。"),
|
| 107 |
+
("太阳从哪边升起", "太阳从东边升起,西边落下。"),
|
| 108 |
+
("中国的首都是哪里", "中国的首都是北京。"),
|
| 109 |
+
("一年有哪四个季节", "一年有春、夏、秋、冬四个季节。"),
|
| 110 |
+
("彩虹有几种颜色", "彩虹有红橙黄绿蓝靛紫 7 种颜色。"),
|
| 111 |
+
("水在多少度结冰", "标准大气压下,水在 0 摄氏度结冰。"),
|
| 112 |
+
("水在多少度沸腾", "标准大气压下,水在 100 摄氏度沸腾。"),
|
| 113 |
+
("世界上最大的海洋是哪个", "世界上最大的海洋是太平洋。"),
|
| 114 |
+
("熊猫最爱吃什么", "熊猫最爱吃竹子。"),
|
| 115 |
+
("蜜蜂酿什么", "蜜蜂采花蜜酿成蜂蜜。"),
|
| 116 |
+
("植物生长需要什么", "植物生长一般需要阳光、水分和空气。"),
|
| 117 |
+
("十二生肖排第一的是什么", "十二生肖排第一的是鼠。"),
|
| 118 |
+
("中国最长的河流是哪条", "中国最长的河流是长江。"),
|
| 119 |
+
("世界上最高的山峰是哪座", "世界上最高的山峰是珠穆朗玛峰。"),
|
| 120 |
+
("人的正常体温大约是多少", "人的正常体温大约是 37 摄氏度。"),
|
| 121 |
+
("一天有多少个小时", "一天有 24 个小时。"),
|
| 122 |
+
("一公里等于多少米", "一公里等于 1000 米。"),
|
| 123 |
+
]
|
| 124 |
+
TRANSLATE = [
|
| 125 |
+
("苹果", "apple"), ("书", "book"), ("猫", "cat"), ("狗", "dog"),
|
| 126 |
+
("太阳", "sun"), ("月亮", "moon"), ("水", "water"), ("火", "fire"),
|
| 127 |
+
("山", "mountain"), ("河", "river"), ("树", "tree"), ("花", "flower"),
|
| 128 |
+
("鸟", "bird"), ("鱼", "fish"), ("学校", "school"), ("老师", "teacher"),
|
| 129 |
+
("学生", "student"), ("医生", "doctor"), ("医院", "hospital"), ("汽车", "car"),
|
| 130 |
+
("火车", "train"), ("飞机", "airplane"), ("手机", "phone"), ("电脑", "computer"),
|
| 131 |
+
("蓝色", "blue"), ("红色", "red"), ("绿色", "green"), ("白色", "white"),
|
| 132 |
+
("星期一", "Monday"), ("星期天", "Sunday"), ("早上", "morning"), ("晚上", "evening"),
|
| 133 |
+
("米饭", "rice"), ("面包", "bread"), ("鸡蛋", "egg"), ("牛奶", "milk"),
|
| 134 |
+
("茶", "tea"), ("咖啡", "coffee"), ("朋友", "friend"), ("家庭", "family"),
|
| 135 |
+
]
|
| 136 |
+
ANTONYMS = [
|
| 137 |
+
("大", "小"), ("多", "少"), ("高", "矮"), ("长", "短"), ("快", "慢"),
|
| 138 |
+
("上", "下"), ("左", "右"), ("前", "后"), ("开", "关"), ("好", "坏"),
|
| 139 |
+
("新", "旧"), ("冷", "热"), ("黑", "白"), ("买", "卖"), ("来", "去"),
|
| 140 |
+
("远", "近"), ("深", "浅"), ("宽", "窄"), ("厚", "薄"), ("轻", "重"),
|
| 141 |
+
("满", "空"), ("强", "弱"), ("甜", "苦"), ("明亮", "黑暗"), ("开始", "结束"),
|
| 142 |
+
("���功", "失败"), ("勇敢", "胆怯"), ("勤劳", "懒惰"), ("干净", "肮脏"), ("整齐", "杂乱"),
|
| 143 |
+
]
|
| 144 |
+
# 工程问题预解: (a, b, t) 满足 1/a + 1/b = 1/t, 即 (a-t)(b-t) = t^2
|
| 145 |
+
WORK_PROBLEMS = [(3, 6, 2), (4, 4, 2), (4, 12, 3), (6, 6, 3), (5, 20, 4),
|
| 146 |
+
(6, 12, 4), (8, 8, 4), (6, 30, 5), (10, 10, 5), (7, 42, 6),
|
| 147 |
+
(8, 24, 6), (9, 18, 6), (10, 15, 6), (12, 12, 6)]
|
| 148 |
+
|
| 149 |
+
|
| 150 |
+
def _gen_easy(rng):
|
| 151 |
+
kind = rng.randrange(4)
|
| 152 |
+
if kind == 0:
|
| 153 |
+
q, a = rng.choice(EASY_GREET)
|
| 154 |
+
elif kind == 1:
|
| 155 |
+
q, a = rng.choice(COMMON_SENSE)
|
| 156 |
+
elif kind == 2:
|
| 157 |
+
zh, en = rng.choice(TRANSLATE)
|
| 158 |
+
if rng.random() < 0.5:
|
| 159 |
+
q, a = f"把'{zh}'翻译成英文", f"{zh} 的英文是 {en}。"
|
| 160 |
+
else:
|
| 161 |
+
q, a = f"把'{en}'翻译成中文", f"{en} 的中文是 {zh}。"
|
| 162 |
+
else:
|
| 163 |
+
op = rng.choice(["+", "-"])
|
| 164 |
+
x, y = rng.randint(1, 9), rng.randint(1, 9)
|
| 165 |
+
if op == "-":
|
| 166 |
+
x, y = max(x, y), min(x, y)
|
| 167 |
+
res = x + y if op == "+" else x - y
|
| 168 |
+
q, a = f"{x} {op} {y} 等于几?", f"{x} {op} {y} = {res}。"
|
| 169 |
+
return q, a
|
| 170 |
+
|
| 171 |
+
|
| 172 |
+
def _gen_medium(rng):
|
| 173 |
+
kind = rng.randrange(5)
|
| 174 |
+
if kind == 0: # 两位数乘一位数 / 三位数加减
|
| 175 |
+
if rng.random() < 0.5:
|
| 176 |
+
a, b = rng.randint(10, 99), rng.randint(2, 9)
|
| 177 |
+
q, a_ = f"{a} 乘 {b} 等于多少?", f"{a} × {b} = {a * b}。"
|
| 178 |
+
else:
|
| 179 |
+
a, b = rng.randint(100, 999), rng.randint(100, 999)
|
| 180 |
+
if rng.random() < 0.5:
|
| 181 |
+
q, a_ = f"{a} 加 {b} 等于多少?", f"{a} + {b} = {a + b}。"
|
| 182 |
+
else:
|
| 183 |
+
hi, lo = max(a, b), min(a, b)
|
| 184 |
+
q, a_ = f"{hi} 减 {lo} 等于多少?", f"{hi} - {lo} = {hi - lo}。"
|
| 185 |
+
elif kind == 1: # 单位换算
|
| 186 |
+
table = [("米", "厘米", 100), ("米", "毫米", 1000), ("千克", "克", 1000),
|
| 187 |
+
("元", "角", 10), ("小时", "分钟", 60), ("分", "秒", 60)]
|
| 188 |
+
big, small, rate = rng.choice(table)
|
| 189 |
+
v = rng.randint(2, 99)
|
| 190 |
+
if rng.random() < 0.5:
|
| 191 |
+
q, a_ = f"{v} {big} 等于多少 {small}?", f"{v} {big} = {v * rate} {small}。"
|
| 192 |
+
else:
|
| 193 |
+
q, a_ = (f"{v * rate} {small} 等于多少 {big}?",
|
| 194 |
+
f"{v * rate} {small} = {v} {big}。")
|
| 195 |
+
elif kind == 2: # 单价×数量
|
| 196 |
+
p, n = rng.randint(2, 20), rng.randint(2, 12)
|
| 197 |
+
q, a_ = (f"一支笔 {p} 元,买 {n} 支一共需要多少钱?",
|
| 198 |
+
f"单价 {p} 元,数量 {n} 支,一共 {p} × {n} = {p * n} 元。")
|
| 199 |
+
elif kind == 3: # 速度×时间
|
| 200 |
+
v, t = rng.choice([20, 30, 40, 50, 60, 70, 80]), rng.randint(2, 9)
|
| 201 |
+
q, a_ = (f"一辆汽车每小时行驶 {v} 千米,行驶 {t} 小时,一共行驶多少千米?",
|
| 202 |
+
f"速度 {v} 千米/小时 × 时间 {t} 小时 = {v * t} 千米。")
|
| 203 |
+
else: # 反义词
|
| 204 |
+
w, ant = rng.choice(ANTONYMS)
|
| 205 |
+
q, a_ = f"'{w}'的反义词是什么?", f"'{w}'的反义词是'{ant}'。"
|
| 206 |
+
return q, a_
|
| 207 |
+
|
| 208 |
+
|
| 209 |
+
def _gen_hard(rng):
|
| 210 |
+
kind = rng.randrange(6)
|
| 211 |
+
if kind == 0: # 三位数×两位数, 答案带分解步骤
|
| 212 |
+
a, b = rng.randint(100, 999), rng.randint(10, 99)
|
| 213 |
+
tens, ones = b // 10 * 10, b % 10
|
| 214 |
+
q = f"{a} 乘 {b} 等于多少?"
|
| 215 |
+
a_ = (f"{a} × {b} = {a * b}。把 {b} 拆成 {tens} + {ones}:"
|
| 216 |
+
f"{a} × {tens} = {a * tens},{a} × {ones} = {a * ones},"
|
| 217 |
+
f"相加得 {a * tens + a * ones} = {a * b}。")
|
| 218 |
+
elif kind == 1: # 鸡兔同笼 (反向构造保证整数解)
|
| 219 |
+
c, r = rng.randint(2, 20), rng.randint(2, 20)
|
| 220 |
+
heads, feet = c + r, 2 * c + 4 * r
|
| 221 |
+
q = f"鸡兔同笼,从上面数有 {heads} 个头,从下面数有 {feet} 只脚,鸡和兔各有几只?"
|
| 222 |
+
a_ = (f"设鸡 x 只,兔 {heads} - x 只。列方程 2x + 4({heads} - x) = {feet},"
|
| 223 |
+
f"化简得 {4 * heads} - 2x = {feet},x = {c}。"
|
| 224 |
+
f"所以鸡 {c} 只,兔 {r} 只。")
|
| 225 |
+
elif kind == 2: # 相遇问题
|
| 226 |
+
v1, v2 = rng.choice([20, 30, 40, 50, 60]), rng.choice([20, 30, 40, 50, 60])
|
| 227 |
+
t = rng.randint(2, 6)
|
| 228 |
+
d = (v1 + v2) * t
|
| 229 |
+
q = (f"甲、乙两车分别从相距 {d} 千米的两地同时出发,相向而行。"
|
| 230 |
+
f"甲车每小时行 {v1} 千米,乙车每小时行 {v2} 千米,几小时后两车相遇?")
|
| 231 |
+
a_ = (f"两车速度和为 {v1} + {v2} = {v1 + v2} 千米/小时。"
|
| 232 |
+
f"相遇时间 = 距离 ÷ 速度和 = {d} ÷ {v1 + v2} = {t} 小时。")
|
| 233 |
+
elif kind == 3: # 工程问题 (预解表保证答案整洁)
|
| 234 |
+
a, b, t = rng.choice(WORK_PROBLEMS)
|
| 235 |
+
q = f"一个水池,甲管单独注满需要 {a} 小时,乙管单独注满需要 {b} 小时,两管同时开需要几小时注满?"
|
| 236 |
+
a_ = (f"设总量为 1。甲管每小时注 1/{a},乙管每小时注 1/{b},"
|
| 237 |
+
f"两管同开每小时注 1/{a} + 1/{b} = {a + b}/{a * b}。"
|
| 238 |
+
f"注满需要 {a * b} ÷ {a + b} = {t} 小时。")
|
| 239 |
+
elif kind == 4: # 数列推理 (等差/等比/二阶)
|
| 240 |
+
sub = rng.randrange(3)
|
| 241 |
+
if sub == 0:
|
| 242 |
+
a1, d = rng.randint(2, 20), rng.randint(2, 12)
|
| 243 |
+
seq = [a1 + i * d for i in range(5)]
|
| 244 |
+
q, a_ = (f"找规律:{', '.join(map(str, seq))},下一个数是多少?",
|
| 245 |
+
f"相邻两项的差恒为 {d},下一个数是 {seq[-1]} + {d} = {seq[-1] + d}。")
|
| 246 |
+
elif sub == 1:
|
| 247 |
+
a1, r = rng.randint(1, 4), rng.choice([2, 3])
|
| 248 |
+
seq, v = [], a1
|
| 249 |
+
for _ in range(5):
|
| 250 |
+
seq.append(v)
|
| 251 |
+
v *= r
|
| 252 |
+
q, a_ = (f"找规律:{', '.join(map(str, seq))},下一个数是多少?",
|
| 253 |
+
f"相邻两项的比为 {r},下一个数是 {seq[-1]} × {r} = {seq[-1] * r}。")
|
| 254 |
+
else:
|
| 255 |
+
a1, d0, e = rng.randint(1, 9), rng.randint(1, 5), rng.randint(1, 4)
|
| 256 |
+
seq, v, dd = [], a1, d0
|
| 257 |
+
for _ in range(5):
|
| 258 |
+
seq.append(v)
|
| 259 |
+
v += dd
|
| 260 |
+
dd += e
|
| 261 |
+
q, a_ = (f"找规律:{', '.join(map(str, seq))},下一个数是多少?",
|
| 262 |
+
f"相邻两项的差依次是 {', '.join(str(d0 + i * e) for i in range(4))},"
|
| 263 |
+
f"差递增 {e}。下一个差是 {dd},下一个数是 {seq[-1]} + {dd} = {seq[-1] + dd}。")
|
| 264 |
+
else: # 除法余数 / 浓度
|
| 265 |
+
if rng.random() < 0.5:
|
| 266 |
+
divisor = rng.randint(6, 12)
|
| 267 |
+
r_ = rng.randint(1, divisor - 1)
|
| 268 |
+
quot = rng.randint(10, 99)
|
| 269 |
+
n = divisor * quot + r_
|
| 270 |
+
q = f"{n} 除以 {divisor},商和余数分别是多少?"
|
| 271 |
+
a_ = f"{n} ÷ {divisor} = {quot} 余 {r_},即商是 {quot},余数是 {r_}。"
|
| 272 |
+
else:
|
| 273 |
+
c = rng.choice([10, 20, 25, 50])
|
| 274 |
+
m = rng.choice([200, 400, 500, 800, 1000])
|
| 275 |
+
sugar = m * c // 100
|
| 276 |
+
q = f"{m} 克糖水中含糖 {sugar} 克,这种糖水的浓度是多少?"
|
| 277 |
+
a_ = f"浓度 = 糖 ÷ 糖水 = {sugar} ÷ {m} = {c}%。"
|
| 278 |
+
return q, a_
|
| 279 |
+
|
| 280 |
+
|
| 281 |
+
def make_mix_data(train_path: str, n_easy=220, n_medium=200, n_hard=180,
|
| 282 |
+
n_eval_per_bucket=30, seed=7):
|
| 283 |
+
"""生成难度分级数据集: train_path (训练) + train_path 去掉 .jsonl 加 .eval.jsonl (评估)。
|
| 284 |
+
每行: {"messages": [...], "difficulty": "easy|medium|hard"}"""
|
| 285 |
+
rng = random.Random(seed)
|
| 286 |
+
eval_path = train_path[:-len(".jsonl")] + ".eval.jsonl" if train_path.endswith(".jsonl") \
|
| 287 |
+
else train_path + ".eval.jsonl"
|
| 288 |
+
gens = {"easy": _gen_easy, "medium": _gen_medium, "hard": _gen_hard}
|
| 289 |
+
|
| 290 |
+
def fill(bucket, n, used, salt):
|
| 291 |
+
rows, tries = [], 0
|
| 292 |
+
r2 = random.Random(seed + salt)
|
| 293 |
+
while len(rows) < n and tries < n * 50:
|
| 294 |
+
tries += 1
|
| 295 |
+
q, a = gens[bucket](r2)
|
| 296 |
+
if q in used:
|
| 297 |
+
continue
|
| 298 |
+
used.add(q)
|
| 299 |
+
rows.append({"messages": [{"role": "user", "content": q},
|
| 300 |
+
{"role": "assistant", "content": a}],
|
| 301 |
+
"difficulty": bucket})
|
| 302 |
+
return rows
|
| 303 |
+
|
| 304 |
+
used = set()
|
| 305 |
+
buckets = ["easy", "medium", "hard"]
|
| 306 |
+
train_rows = {b: fill(b, n, used, i)
|
| 307 |
+
for i, (b, n) in enumerate(zip(buckets, [n_easy, n_medium, n_hard]))}
|
| 308 |
+
eval_rows = {b: fill(b, n_eval_per_bucket, used, 100) for b in buckets}
|
| 309 |
+
|
| 310 |
+
for path, rows in [(train_path, train_rows), (eval_path, eval_rows)]:
|
| 311 |
+
with open(path, "w", encoding="utf-8") as f:
|
| 312 |
+
for b in ("easy", "medium", "hard"):
|
| 313 |
+
for r in rows[b]:
|
| 314 |
+
f.write(json.dumps(r, ensure_ascii=False) + "\n")
|
| 315 |
+
counts = {b: len(rows[b]) for b in rows}
|
| 316 |
+
print(f"已写入 {path} {counts}")
|
| 317 |
+
|
| 318 |
+
|
| 319 |
+
def load_jsonl(path: str) -> list[dict]:
|
| 320 |
+
rows = []
|
| 321 |
+
with open(path, encoding="utf-8") as f:
|
| 322 |
+
for line in f:
|
| 323 |
+
line = line.strip()
|
| 324 |
+
if line:
|
| 325 |
+
rows.append(json.loads(line))
|
| 326 |
+
return rows
|
| 327 |
+
|
| 328 |
+
|
| 329 |
+
def _apply_chat_ids(tok, msgs, add_generation_prompt=False) -> list[int]:
|
| 330 |
+
"""apply_chat_template 跨版本归一化: 统一返回 list[int]
|
| 331 |
+
(transformers 5.x 返回 BatchEncoding, 旧版返回 list[int], fast 分支返回 Encoding)"""
|
| 332 |
+
out = tok.apply_chat_template(msgs, tokenize=True,
|
| 333 |
+
add_generation_prompt=add_generation_prompt)
|
| 334 |
+
if hasattr(out, "ids"): # tokenizers.Encoding
|
| 335 |
+
return list(out.ids)
|
| 336 |
+
if hasattr(out, "input_ids"): # BatchEncoding
|
| 337 |
+
out = out["input_ids"]
|
| 338 |
+
if isinstance(out, torch.Tensor):
|
| 339 |
+
out = out.tolist()
|
| 340 |
+
if out and isinstance(out[0], (list, tuple)): # 嵌套 -> 单条
|
| 341 |
+
out = out[0]
|
| 342 |
+
return [int(t) for t in out]
|
| 343 |
+
|
| 344 |
+
|
| 345 |
+
class PonderDataset(Dataset):
|
| 346 |
+
"""jsonl -> token ids。chat 格式走 chat_template 并只对回答部分计 loss;
|
| 347 |
+
纯文本格式整句计 loss。右侧 padding, labels 在 pad 处置 -100。"""
|
| 348 |
+
|
| 349 |
+
def __init__(self, rows, tokenizer, max_len=512, text_field=None):
|
| 350 |
+
tok = tokenizer
|
| 351 |
+
self.tok = tokenizer
|
| 352 |
+
self.max_len = max_len
|
| 353 |
+
self.examples = []
|
| 354 |
+
pad_id = tokenizer.pad_token_id or 0
|
| 355 |
+
for r in rows:
|
| 356 |
+
if text_field and text_field in r:
|
| 357 |
+
enc = tokenizer(r[text_field], truncation=True, max_length=max_len)
|
| 358 |
+
ids, label_mask = enc["input_ids"], [1] * len(enc["input_ids"])
|
| 359 |
+
else:
|
| 360 |
+
msgs = r["messages"]
|
| 361 |
+
# 逐轮拼接: assistant 回答的 token 区间计 loss, 其余位置置 -100
|
| 362 |
+
ids: list[int] = []
|
| 363 |
+
label_mask: list[int] = []
|
| 364 |
+
for i, m in enumerate(msgs):
|
| 365 |
+
cur = _apply_chat_ids(tok, msgs[: i + 1])
|
| 366 |
+
new = cur[len(ids):] if cur[: len(ids)] == ids else cur
|
| 367 |
+
label_mask += ([1] if m["role"] == "assistant" else [0]) * len(new)
|
| 368 |
+
ids = cur
|
| 369 |
+
ids = ids[:max_len]
|
| 370 |
+
label_mask = label_mask[: len(ids)]
|
| 371 |
+
labels = [t if mk == 1 else -100 for t, mk in zip(ids, label_mask)]
|
| 372 |
+
# difficulty 透传给 PonderLoader (--difficulty-prior 需要)
|
| 373 |
+
self.examples.append({"input_ids": ids, "labels": labels,
|
| 374 |
+
"difficulty": r.get("difficulty", "unknown")})
|
| 375 |
+
|
| 376 |
+
def __len__(self):
|
| 377 |
+
return len(self.examples)
|
| 378 |
+
|
| 379 |
+
def __getitem__(self, i):
|
| 380 |
+
return self.examples[i]
|
| 381 |
+
|
| 382 |
+
|
| 383 |
+
def collate(batch, pad_id=1):
|
| 384 |
+
maxlen = max(len(b["input_ids"]) for b in batch)
|
| 385 |
+
input_ids, labels, attn = [], [], []
|
| 386 |
+
for b in batch:
|
| 387 |
+
n = len(b["input_ids"])
|
| 388 |
+
pad = maxlen - n
|
| 389 |
+
input_ids.append(b["input_ids"] + [pad_id] * pad)
|
| 390 |
+
labels.append(b["labels"] + [-100] * pad)
|
| 391 |
+
attn.append([1] * n + [0] * pad)
|
| 392 |
+
return (torch.tensor(input_ids), torch.tensor(labels), torch.tensor(attn))
|
| 393 |
+
|
| 394 |
+
|
| 395 |
+
class PonderLoader:
|
| 396 |
+
"""难度感知 DataLoader: 同难度样本组成同 batch, 迭代产出
|
| 397 |
+
(input_ids, labels, attn, difficulty)。配合 --difficulty-prior 使用 ——
|
| 398 |
+
同质 batch 使得按难度动态设置 config.ponder_prior_p 成为可能
|
| 399 |
+
(KL 先验把 easy 拉向快停 / hard 拉向多想, 语义分化由此而来)。
|
| 400 |
+
不用 --difficulty-prior 时走普通 DataLoader, 行为不变。"""
|
| 401 |
+
|
| 402 |
+
def __init__(self, ds, batch_size, pad_id, seed=7):
|
| 403 |
+
self.ds, self.bs, self.pad_id, self.seed = ds, batch_size, pad_id, seed
|
| 404 |
+
self.epoch = 0
|
| 405 |
+
self.by_bucket: dict[str, list[int]] = {}
|
| 406 |
+
for i, ex in enumerate(ds.examples):
|
| 407 |
+
self.by_bucket.setdefault(ex.get("difficulty", "unknown"), []).append(i)
|
| 408 |
+
|
| 409 |
+
def __len__(self):
|
| 410 |
+
return sum((len(v) + self.bs - 1) // self.bs for v in self.by_bucket.values())
|
| 411 |
+
|
| 412 |
+
def __iter__(self):
|
| 413 |
+
rng = random.Random(self.seed + self.epoch)
|
| 414 |
+
self.epoch += 1
|
| 415 |
+
batches = []
|
| 416 |
+
for d, idxs in self.by_bucket.items():
|
| 417 |
+
idxs = idxs[:]
|
| 418 |
+
rng.shuffle(idxs)
|
| 419 |
+
batches += [(d, idxs[i:i + self.bs])
|
| 420 |
+
for i in range(0, len(idxs), self.bs)]
|
| 421 |
+
rng.shuffle(batches) # batch 间顺序打乱 (batch 内同质)
|
| 422 |
+
for d, idxs in batches:
|
| 423 |
+
yield (*collate([self.ds[i] for i in idxs], self.pad_id), d)
|
| 424 |
+
|
| 425 |
+
|
| 426 |
+
# ========================================================================== #
|
| 427 |
+
# 模型 / 参数分组
|
| 428 |
+
# ========================================================================== #
|
| 429 |
+
def resolve_dtype(arg, device):
|
| 430 |
+
"""--dtype 解析。auto: cpu→fp32; cuda→原生 bf16 (sm80+) 否则 fp16。
|
| 431 |
+
注意: 不能用 torch.cuda.is_bf16_supported() 判断 T4 —— 新版 torch 对 sm75
|
| 432 |
+
返回 True 但那是模拟 (emulated) bf16, 没有硬件 tensor core, 训练极慢。
|
| 433 |
+
T4 (sm75) → fp16 + GradScaler; 30/40/50 系 (sm80+) → 原生 bf16。"""
|
| 434 |
+
if arg == "float32":
|
| 435 |
+
return torch.float32
|
| 436 |
+
if arg == "float16":
|
| 437 |
+
return torch.float16
|
| 438 |
+
if arg == "bfloat16":
|
| 439 |
+
return torch.bfloat16
|
| 440 |
+
if device == "cpu":
|
| 441 |
+
return torch.float32
|
| 442 |
+
major, _ = torch.cuda.get_device_capability(device)
|
| 443 |
+
return torch.bfloat16 if major >= 8 and torch.cuda.is_bf16_supported() else torch.float16
|
| 444 |
+
|
| 445 |
+
|
| 446 |
+
def build_model(path, args, device, dtype=None):
|
| 447 |
+
"""dtype: torch dtype; None 时按 --dtype 解析 (auto: cuda 无 bf16 支持则 fp16, T4 即此路径)."""
|
| 448 |
+
dtype = dtype or resolve_dtype(args.dtype, device)
|
| 449 |
+
print(f"模型精度: {dtype} (device={device})")
|
| 450 |
+
cfg = PonderLlamaConfig.from_llama(
|
| 451 |
+
LlamaConfig.from_pretrained(path),
|
| 452 |
+
ponder_signal="learned", # 训练必须用可学习信号
|
| 453 |
+
max_ponder_steps=args.max_steps,
|
| 454 |
+
ponder_epsilon=args.epsilon,
|
| 455 |
+
ponder_prior_p=args.prior_p,
|
| 456 |
+
ponder_loss_beta=args.beta,
|
| 457 |
+
ponder_start_layer=args.start_layer,
|
| 458 |
+
ponder_all_positions=True, # 训练时所有位置独立 halting
|
| 459 |
+
ponder_head_hidden=getattr(args, "head_hidden", 0), # 方向①: >0 时 MLP 头
|
| 460 |
+
ponder_step_ce_weight=getattr(args, "step_ce_weight", 0.0), # 方向②: per-step CE
|
| 461 |
+
)
|
| 462 |
+
model = PonderLlamaForCausalLM.from_pretrained(
|
| 463 |
+
path, config=cfg, torch_dtype=dtype,
|
| 464 |
+
low_cpu_mem_usage=True, attn_implementation="sdpa")
|
| 465 |
+
# halting 头始终用 fp32 保证数值稳定
|
| 466 |
+
model.ponder_head = model.ponder_head.float()
|
| 467 |
+
# 初始停机概率: sigmoid(bias)。bias=-1 → λ≈0.27, 期望思考步数 ≈ 3.7
|
| 468 |
+
# 兼容 Linear (方向① 默认) 与 Sequential/MLP (方向① --head-hidden>0) 两种头
|
| 469 |
+
last_linear = [m for m in model.ponder_head.modules()
|
| 470 |
+
if isinstance(m, torch.nn.Linear)][-1]
|
| 471 |
+
with torch.no_grad():
|
| 472 |
+
last_linear.bias.fill_(args.head_bias_init)
|
| 473 |
+
if getattr(args, "head_hidden", 0) and args.head_hidden > 0:
|
| 474 |
+
print(f"方向①: halting 头升级 MLP (hidden={args.head_hidden}, "
|
| 475 |
+
f"参数 {sum(p.numel() for p in model.ponder_head.parameters()):,})")
|
| 476 |
+
model.to(device)
|
| 477 |
+
if args.grad_ckpt and device != "cpu":
|
| 478 |
+
model.gradient_checkpointing_enable()
|
| 479 |
+
return model
|
| 480 |
+
|
| 481 |
+
|
| 482 |
+
def setup_trainables(model, args):
|
| 483 |
+
"""返回 (可训练参数组, 冻结说明). 三种模式对应不同的梯度开放范围."""
|
| 484 |
+
n = model.config.num_hidden_layers
|
| 485 |
+
s = model.config.ponder_start_layer
|
| 486 |
+
# 1) 全部冻结
|
| 487 |
+
for p in model.parameters():
|
| 488 |
+
p.requires_grad_(False)
|
| 489 |
+
groups = []
|
| 490 |
+
# 2) halting 头 (总是可训, fp32)
|
| 491 |
+
for p in model.ponder_head.parameters():
|
| 492 |
+
p.requires_grad_(True)
|
| 493 |
+
groups.append({"params": [p], "lr": args.lr_head, "name": "head"})
|
| 494 |
+
frozen_note = ["backbone"]
|
| 495 |
+
# 3) 思考块 / LoRA
|
| 496 |
+
if args.train_mode == "head+block":
|
| 497 |
+
for layer in model.model.layers[s:n]:
|
| 498 |
+
for p in layer.parameters():
|
| 499 |
+
p.requires_grad_(True)
|
| 500 |
+
groups.append({"params": [p for layer in model.model.layers[s:n]
|
| 501 |
+
for p in layer.parameters()],
|
| 502 |
+
"lr": args.lr_block, "name": "block"})
|
| 503 |
+
frozen_note = ["layers 0.." + str(s - 1)]
|
| 504 |
+
elif args.train_mode == "head+lora":
|
| 505 |
+
try:
|
| 506 |
+
from peft import LoraConfig, get_peft_model, PeftModel
|
| 507 |
+
except ImportError as e:
|
| 508 |
+
raise SystemExit("head+lora 需要 peft: pip install peft") from e
|
| 509 |
+
if getattr(args, "init_adapter", None):
|
| 510 |
+
# 失败自举第三轮: 加载已训 LoRA 继续训 (head 重置为新初始化)
|
| 511 |
+
model = PeftModel.from_pretrained(model, args.init_adapter,
|
| 512 |
+
is_trainable=True)
|
| 513 |
+
print(f"已加载并继续训练 LoRA: {args.init_adapter}")
|
| 514 |
+
else:
|
| 515 |
+
scope = (list(range(s, n)) if args.lora_scope == "ponder" else list(range(n)))
|
| 516 |
+
target_names = set()
|
| 517 |
+
for i in scope:
|
| 518 |
+
for mod in ("q_proj", "k_proj", "v_proj", "o_proj"):
|
| 519 |
+
target_names.add(f"model.layers.{i}.self_attn.{mod}")
|
| 520 |
+
for mod in ("gate_proj", "up_proj", "down_proj"):
|
| 521 |
+
target_names.add(f"model.layers.{i}.mlp.{mod}")
|
| 522 |
+
lcfg = LoraConfig(r=args.lora_r, lora_alpha=args.lora_r * 2, lora_dropout=0.0,
|
| 523 |
+
bias="none", task_type="CAUSAL_LM",
|
| 524 |
+
target_modules=list(target_names))
|
| 525 |
+
model = get_peft_model(model, lcfg)
|
| 526 |
+
# 重新分组 (peft 包装后参数引用变化)
|
| 527 |
+
groups = [g for g in groups if g["name"] == "head"]
|
| 528 |
+
model.base_model.model.ponder_head.requires_grad_(True) # 保险
|
| 529 |
+
lora_params = [p for pn, p in model.named_parameters()
|
| 530 |
+
if p.requires_grad and "lora_" in pn]
|
| 531 |
+
groups.append({"params": lora_params, "lr": args.lr_lora, "name": "lora"})
|
| 532 |
+
frozen_note = [f"backbone (+LoRA r={args.lora_r} on {args.lora_scope} layers)"]
|
| 533 |
+
elif args.train_mode == "head" and getattr(args, "init_adapter", None):
|
| 534 |
+
# 两阶段 (失败自举第二阶段): 加载已训 LoRA 并冻结, 只重训 halting 头 ——
|
| 535 |
+
# 表征固定后, head 只能在"1 步不够"的位置上学习停机 (排除 LoRA 同步快停的混淆)
|
| 536 |
+
try:
|
| 537 |
+
from peft import PeftModel
|
| 538 |
+
except ImportError as e:
|
| 539 |
+
raise SystemExit("--init-adapter 需要 peft: pip install peft") from e
|
| 540 |
+
model = PeftModel.from_pretrained(model, args.init_adapter)
|
| 541 |
+
for pn_, p_ in model.named_parameters():
|
| 542 |
+
if "lora_" in pn_:
|
| 543 |
+
p_.requires_grad_(False)
|
| 544 |
+
# head 参数引用在 peft 包装后不变 (同一张量), groups 中的引用仍有效;
|
| 545 |
+
# peft from_pretrained 默认全部 requires_grad=False, 显式重开 head
|
| 546 |
+
for p in model.ponder_head.parameters():
|
| 547 |
+
p.requires_grad_(True)
|
| 548 |
+
frozen_note = [f"backbone + frozen LoRA from {args.init_adapter}"]
|
| 549 |
+
return model, groups, frozen_note
|
| 550 |
+
|
| 551 |
+
|
| 552 |
+
# ========================================================================== #
|
| 553 |
+
# 训练循环
|
| 554 |
+
# ========================================================================== #
|
| 555 |
+
def train_one_epoch(model, loader, optimizer, scheduler, device,
|
| 556 |
+
accum=1, clip=1.0, log_every=0, global_step=0, scaler=None,
|
| 557 |
+
prior_map=None, step_ce_weight=0.0):
|
| 558 |
+
model.train()
|
| 559 |
+
stats = {"loss": 0.0, "ce": 0.0, "kl": 0.0, "sce": 0.0, "steps": 0.0, "nb": 0}
|
| 560 |
+
optimizer.zero_grad(set_to_none=True)
|
| 561 |
+
t0 = time.time()
|
| 562 |
+
for it, item in enumerate(loader):
|
| 563 |
+
if len(item) == 4: # PonderLoader: 带 difficulty
|
| 564 |
+
input_ids, labels, attn, diff = item
|
| 565 |
+
pg = (prior_map or {}).get(diff)
|
| 566 |
+
if pg is not None: # 难度条件化先验: 同 batch 生效
|
| 567 |
+
model.config.ponder_prior_p = pg
|
| 568 |
+
else:
|
| 569 |
+
input_ids, labels, attn = item
|
| 570 |
+
input_ids, labels, attn = input_ids.to(device), labels.to(device), attn.to(device)
|
| 571 |
+
out = model(input_ids=input_ids, attention_mask=attn, labels=labels,
|
| 572 |
+
output_ponder=True)
|
| 573 |
+
# 注意: forward 的 loss = CE + beta*KL (+ step_ce_weight*per-step CE)
|
| 574 |
+
loss_to_back = out.loss / accum
|
| 575 |
+
if scaler is not None:
|
| 576 |
+
scaler.scale(loss_to_back).backward() # fp16 梯度缩放防下溢
|
| 577 |
+
else:
|
| 578 |
+
loss_to_back.backward()
|
| 579 |
+
with torch.no_grad():
|
| 580 |
+
# 思考步数只统计有监督的位置
|
| 581 |
+
mask = labels[:, 1:] != -100
|
| 582 |
+
steps = out.ponder_steps[:, :-1][mask].float().mean() if mask.any() \
|
| 583 |
+
else out.ponder_steps.float().mean()
|
| 584 |
+
sce = out.ponder_step_ce.item() if out.ponder_step_ce is not None else 0.0
|
| 585 |
+
stats["loss"] += out.loss.item()
|
| 586 |
+
stats["ce"] += (out.loss.item()
|
| 587 |
+
- (model.config.ponder_loss_beta * out.ponder_kl.item())
|
| 588 |
+
- step_ce_weight * sce)
|
| 589 |
+
stats["kl"] += out.ponder_kl.item()
|
| 590 |
+
stats["sce"] += sce
|
| 591 |
+
stats["steps"] += steps.item()
|
| 592 |
+
stats["nb"] += 1
|
| 593 |
+
if (it + 1) % accum == 0 or (it + 1) == len(loader):
|
| 594 |
+
if scaler is not None:
|
| 595 |
+
scaler.unscale_(optimizer) # clip 必须在真实梯度上做
|
| 596 |
+
torch.nn.utils.clip_grad_norm_(
|
| 597 |
+
[p for g in optimizer.param_groups for p in g["params"]], clip)
|
| 598 |
+
if scaler is not None:
|
| 599 |
+
scaler.step(optimizer)
|
| 600 |
+
scaler.update()
|
| 601 |
+
else:
|
| 602 |
+
optimizer.step()
|
| 603 |
+
scheduler.step()
|
| 604 |
+
optimizer.zero_grad(set_to_none=True)
|
| 605 |
+
global_step += 1
|
| 606 |
+
if log_every and global_step % log_every == 0:
|
| 607 |
+
k = max(stats["nb"], 1)
|
| 608 |
+
sce_part = (f" | stepCE {stats['sce']/k:6.4f}"
|
| 609 |
+
if step_ce_weight > 0 else "")
|
| 610 |
+
print(f" step {global_step:4d} | loss {stats['loss']/k:7.4f} "
|
| 611 |
+
f"(ce {stats['ce']/k:7.4f} + {model.config.ponder_loss_beta}·kl {stats['kl']/k:6.4f}"
|
| 612 |
+
f"{sce_part})"
|
| 613 |
+
f" | ponder_steps {stats['steps']/k:5.2f} | lr {scheduler.get_last_lr()[0]:.2e}"
|
| 614 |
+
f" | {time.time()-t0:6.1f}s", flush=True)
|
| 615 |
+
stats = {"loss": 0.0, "ce": 0.0, "kl": 0.0, "sce": 0.0, "steps": 0.0, "nb": 0}
|
| 616 |
+
t0 = time.time()
|
| 617 |
+
return global_step
|
| 618 |
+
|
| 619 |
+
|
| 620 |
+
def save_model(model, args, out_dir):
|
| 621 |
+
os.makedirs(out_dir, exist_ok=True)
|
| 622 |
+
core = model.get_base_model() if hasattr(model, "get_base_model") else model
|
| 623 |
+
# ponder_head 总是单独保存 (peft 适配器不包含它)。
|
| 624 |
+
# 键名 = state_dict 去掉 ponder_head. 前缀: Linear 头 → weight/bias (向后兼容);
|
| 625 |
+
# MLP 头 (方向①) → 0.weight/0.bias/2.weight/2.bias
|
| 626 |
+
from safetensors.torch import save_file
|
| 627 |
+
head_sd = {k: v.data.to(torch.float32).cpu()
|
| 628 |
+
for k, v in core.ponder_head.state_dict().items()}
|
| 629 |
+
save_file(head_sd, os.path.join(out_dir, "ponder_head.safetensors"),
|
| 630 |
+
metadata={"note": f"signal=learned, K={args.max_steps}, p_g={args.prior_p}, "
|
| 631 |
+
f"head_hidden={getattr(args, 'head_hidden', 0)}"})
|
| 632 |
+
if args.train_mode == "head+lora":
|
| 633 |
+
model.save_pretrained(out_dir) # adapter_model.safetensors
|
| 634 |
+
print(f"已保存 LoRA 适配器 + ponder_head 到 {out_dir}")
|
| 635 |
+
elif args.train_mode == "head+block":
|
| 636 |
+
core.save_pretrained(out_dir)
|
| 637 |
+
print(f"已保存全量模型 (含新 ponder_head) 到 {out_dir}")
|
| 638 |
+
else:
|
| 639 |
+
print(f"已保存 ponder_head.safetensors 到 {out_dir} "
|
| 640 |
+
f"(推理时 from_ponder 加载后 model.ponder_head.load_state_dict)")
|
| 641 |
+
|
| 642 |
+
|
| 643 |
+
# ========================================================================== #
|
| 644 |
+
# main
|
| 645 |
+
# ========================================================================== #
|
| 646 |
+
def main():
|
| 647 |
+
ap = argparse.ArgumentParser(description=__doc__.split("\n")[1])
|
| 648 |
+
ap.add_argument("--model", default=None,
|
| 649 |
+
help="模型路径; 仅生成数据 (--make-demo-data/--make-mix-data) 时可省略")
|
| 650 |
+
ap.add_argument("--data", default=None, help="jsonl 训练数据; 不给且无 --make-* 时用内置演示数据")
|
| 651 |
+
ap.add_argument("--make-demo-data", default=None, metavar="PATH")
|
| 652 |
+
ap.add_argument("--make-mix-data", default=None, metavar="PATH",
|
| 653 |
+
help="生成难度分级数据集 (easy/medium/hard + .eval.jsonl 评估集)")
|
| 654 |
+
ap.add_argument("--n-easy", type=int, default=220)
|
| 655 |
+
ap.add_argument("--n-medium", type=int, default=200)
|
| 656 |
+
ap.add_argument("--n-hard", type=int, default=180)
|
| 657 |
+
ap.add_argument("--n-eval-per-bucket", type=int, default=30)
|
| 658 |
+
ap.add_argument("--seed", type=int, default=7)
|
| 659 |
+
ap.add_argument("--text-field", default=None, help="纯文本字段名 (默认按 chat messages 解析)")
|
| 660 |
+
ap.add_argument("--train-mode", default="head", choices=["head", "head+lora", "head+block"])
|
| 661 |
+
ap.add_argument("--output", default="out/ponder")
|
| 662 |
+
ap.add_argument("--epochs", type=float, default=2.0)
|
| 663 |
+
ap.add_argument("--batch-size", type=int, default=2)
|
| 664 |
+
ap.add_argument("--accum", type=int, default=4)
|
| 665 |
+
ap.add_argument("--max-len", type=int, default=512)
|
| 666 |
+
ap.add_argument("--lr-head", type=float, default=1e-3)
|
| 667 |
+
ap.add_argument("--lr-lora", type=float, default=2e-4)
|
| 668 |
+
ap.add_argument("--lr-block", type=float, default=1e-5)
|
| 669 |
+
ap.add_argument("--warmup-ratio", type=float, default=0.03)
|
| 670 |
+
ap.add_argument("--clip", type=float, default=1.0)
|
| 671 |
+
ap.add_argument("--max-steps", type=int, default=8, help="思考步数上限 K")
|
| 672 |
+
ap.add_argument("--epsilon", type=float, default=0.01)
|
| 673 |
+
ap.add_argument("--prior-p", type=float, default=0.4, help="几何先验 p_g (越小越鼓励多想)")
|
| 674 |
+
ap.add_argument("--difficulty-prior", default=None,
|
| 675 |
+
help="难度条件化先验, 如 easy:0.7,medium:0.4,hard:0.15 —— "
|
| 676 |
+
"同难度样本组成同 batch, KL 先验按难度调整 "
|
| 677 |
+
"(hard 用小 p_g 鼓励多想; 需数据带 difficulty 字段)")
|
| 678 |
+
ap.add_argument("--beta", type=float, default=0.05, help="KL 正则权重")
|
| 679 |
+
ap.add_argument("--head-bias-init", type=float, default=-1.0, help="初始停机偏置, sigmoid(-1)≈0.27")
|
| 680 |
+
ap.add_argument("--head-hidden", type=int, default=0,
|
| 681 |
+
help="方向①: >0 时 halting 头升级 MLP (Linear->GELU->Linear), "
|
| 682 |
+
"打破线性可分上限 (round4: 线性头在冻结特征上已到极限)")
|
| 683 |
+
ap.add_argument("--step-ce-weight", type=float, default=0.0,
|
| 684 |
+
help="方向②: per-step CE 权重 —— 每步思考表示单独算解码 CE 按 w 加权, "
|
| 685 |
+
"让后几步表示'可用', 回收 mix 摊薄代价 (需 head+lora/head+block 模式)")
|
| 686 |
+
ap.add_argument("--start-layer", type=int, default=None, help="思考块起始层 (默认最后8层)")
|
| 687 |
+
ap.add_argument("--lora-r", type=int, default=16)
|
| 688 |
+
ap.add_argument("--lora-scope", default="ponder", choices=["ponder", "all"])
|
| 689 |
+
ap.add_argument("--init-adapter", default=None,
|
| 690 |
+
help="两阶段训练: 加载已训 LoRA 适配器目录并冻结 (配 --train-mode head), "
|
| 691 |
+
"只重训 halting 头 (失败自举第二阶段)")
|
| 692 |
+
ap.add_argument("--dtype", default="auto", choices=["auto", "float16", "bfloat16", "float32"],
|
| 693 |
+
help="auto: T4 等 pre-Ampere 卡自动 fp16, 30/40/50 系 bf16, cpu fp32")
|
| 694 |
+
ap.add_argument("--grad-ckpt", action="store_true", default=None,
|
| 695 |
+
help="思考块梯度重算 (16GB 显存训练强烈建议开启)")
|
| 696 |
+
ap.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
|
| 697 |
+
ap.add_argument("--log-every", type=int, default=10)
|
| 698 |
+
args = ap.parse_args()
|
| 699 |
+
|
| 700 |
+
if args.make_mix_data:
|
| 701 |
+
make_mix_data(args.make_mix_data, n_easy=args.n_easy, n_medium=args.n_medium,
|
| 702 |
+
n_hard=args.n_hard, n_eval_per_bucket=args.n_eval_per_bucket,
|
| 703 |
+
seed=args.seed)
|
| 704 |
+
if not args.data:
|
| 705 |
+
return
|
| 706 |
+
if args.make_demo_data:
|
| 707 |
+
make_demo_data(args.make_demo_data)
|
| 708 |
+
if not args.data:
|
| 709 |
+
return
|
| 710 |
+
if not args.model:
|
| 711 |
+
ap.error("--model 为必填 (仅在仅使用 --make-demo-data / --make-mix-data 时可省略)")
|
| 712 |
+
|
| 713 |
+
# ---- 数据 ----
|
| 714 |
+
rows = load_jsonl(args.data) if args.data else DEMO_DATA
|
| 715 |
+
print(f"训练样本: {len(rows)} 条 (mode={args.train_mode})")
|
| 716 |
+
tok = AutoTokenizer.from_pretrained(args.model, trust_remote_code=True)
|
| 717 |
+
if tok.pad_token_id is None:
|
| 718 |
+
tok.pad_token = tok.eos_token
|
| 719 |
+
ds = PonderDataset(rows, tok, max_len=args.max_len, text_field=args.text_field)
|
| 720 |
+
prior_map = None
|
| 721 |
+
if args.difficulty_prior:
|
| 722 |
+
prior_map = {}
|
| 723 |
+
for part in args.difficulty_prior.split(","):
|
| 724 |
+
k, v = part.split(":")
|
| 725 |
+
prior_map[k.strip()] = float(v)
|
| 726 |
+
n_diff = sum(1 for ex in ds.examples if ex.get("difficulty") in prior_map)
|
| 727 |
+
print(f"难度条件化先验: {prior_map} | 命中样本 {n_diff}/{len(ds.examples)}")
|
| 728 |
+
if n_diff == 0:
|
| 729 |
+
raise SystemExit("--difficulty-prior 需要数据带 difficulty 字段 (如 make_gsm8k_zh_data.py / --make-mix-data 的产物)")
|
| 730 |
+
loader = PonderLoader(ds, args.batch_size, tok.pad_token_id or 1, args.seed)
|
| 731 |
+
else:
|
| 732 |
+
loader = DataLoader(ds, batch_size=args.batch_size, shuffle=True,
|
| 733 |
+
collate_fn=lambda b: collate(b, tok.pad_token_id or 1))
|
| 734 |
+
|
| 735 |
+
# ---- 模型 ----
|
| 736 |
+
if args.step_ce_weight > 0 and args.train_mode == "head":
|
| 737 |
+
print("[警告] --step-ce-weight>0 但 train-mode=head: 头不改变表示且权重已 detach, "
|
| 738 |
+
"per-step CE 将无梯度作用对象; 建议搭配 --train-mode head+lora / head+block")
|
| 739 |
+
dtype = resolve_dtype(args.dtype, args.device)
|
| 740 |
+
model = build_model(args.model, args, args.device, dtype=dtype)
|
| 741 |
+
model, groups, frozen_note = setup_trainables(model, args)
|
| 742 |
+
n_train = sum(p.numel() for g in groups for p in g["params"])
|
| 743 |
+
print(f"可训练参数: {n_train:,} ({[g['name'] for g in groups]}), 冻结: {frozen_note}")
|
| 744 |
+
|
| 745 |
+
# ---- 优化器 ----
|
| 746 |
+
optimizer = torch.optim.AdamW(
|
| 747 |
+
[{"params": g["params"], "lr": g["lr"]} for g in groups], weight_decay=0.01)
|
| 748 |
+
total_updates = max(1, int(len(loader) / args.accum * args.epochs))
|
| 749 |
+
scheduler = get_cosine_schedule_with_warmup(
|
| 750 |
+
optimizer, int(total_updates * args.warmup_ratio), total_updates)
|
| 751 |
+
# fp16 训练需要 GradScaler 防梯度下溢 (T4); bf16/fp32 不需要
|
| 752 |
+
scaler = None
|
| 753 |
+
if args.device.startswith("cuda") and dtype == torch.float16:
|
| 754 |
+
scaler = torch.amp.GradScaler(args.device)
|
| 755 |
+
print("已启用 GradScaler (fp16)")
|
| 756 |
+
|
| 757 |
+
# ---- 训练 ----
|
| 758 |
+
gs = 0
|
| 759 |
+
for ep in range(math.ceil(args.epochs)):
|
| 760 |
+
print(f"--- epoch {ep + 1}/{math.ceil(args.epochs)} ---")
|
| 761 |
+
gs = train_one_epoch(model, loader, optimizer, scheduler, args.device,
|
| 762 |
+
accum=args.accum, clip=args.clip,
|
| 763 |
+
log_every=args.log_every, global_step=gs, scaler=scaler,
|
| 764 |
+
prior_map=prior_map, step_ce_weight=args.step_ce_weight)
|
| 765 |
+
|
| 766 |
+
save_model(model, args, args.output)
|
| 767 |
+
print("完成。推理加载示例:")
|
| 768 |
+
print(" model = PonderLlamaForCausalLM.from_ponder('<模型路径>',")
|
| 769 |
+
print(f" ponder_kwargs={{'ponder_signal':'learned','max_ponder_steps':{args.max_steps}}})")
|
| 770 |
+
print(" head = load_file('<output>/ponder_head.safetensors')")
|
| 771 |
+
print(" model.ponder_head.load_state_dict(head)")
|
| 772 |
+
|
| 773 |
+
|
| 774 |
+
if __name__ == "__main__":
|
| 775 |
+
main()
|