pondernet: add difficulty-conditioned prior (--difficulty-prior, PonderLoader) + make_gsm8k_zh_data.py (HF dataset -> difficulty-bucketed CoT jsonl) + README section 10
Browse files- pondernet/README.md +46 -0
- pondernet/make_gsm8k_zh_data.py +117 -0
- pondernet/train_ponder_head.py +64 -6
pondernet/README.md
CHANGED
|
@@ -344,3 +344,49 @@ python eval_ponder.py --model models/MiniCPM5-2B --data data/mix.eval.jsonl \
|
|
| 344 |
注意:新版 torch 的 `is_bf16_supported()` 对 T4 (sm75) 返回 True 但那是模拟 bf16
|
| 345 |
(无硬件 tensor core),训练脚本已改为按 compute capability 判断,T4 自动落
|
| 346 |
fp16 + GradScaler。
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 344 |
注意:新版 torch 的 `is_bf16_supported()` 对 T4 (sm75) 返回 True 但那是模拟 bf16
|
| 345 |
(无硬件 tensor core),训练脚本已改为按 compute capability 判断,T4 自动落
|
| 346 |
fp16 + GradScaler。
|
| 347 |
+
|
| 348 |
+
## 十、"难题多想"语义分化专项:数据集选型与方案(2026-09-08)
|
| 349 |
+
|
| 350 |
+
9.4 指出核心矛盾:模板化短答案使 hard 桶 per-token 不确定性反而低。破局需要
|
| 351 |
+
**真实多步 CoT 数据**(让"1 步解不出"成为 CE 层面的现实)+ **难度→先验的显式通路**。
|
| 352 |
+
|
| 353 |
+
### 10.1 HF 数据集选型(实查字段与规模)
|
| 354 |
+
|
| 355 |
+
| 数据集 | 规模 | 语言 | 关键字段 | 适配点 |
|
| 356 |
+
|---|---|---|---|---|
|
| 357 |
+
| `meta-math/MetaMathQA_GSM8K_zh` | 240k | 中文 | `query_zh` / `response_zh` / `type` | 中文逐步 CoT 解答,替换模板答案的首选 |
|
| 358 |
+
| `swulling/gsm8k_chinese` | 7.5k | 中文题+英文解 | `answer` 内含 `<<48/2=24>>` 步标注 | 步标注可直接数出推理步数 = 现成难度标签 |
|
| 359 |
+
| `AI-MO/NuminaMath-CoT` | 860k | 英文为主 | `source`(gsm8k/amc_aime/olympiads…) | source 即天然难度梯度,适合跨级对比 |
|
| 360 |
+
| `rabinadk1/competition_math-level_5` 等 | MATH 分级子集 | 英文 | level 1–5 | 官方难度分级,可直接当 difficulty 字段 |
|
| 361 |
+
|
| 362 |
+
### 10.2 专门方案(按改动成本排序,1+2 组合拳已落地)
|
| 363 |
+
|
| 364 |
+
1. **CoT 数据替换(数据侧,首选)**:`make_gsm8k_zh_data.py` 从 MetaMathQA_GSM8K_zh
|
| 365 |
+
抽样,按"解答中等式行数"自动分桶(≤2 步 easy / 3 步 medium / ≥4 步 hard),
|
| 366 |
+
生成带 difficulty 字段的 jsonl。多步 CoT 的中间 token 无法靠 1 步记住,
|
| 367 |
+
halting 头将在真实的预测压力下学到"多想"。
|
| 368 |
+
2. **难度条件化先验(损失侧,已实现)**:`--difficulty-prior easy:0.7,medium:0.4,hard:0.15`。
|
| 369 |
+
实现:`PonderLoader` 把同难度样本组成同 batch,训练循环在每个 batch 前把
|
| 370 |
+
`config.ponder_prior_p` 切换为该难度的 p_g —— KL 正则随之把 easy 拉向快停、
|
| 371 |
+
hard 拉向多想。这是 PonderNet 几何先验的 per-sample 扩展,ponder_llama.py 零改动。
|
| 372 |
+
3. **两阶段训练(调度侧)**:先 LoRA 拟合 CoT 任务(CE 降下来),再冻结 LoRA
|
| 373 |
+
单独训 halting 头 —— 避免"快停"在训练早期成为 CE+KL 的联合最优解。
|
| 374 |
+
4. **失败自举(AdaptThink 式,进阶)**:先用 K=1 跑一遍训练集,答错的样本构成
|
| 375 |
+
"必须多想"集合,对其施加更小的 p_g(或直接监督步数目标)—— 用模型自己的
|
| 376 |
+
失败案例定义难度,比启发式分桶更精准。
|
| 377 |
+
|
| 378 |
+
### 10.3 一键复现(方案 1+2)
|
| 379 |
+
|
| 380 |
+
```bash
|
| 381 |
+
pip install datasets
|
| 382 |
+
python make_gsm8k_zh_data.py --out-dir data_zh --n-train-per-bucket 800 --n-eval-per-bucket 50
|
| 383 |
+
python train_ponder_head.py --model <MODEL_DIR> --train-mode head+lora \
|
| 384 |
+
--data data_zh/zh_cot.jsonl --difficulty-prior easy:0.7,medium:0.4,hard:0.15 \
|
| 385 |
+
--epochs 2 --batch-size 2 --accum 4 --grad-ckpt --dtype float16 --output out/ponder-zh
|
| 386 |
+
python eval_ponder.py --model <MODEL_DIR> --data data_zh/zh_cot.eval.jsonl \
|
| 387 |
+
--head out/ponder-zh/ponder_head.safetensors --adapter out/ponder-zh \
|
| 388 |
+
--max-steps 8 --tag zh-cot
|
| 389 |
+
```
|
| 390 |
+
|
| 391 |
+
预期:easy 桶步数 →~1,hard 桶步数显著上移(p_g=0.15 的几何先验期望 ≈6.7 步),
|
| 392 |
+
出现 `easy < medium < hard` 的语义分化;评估脚本按 difficulty 分桶可直接检验。
|
pondernet/make_gsm8k_zh_data.py
ADDED
|
@@ -0,0 +1,117 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# -*- coding: utf-8 -*-
|
| 2 |
+
"""
|
| 3 |
+
make_gsm8k_zh_data.py — 从 HF 拉取中文 CoT 数学数据, 按推理步数分难度桶
|
| 4 |
+
|
| 5 |
+
解决"T4 实测中 hard 桶答案模板化 → 步数分化不出现"的问题:
|
| 6 |
+
用真实多步推理数据 (MetaMathQA_GSM8K_zh, 响应为中文逐步 CoT) 替换模板化短答案,
|
| 7 |
+
使"1 步解不出"成为 CE 层面的现实; 难度桶按解答步数自动划分 (反向构造可靠性高)。
|
| 8 |
+
|
| 9 |
+
数据源 (HF):
|
| 10 |
+
meta-math/MetaMathQA_GSM8K_zh — 字段 query_zh / response_zh / type, 240k 条
|
| 11 |
+
分桶规则: response_zh 中"计算行"数 (含 = 的行) ≈ 推理步数
|
| 12 |
+
easy: 1-2 步 medium: 3 步 hard: >=4 步
|
| 13 |
+
|
| 14 |
+
用法 (T4/Colab 或任意有网环境):
|
| 15 |
+
pip install datasets
|
| 16 |
+
python make_gsm8k_zh_data.py --out-dir data_zh --n-train-per-bucket 800 \
|
| 17 |
+
--n-eval-per-bucket 50 --max-per-bucket 4000
|
| 18 |
+
输出: data_zh/zh_cot.jsonl + data_zh/zh_cot.eval.jsonl
|
| 19 |
+
每行: {"messages": [...], "difficulty": "easy|medium|hard", "n_steps": int}
|
| 20 |
+
"""
|
| 21 |
+
from __future__ import annotations
|
| 22 |
+
|
| 23 |
+
import argparse
|
| 24 |
+
import json
|
| 25 |
+
import os
|
| 26 |
+
import random
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
def count_reasoning_steps(response: str) -> int:
|
| 30 |
+
"""解答步数 ≈ 含等号的行数 (CoT 中每个计算步至少一个 =)"""
|
| 31 |
+
lines = [ln for ln in response.replace("。", "\n").split("\n")
|
| 32 |
+
if "=" in ln and any(ch.isdigit() for ch in ln)]
|
| 33 |
+
return max(1, len(lines))
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
def bucket_of(n_steps: int) -> str:
|
| 37 |
+
if n_steps <= 2:
|
| 38 |
+
return "easy"
|
| 39 |
+
if n_steps == 3:
|
| 40 |
+
return "medium"
|
| 41 |
+
return "hard"
|
| 42 |
+
|
| 43 |
+
|
| 44 |
+
def to_messages(query: str, response: str) -> list[dict]:
|
| 45 |
+
# 规范化答案结尾: 确保最后一行有"答案是 X"格式 (MetaMathQA 自带, 保持原样)
|
| 46 |
+
return [{"role": "user", "content": query.strip()},
|
| 47 |
+
{"role": "assistant", "content": response.strip()}]
|
| 48 |
+
|
| 49 |
+
|
| 50 |
+
def main():
|
| 51 |
+
ap = argparse.ArgumentParser()
|
| 52 |
+
ap.add_argument("--out-dir", default="data_zh")
|
| 53 |
+
ap.add_argument("--n-train-per-bucket", type=int, default=800)
|
| 54 |
+
ap.add_argument("--n-eval-per-bucket", type=int, default=50)
|
| 55 |
+
ap.add_argument("--max-per-bucket", type=int, default=4000,
|
| 56 |
+
help="每桶最多扫描收集的候选数 (控制内存)")
|
| 57 |
+
ap.add_argument("--max-chars", type=int, default=1200,
|
| 58 |
+
help="response_zh 最大字符数 (控制序列长度/T4 显存)")
|
| 59 |
+
ap.add_argument("--seed", type=int, default=42)
|
| 60 |
+
ap.add_argument("--dataset", default="meta-math/MetaMathQA_GSM8K_zh")
|
| 61 |
+
args = ap.parse_args()
|
| 62 |
+
|
| 63 |
+
from datasets import load_dataset
|
| 64 |
+
|
| 65 |
+
rng = random.Random(args.seed)
|
| 66 |
+
print(f"加载 {args.dataset} (streaming) ...")
|
| 67 |
+
ds = load_dataset(args.dataset, split="train", streaming=True)
|
| 68 |
+
|
| 69 |
+
buckets = {"easy": [], "medium": [], "hard": []}
|
| 70 |
+
cap = args.max_per_bucket
|
| 71 |
+
scanned = 0
|
| 72 |
+
for ex in ds:
|
| 73 |
+
scanned += 1
|
| 74 |
+
if scanned % 20000 == 0:
|
| 75 |
+
print(f" 已扫描 {scanned} 条, 桶计数: "
|
| 76 |
+
f"{ {k: len(v) for k, v in buckets.items()} }")
|
| 77 |
+
q, r = (ex.get("query_zh") or "").strip(), (ex.get("response_zh") or "").strip()
|
| 78 |
+
if not q or not r or len(r) > args.max_chars or len(q) > 400:
|
| 79 |
+
continue
|
| 80 |
+
if "答案是" not in r and "答案:" not in r and "####" not in r:
|
| 81 |
+
continue
|
| 82 |
+
n = count_reasoning_steps(r)
|
| 83 |
+
b = bucket_of(n)
|
| 84 |
+
if len(buckets[b]) >= cap:
|
| 85 |
+
continue
|
| 86 |
+
buckets[b].append({"messages": to_messages(q, r),
|
| 87 |
+
"difficulty": b, "n_steps": n})
|
| 88 |
+
if all(len(v) >= cap for v in buckets.values()):
|
| 89 |
+
break
|
| 90 |
+
print(f"扫描完成 ({scanned} 条), 各桶候选: "
|
| 91 |
+
f"{ {k: len(v) for k, v in buckets.items()} }")
|
| 92 |
+
|
| 93 |
+
os.makedirs(args.out_dir, exist_ok=True)
|
| 94 |
+
train_path = os.path.join(args.out_dir, "zh_cot.jsonl")
|
| 95 |
+
eval_path = os.path.join(args.out_dir, "zh_cot.eval.jsonl")
|
| 96 |
+
for path, n_take in [(train_path, args.n_train_per_bucket),
|
| 97 |
+
(eval_path, args.n_eval_per_bucket)]:
|
| 98 |
+
out = []
|
| 99 |
+
for b, pool in buckets.items():
|
| 100 |
+
pool = pool[:]
|
| 101 |
+
rng.shuffle(pool)
|
| 102 |
+
out.extend(pool[:n_take])
|
| 103 |
+
rng.shuffle(out)
|
| 104 |
+
with open(path, "w", encoding="utf-8") as f:
|
| 105 |
+
for r in out:
|
| 106 |
+
f.write(json.dumps(r, ensure_ascii=False) + "\n")
|
| 107 |
+
from collections import Counter
|
| 108 |
+
print(f"写入 {path}: {len(out)} 条 {dict(Counter(r['difficulty'] for r in out))}")
|
| 109 |
+
|
| 110 |
+
print("完成。训练建议 (难度条件化先验, 让 hard 学会多想):")
|
| 111 |
+
print(" python train_ponder_head.py --model <MODEL> --train-mode head+lora "
|
| 112 |
+
"--data data_zh/zh_cot.jsonl --difficulty-prior easy:0.7,medium:0.4,hard:0.15 "
|
| 113 |
+
"--epochs 2 --batch-size 2 --accum 4 --grad-ckpt --dtype float16")
|
| 114 |
+
|
| 115 |
+
|
| 116 |
+
if __name__ == "__main__":
|
| 117 |
+
main()
|
pondernet/train_ponder_head.py
CHANGED
|
@@ -369,7 +369,9 @@ class PonderDataset(Dataset):
|
|
| 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 |
-
|
|
|
|
|
|
|
| 373 |
|
| 374 |
def __len__(self):
|
| 375 |
return len(self.examples)
|
|
@@ -390,6 +392,37 @@ def collate(batch, pad_id=1):
|
|
| 390 |
return (torch.tensor(input_ids), torch.tensor(labels), torch.tensor(attn))
|
| 391 |
|
| 392 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 393 |
# ========================================================================== #
|
| 394 |
# 模型 / 参数分组
|
| 395 |
# ========================================================================== #
|
|
@@ -491,12 +524,20 @@ def setup_trainables(model, args):
|
|
| 491 |
# 训练循环
|
| 492 |
# ========================================================================== #
|
| 493 |
def train_one_epoch(model, loader, optimizer, scheduler, device,
|
| 494 |
-
accum=1, clip=1.0, log_every=0, global_step=0, scaler=None
|
|
|
|
| 495 |
model.train()
|
| 496 |
stats = {"loss": 0.0, "ce": 0.0, "kl": 0.0, "steps": 0.0, "nb": 0}
|
| 497 |
optimizer.zero_grad(set_to_none=True)
|
| 498 |
t0 = time.time()
|
| 499 |
-
for it,
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 500 |
input_ids, labels, attn = input_ids.to(device), labels.to(device), attn.to(device)
|
| 501 |
out = model(input_ids=input_ids, attention_mask=attn, labels=labels,
|
| 502 |
output_ponder=True)
|
|
@@ -592,6 +633,10 @@ def main():
|
|
| 592 |
ap.add_argument("--max-steps", type=int, default=8, help="思考步数上限 K")
|
| 593 |
ap.add_argument("--epsilon", type=float, default=0.01)
|
| 594 |
ap.add_argument("--prior-p", type=float, default=0.4, help="几何先验 p_g (越小越鼓励多想)")
|
|
|
|
|
|
|
|
|
|
|
|
|
| 595 |
ap.add_argument("--beta", type=float, default=0.05, help="KL 正则权重")
|
| 596 |
ap.add_argument("--head-bias-init", type=float, default=-1.0, help="初始停机偏置, sigmoid(-1)≈0.27")
|
| 597 |
ap.add_argument("--start-layer", type=int, default=None, help="思考块起始层 (默认最后8层)")
|
|
@@ -625,8 +670,20 @@ def main():
|
|
| 625 |
if tok.pad_token_id is None:
|
| 626 |
tok.pad_token = tok.eos_token
|
| 627 |
ds = PonderDataset(rows, tok, max_len=args.max_len, text_field=args.text_field)
|
| 628 |
-
|
| 629 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 630 |
|
| 631 |
# ---- 模型 ----
|
| 632 |
dtype = resolve_dtype(args.dtype, args.device)
|
|
@@ -653,7 +710,8 @@ def main():
|
|
| 653 |
print(f"--- epoch {ep + 1}/{math.ceil(args.epochs)} ---")
|
| 654 |
gs = train_one_epoch(model, loader, optimizer, scheduler, args.device,
|
| 655 |
accum=args.accum, clip=args.clip,
|
| 656 |
-
log_every=args.log_every, global_step=gs, scaler=scaler
|
|
|
|
| 657 |
|
| 658 |
save_model(model, args, args.output)
|
| 659 |
print("完成。推理加载示例:")
|
|
|
|
| 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)
|
|
|
|
| 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 |
# ========================================================================== #
|
|
|
|
| 524 |
# 训练循环
|
| 525 |
# ========================================================================== #
|
| 526 |
def train_one_epoch(model, loader, optimizer, scheduler, device,
|
| 527 |
+
accum=1, clip=1.0, log_every=0, global_step=0, scaler=None,
|
| 528 |
+
prior_map=None):
|
| 529 |
model.train()
|
| 530 |
stats = {"loss": 0.0, "ce": 0.0, "kl": 0.0, "steps": 0.0, "nb": 0}
|
| 531 |
optimizer.zero_grad(set_to_none=True)
|
| 532 |
t0 = time.time()
|
| 533 |
+
for it, item in enumerate(loader):
|
| 534 |
+
if len(item) == 4: # PonderLoader: 带 difficulty
|
| 535 |
+
input_ids, labels, attn, diff = item
|
| 536 |
+
pg = (prior_map or {}).get(diff)
|
| 537 |
+
if pg is not None: # 难度条件化先验: 同 batch 生效
|
| 538 |
+
model.config.ponder_prior_p = pg
|
| 539 |
+
else:
|
| 540 |
+
input_ids, labels, attn = item
|
| 541 |
input_ids, labels, attn = input_ids.to(device), labels.to(device), attn.to(device)
|
| 542 |
out = model(input_ids=input_ids, attention_mask=attn, labels=labels,
|
| 543 |
output_ponder=True)
|
|
|
|
| 633 |
ap.add_argument("--max-steps", type=int, default=8, help="思考步数上限 K")
|
| 634 |
ap.add_argument("--epsilon", type=float, default=0.01)
|
| 635 |
ap.add_argument("--prior-p", type=float, default=0.4, help="几何先验 p_g (越小越鼓励多想)")
|
| 636 |
+
ap.add_argument("--difficulty-prior", default=None,
|
| 637 |
+
help="难度条件化先验, 如 easy:0.7,medium:0.4,hard:0.15 —— "
|
| 638 |
+
"同难度样本组成同 batch, KL 先验按难度调整 "
|
| 639 |
+
"(hard 用小 p_g 鼓励多想; 需数据带 difficulty 字段)")
|
| 640 |
ap.add_argument("--beta", type=float, default=0.05, help="KL 正则权重")
|
| 641 |
ap.add_argument("--head-bias-init", type=float, default=-1.0, help="初始停机偏置, sigmoid(-1)≈0.27")
|
| 642 |
ap.add_argument("--start-layer", type=int, default=None, help="思考块起始层 (默认最后8层)")
|
|
|
|
| 670 |
if tok.pad_token_id is None:
|
| 671 |
tok.pad_token = tok.eos_token
|
| 672 |
ds = PonderDataset(rows, tok, max_len=args.max_len, text_field=args.text_field)
|
| 673 |
+
prior_map = None
|
| 674 |
+
if args.difficulty_prior:
|
| 675 |
+
prior_map = {}
|
| 676 |
+
for part in args.difficulty_prior.split(","):
|
| 677 |
+
k, v = part.split(":")
|
| 678 |
+
prior_map[k.strip()] = float(v)
|
| 679 |
+
n_diff = sum(1 for ex in ds.examples if ex.get("difficulty") in prior_map)
|
| 680 |
+
print(f"难度条件化先验: {prior_map} | 命中样本 {n_diff}/{len(ds.examples)}")
|
| 681 |
+
if n_diff == 0:
|
| 682 |
+
raise SystemExit("--difficulty-prior 需要数据带 difficulty 字段 (如 make_gsm8k_zh_data.py / --make-mix-data 的产物)")
|
| 683 |
+
loader = PonderLoader(ds, args.batch_size, tok.pad_token_id or 1, args.seed)
|
| 684 |
+
else:
|
| 685 |
+
loader = DataLoader(ds, batch_size=args.batch_size, shuffle=True,
|
| 686 |
+
collate_fn=lambda b: collate(b, tok.pad_token_id or 1))
|
| 687 |
|
| 688 |
# ---- 模型 ----
|
| 689 |
dtype = resolve_dtype(args.dtype, args.device)
|
|
|
|
| 710 |
print(f"--- epoch {ep + 1}/{math.ceil(args.epochs)} ---")
|
| 711 |
gs = train_one_epoch(model, loader, optimizer, scheduler, args.device,
|
| 712 |
accum=args.accum, clip=args.clip,
|
| 713 |
+
log_every=args.log_every, global_step=gs, scaler=scaler,
|
| 714 |
+
prior_map=prior_map)
|
| 715 |
|
| 716 |
save_model(model, args, args.output)
|
| 717 |
print("完成。推理加载示例:")
|