tchbcb commited on
Commit
3e2da9b
·
verified ·
1 Parent(s): 4ca75d6

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 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
- self.examples.append({"input_ids": ids, "labels": labels})
 
 
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, (input_ids, labels, attn) in enumerate(loader):
 
 
 
 
 
 
 
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
- loader = DataLoader(ds, batch_size=args.batch_size, shuffle=True,
629
- collate_fn=lambda b: collate(b, tok.pad_token_id or 1))
 
 
 
 
 
 
 
 
 
 
 
 
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("完成。推理加载示例:")