Files
hexone2086 5ccd3df874 docs: add TODO.md
Generated by Mistral Vibe.
Co-Authored-By: Mistral Vibe <vibe@mistral.ai>
2026-07-09 15:50:25 +08:00

3.8 KiB

TODO — 训练/评估改进计划


P1 · 截断反传(TBPTT):修 GRU 隐状态长程失配(新)

问题诊断(2025-07-09)

diag_state_mismatch.py 确认核心问题是 GRU 隐状态训练/推理分布失配:

checkpoint A 无状态(h=0 每 8 帧重置) B 有状态(h 持续滚动) B/A
epoch 4 (val 7.91) RMSE 0.14~4.0 RMSE 1.6~4.3 1.1~1.4
epoch 180 (val 9.17) RMSE 0.10~0.98 RMSE 2.5~5.4 2.8~38.9

训练越久,分叉越大。epoch 180 模型在窗口内几乎背下数据(A=0.14),但一滚动就崩(B=5.4)。训练 loss 越低,长程越差

根因:训练时每个序列从 h=0 开始(seq_len=8),GRU 只见「零后 ≤8 步」的状态分布。推理时 h 滚动数百帧,进入 OOD 区域,预测崩溃。val_loss 从 7.91 升到 9.17 进一步印证过拟合。

修法:TBPTT

跨 batch 传递 detach 后的隐状态,让 GRU 暴露于长程状态:

batch_k:     frames [k,  k+7]  → GRU(fused, h_prev) → h_curr  (detach → batch_{k+1})
batch_{k+1}: frames [k+8, k+15] → GRU(fused, h_detached) → h_next  (detach → ...)

具体改动

  1. model.py::VelocityPredictionModel.forward

    • 签名: forward(self, events, tilt, h=None) -> (v_body, h_new)
    • gru_out, h_new = self.gru(fused, h)
    • model.step 保持不变
  2. train.py::train_one_epoch

    • 接受 h_state 参数,初值为 None
    • pred_seq, h_new = model(events, tilt, h_state)
    • h_state = h_new.detach() 传出供下一 batch
    • 每个 epoch 重置 h_state = None
  3. train.py::validate

    • 保持原样传 h=None(每 batch 独立窗口,与当前评估一致)
  4. dataset.py

    • 当前 create_train_loader 使用 shuffle=1000 → 序列打乱 → 无法跨 batch 传 h
    • 方案:新增 create_tbptt_loader:
      • shuffle=0, deterministic=True
      • stride=seq_len(非重叠)
      • num_workers=0(确保严格时序)
      • 场景级 shuffle:在加载前打乱场景列表,但每个场景内帧序严格
    • 直接将原始 _build_pipeline 参数由 stride=1 改为 stride=seq_len,关闭 shuffle

注意事项

  • 梯度只回传到当前 batch 的 seq_len 步,不回溯更早 batch(detach 保证)
  • 首 batch h=None 仍从零初始化
  • BN 不受影响(CNN 在 per-frame 维度独立运行)
  • TBPTT loader 帧数 = floor(scene_frames / seq_len) * seq_len,末尾不足部分丢弃
  • 需起新 run_id(旧 checkpoint 全部在 h=0 重置模式下训练,不兼容)

验证

  • diag_state_mismatch.py 重新评估:预期 B/A < 1.5
  • 同一场景下 stateful RMSE 接近 stateless RMSE

P2 · 考虑中(未纳入本次)

2.1 归一化被注释

确认训练/评估两端 self-consistent(输出+target 都是原始 m/s),暂保留。

2.2 验证集/测试集划分

当前划分可调,涉及 config.py 场景列表调整和重训。

2.3 速度平滑正则

((pred[:, 1:] - pred[:, :-1]) ** 2).mean() 加权 lambda=0.01。建议 TBPTT 收敛后再加。

2.4 Loss mask 前几步

TBPTT 后 GRU 状态连贯,此问题可能自然缓解,暂不动。


落地顺序

  1. model.py::forward 签名,返回 (pred, h_new)
  2. train.py train + validate 适配新签名
  3. 新增/修改 loader 实现 TBPTT 模式(不 shuffle,非重叠)
  4. 起新 run_id 训练,观察 val_loss 走势
  5. diag_state_mismatch.py 验证 B/A 比值

已完成

以下为旧版 TODO 中已完成的项目,保留归档。

P1 · 序列级损失监督(已完成,2025-07)

已在 train.py:54-59 + model.py:149-153 实现:

  • model.forward 返回 (B, S, 2),对整段 gru_out 过 head
  • 全序列逐 step MSE loss (B, S, 2) → (B, S) → scalar
  • TensorBoard loss_step_{i:02d} 分位数记录
  • 验证通过:输出 shape (B, 8, 2),TensorBoard 观察 loss_step 分布正常