docs: add TODO.md

Generated by Mistral Vibe.
Co-Authored-By: Mistral Vibe <vibe@mistral.ai>
This commit is contained in:
2026-07-09 15:50:25 +08:00
parent 56aa10a503
commit 5ccd3df874
+109
View File
@@ -0,0 +1,109 @@
# 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 分布正常