tchbcb commited on
Commit
5da4700
·
verified ·
1 Parent(s): fe6d79c

round5: 全前向训练: head / head+lora / head+block + 方向①②开关

Browse files
Files changed (1) hide show
  1. 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()