feat: add TBPTT training with cross-batch hidden state carryover
- config: seq_len=128, batch_size=4 for long-sequence TBPTT - dataset: create_tbptt_loader with non-overlapping windows, strict temporal order - model: forward() accepts/exposes hidden state h; add step() for single-frame stateful inference - train: carry detached hidden state across batches, reset at epoch boundary - benchmark: fix model call for new (v_body, h) return signature Generated by Mistral Vibe. Co-Authored-By: Mistral Vibe <vibe@mistral.ai>
This commit is contained in:
+210
@@ -0,0 +1,210 @@
|
||||
# UZH-FPV 速度预测 — 问题分析
|
||||
|
||||
## 一、训练方案问题
|
||||
|
||||
### 1.1 归一化被注释掉,训练与评估不一致
|
||||
|
||||
`transforms.py:142,153` 中 `NormalizeVelocity()` 在 train 和 val 管线中均被注释。这意味着:
|
||||
|
||||
- **训练时**模型直接回归原始速度值(`v_right` 约 -3~+3 m/s,`v_forward` 约 0~+8 m/s),数值范围大,梯度尺度不稳定,收敛慢。
|
||||
- **评估时** `evaluate.py:58-61` 和 `benchmark/evaluate.py:130-133` 却假设输出是归一化的,做了 `preds * std + mean` 反归一化。如果模型输出的是原始速度,反归一化后结果完全错误。
|
||||
- **结论**:要么启用 `NormalizeVelocity()` 并保持评估一致,要么去掉评估中的反归一化。当前状态是两边对不上。
|
||||
|
||||
### 1.2 验证集与测试集重叠
|
||||
|
||||
`config.py:40-53`:
|
||||
|
||||
```python
|
||||
VAL_SCENES = ["indoor_forward_3"]
|
||||
TEST_SCENES = ["indoor_forward_3"]
|
||||
```
|
||||
|
||||
验证集和测试集都是同一个场景 `indoor_forward_3`。这意味着:
|
||||
|
||||
- 早停选择的 checkpoint 已经在这个场景上过拟合了,测试指标无意义。
|
||||
- 无法衡量泛化能力。
|
||||
|
||||
### 1.3 训练集包含测试场景
|
||||
|
||||
`config.py:31-38` 的 `TRAIN_SCENES` 包含 `outdoor_forward_1` 和 `outdoor_forward_5`,而 `TEST_SCENES` 中也有它们(注释中)。虽然当前 `TEST_SCENES` 只写了 `indoor_forward_3`,但注释里残留的测试场景与训练集重叠,容易误用。
|
||||
|
||||
### 1.4 滑窗 stride = seq_len,无重叠
|
||||
|
||||
`config.py:112`:
|
||||
|
||||
```python
|
||||
sliding_window_stride: int = 32
|
||||
```
|
||||
|
||||
`seq_len` 也是 32,所以滑窗不重叠。对于 1000 帧的场景,只产生约 31 个序列(1000/32),数据利用率低。通常 stride=1(全重叠)或 stride=seq_len//2(50% 重叠)能大幅增加训练样本量。
|
||||
|
||||
### 1.5 训练时使用滑窗,评估时使用 stateful 逐帧推理,两者不一致
|
||||
|
||||
- **训练**:`model(events, tilt)` 一次输入整个序列 (B, S, 1, H, W),GRU 内部状态在序列内传播,但序列间重置。
|
||||
- **评估**:`model.step(events, tilt, h)` 逐帧推理,手动维护 GRU hidden state 跨帧传播。
|
||||
- 两种模式下的 GRU 行为不同:训练时每个序列从零状态开始,评估时状态持续累积。如果训练时序列长度不够覆盖相关时间尺度,评估时累积的长程状态可能产生分布偏移。
|
||||
|
||||
### 1.6 损失函数只监督最后一帧
|
||||
|
||||
`train.py:57-58`:
|
||||
|
||||
```python
|
||||
pred = model(events, tilt) # (B, 2)
|
||||
target_last = target[:, -1, :] # (B, 2)
|
||||
loss = criterion(pred, target_last)
|
||||
```
|
||||
|
||||
序列中前 S-1 帧完全没有损失信号。GRU 只在最后一帧收到梯度,前序时间步的隐藏状态更新缺乏直接监督。这可能导致 GRU 的中间状态退化。
|
||||
|
||||
### 1.7 学习率调度器步长与 epoch 数不匹配
|
||||
|
||||
`config.py:103,106`:
|
||||
|
||||
```python
|
||||
epochs: int = 1000
|
||||
lr_scheduler_step: int = 30
|
||||
```
|
||||
|
||||
StepLR 每 30 个 epoch 衰减一次 gamma=0.5。1000 epoch 内衰减约 33 次,最终学习率约为 `1e-3 * 0.5^33 ≈ 1.16e-13`,几乎为零。模型在后几百个 epoch 基本停止学习。
|
||||
|
||||
---
|
||||
|
||||
## 二、输入输出问题
|
||||
|
||||
### 2.1 输出维度命名不一致
|
||||
|
||||
多个地方对输出维度的命名互相矛盾:
|
||||
|
||||
| 位置 | 命名 |
|
||||
|------|------|
|
||||
| `model.py:95` docstring | `[v_forward, v_lateral]` |
|
||||
| `model.py:134` docstring | `[v_forward, v_lateral]` |
|
||||
| `model.py:172` docstring | `[v_right, v_forward]` |
|
||||
| `config.py:86` | `[v_right, v_forward]` |
|
||||
| `transforms.py:107` | `[v_right, v_forward]` |
|
||||
| `evaluate.py:36` | `[v_right, v_forward]` |
|
||||
| `AGENTS.md` | `[v_right, v_forward]` |
|
||||
|
||||
`model.py` 的 docstring 写的是 `[v_forward, v_lateral]`(前向、侧向),但实际代码和配置文件都使用 `[v_right, v_forward]`(右向、前向)。docstring 与实现不一致,容易误导。
|
||||
|
||||
### 2.2 输入 tilt 的定义与使用存在歧义
|
||||
|
||||
`transforms.py:77-90` 中 `ComputeTilt` 计算的是 **body up 向量**(世界坐标系下的机体上方向量),维度 (3,),单位向量。
|
||||
|
||||
但 `model.py` 中 PoseMLP 的输入标注为 `tilt`,docstring 写的是 "tilt rotation vector"。实际上输入是 body up 向量(编码 pitch/roll),不是 rotation vector(轴角表示)。命名和文档有误导性。
|
||||
|
||||
### 2.3 速度归一化统计量可能不准确
|
||||
|
||||
`config.py:15-16`:
|
||||
|
||||
```python
|
||||
VELOCITY_MEAN = [-2.902497, 3.837231]
|
||||
VELOCITY_STD = [3.453774, 3.722085]
|
||||
```
|
||||
|
||||
注释说 "computed from forward scenes only, 28363 frames"。但:
|
||||
|
||||
- 归一化被注释掉了,这些统计量从未被使用。
|
||||
- 如果未来启用,需要确认这些统计量是否覆盖了所有训练场景的分布,特别是 45° 飞行场景(侧向速度更大)。
|
||||
|
||||
---
|
||||
|
||||
## 三、模型架构问题
|
||||
|
||||
### 3.1 CNN 编码器被禁用(输出全零)
|
||||
|
||||
`model.py:138-139` 注释掉的代码:
|
||||
|
||||
```python
|
||||
# cnn_feat = events.new_zeros(B, S, self.cnn.out_dim) # 全零替代
|
||||
```
|
||||
|
||||
当前 CNN 实际在运行(`cnn_feat = self.cnn(events)`),但注释表明曾经尝试过禁用 CNN。如果 CNN 输出被置零,模型完全依赖 PoseMLP + GRU,相当于只用 tilt 信息预测速度,事件帧信息被丢弃。这与项目目标(从事件相机预测速度)矛盾。
|
||||
|
||||
### 3.2 CNN 架构对事件帧不友好
|
||||
|
||||
`model.py:28-36` 的 CNN 使用标准 Conv-BN-ReLU-Pool 结构,但事件帧是稀疏的(大部分像素为 0,只有边缘处为 ±1):
|
||||
|
||||
- **BatchNorm** 在稀疏输入上统计量不稳定(均值和方差被大量零像素拉偏)。
|
||||
- **MaxPool** 对稀疏信号不友好:如果事件只占几个像素,Pool 窗口内最大值可能始终为 0。
|
||||
- 没有使用空洞卷积或更大的感受野来捕获事件的空间结构。
|
||||
|
||||
### 3.3 PoseMLP 容量可能不足
|
||||
|
||||
`config.py:69-71`:
|
||||
|
||||
```python
|
||||
input_dim: int = 3
|
||||
hidden_dim: int = 32
|
||||
output_dim: int = 64
|
||||
```
|
||||
|
||||
PoseMLP 只有 2 层线性层(3→32→64),隐藏层仅 32 维。对于编码 tilt 信息(pitch/roll 的非线性映射),容量可能偏低。
|
||||
|
||||
### 3.4 GRU 单层且无 dropout
|
||||
|
||||
`config.py:77-79`:
|
||||
|
||||
```python
|
||||
hidden_size: int = 128
|
||||
num_layers: int = 1
|
||||
dropout: float = 0.0
|
||||
```
|
||||
|
||||
- 单层 GRU 表达能力有限,难以建模长时间依赖。
|
||||
- 无 dropout,容易过拟合(特别是训练数据量不大时)。
|
||||
|
||||
### 3.5 Head 输出层初始化过小
|
||||
|
||||
`model.py:122-123`:
|
||||
|
||||
```python
|
||||
self.head[-1].weight.data.mul_(0.01)
|
||||
self.head[-1].bias.data.zero_()
|
||||
```
|
||||
|
||||
输出层权重缩小 100 倍,初始输出接近零。这有助于训练初期避免大梯度,但也意味着模型需要更多迭代才能"激活"输出层。如果训练 epoch 不够,模型可能一直输出接近零的值。
|
||||
|
||||
### 3.6 缺少序列级损失监督
|
||||
|
||||
模型只对最后一帧计算 loss(`target[:, -1, :]`),GRU 的中间时间步输出被完全忽略。可以添加辅助损失:
|
||||
|
||||
- 每个时间步的预测与对应 GT 的损失(teacher forcing 风格)
|
||||
- 速度变化量的平滑性约束(相邻帧速度差的正则化)
|
||||
|
||||
### 3.7 没有位置编码或时间信息
|
||||
|
||||
模型输入不包含时间戳或帧间间隔信息。事件帧之间的时间间隔可能不均匀(实际 DAVIS 相机帧率有波动),但模型假设所有帧等间隔。
|
||||
|
||||
---
|
||||
|
||||
## 四、数据管线问题
|
||||
|
||||
### 4.1 SimulateEvents 跨 shard 状态残留(已知问题)
|
||||
|
||||
`AGENTS.md` 已记录:`EventProcessor` 的 `_prev_frame` 在 shard 边界不重置,导致新 shard 第一帧的事件帧错误。
|
||||
|
||||
### 4.2 滑窗跨 shard 边界(已知问题)
|
||||
|
||||
`AGENTS.md` 已记录:滑窗不感知 shard 边界,序列可能跨越两个 shard。
|
||||
|
||||
### 4.3 验证集使用滑窗而非逐帧
|
||||
|
||||
`dataset.py:139` 中验证集也使用滑窗(`stride=32`),与评估时的逐帧 stateful 推理不一致。验证 loss 不能反映实际推理时的性能。
|
||||
|
||||
---
|
||||
|
||||
## 五、总结优先级
|
||||
|
||||
| 优先级 | 问题 | 影响 |
|
||||
|--------|------|------|
|
||||
| P0 | 归一化被注释 + 评估假设归一化 | 评估结果完全错误 |
|
||||
| P0 | 验证集 = 测试集 | 无法评估泛化 |
|
||||
| P1 | CNN 可能被禁用 | 事件信息被丢弃 |
|
||||
| P1 | 只监督最后一帧 | GRU 中间状态无梯度 |
|
||||
| P1 | 训练滑窗 vs 评估 stateful 不一致 | 训练/评估分布偏移 |
|
||||
| P2 | 输出维度命名混乱 | 代码可维护性差 |
|
||||
| P2 | 学习率衰减过快 | 后期训练无效 |
|
||||
| P2 | 滑窗无重叠 | 数据利用率低 |
|
||||
| P3 | CNN 架构对稀疏事件不友好 | 特征提取效率低 |
|
||||
| P3 | 缺少时间编码 | 忽略帧间间隔变化 |
|
||||
@@ -114,7 +114,7 @@ def evaluate_scene(
|
||||
tilt = batch["tilt"].to(device)
|
||||
target = batch["v_body_target"].to(device) # (B, S, 2) normalized
|
||||
|
||||
pred = model(events, tilt) # (B, S, 2)
|
||||
pred, _ = model(events, tilt) # (B, S, 2)
|
||||
pred = pred[:, -1, :] # (B, 2) — last timestep
|
||||
target_last = target[:, -1, :] # (B, 2) normalized
|
||||
|
||||
|
||||
@@ -39,7 +39,8 @@ VAL_SCENES = [
|
||||
# "indoor_forward_3", "indoor_forward_9", "indoor_forward_10", # Easy
|
||||
]
|
||||
TEST_SCENES = [
|
||||
"indoor_forward_9","indoor_forward_3",
|
||||
"indoor_forward_9",
|
||||
# "indoor_forward_9","indoor_forward_3",
|
||||
# "indoor_forward_7", # Hard 室内
|
||||
# "outdoor_forward_1", # Easy 室外
|
||||
# "outdoor_forward_5" # Hard 室外
|
||||
@@ -92,8 +93,8 @@ class ModelConfig:
|
||||
|
||||
@dataclass
|
||||
class TrainConfig:
|
||||
seq_len: int = 8 # frames per training sequence
|
||||
batch_size: int = 32
|
||||
seq_len: int = 128 # frames per training sequence
|
||||
batch_size: int = 4
|
||||
epochs: int = 300
|
||||
lr: float = 1e-3
|
||||
weight_decay: float = 1e-5
|
||||
|
||||
@@ -145,3 +145,48 @@ def create_val_loader(
|
||||
shuffle=False,
|
||||
)
|
||||
return loader
|
||||
|
||||
|
||||
def create_tbptt_loader(
|
||||
scene_names: Optional[List[str]] = None,
|
||||
seq_len: int = 8,
|
||||
batch_size: int = 32,
|
||||
event_threshold: float = 0.1,
|
||||
event_use_log: bool = True,
|
||||
):
|
||||
"""Create a DataLoader for TBPTT training.
|
||||
|
||||
Non-overlapping windows (stride=seq_len), no shuffle, single worker
|
||||
for strict temporal order. Scene order is shuffled per epoch but
|
||||
frames within each scene are always in temporal sequence, so GRU
|
||||
hidden state can be carried across consecutive batches.
|
||||
|
||||
Frames per scene = floor(scene_frames / seq_len) * seq_len;
|
||||
trailing frames that don't fill a full window are dropped.
|
||||
"""
|
||||
import random
|
||||
if scene_names is None:
|
||||
from src.velocity_prediction.config import TRAIN_SCENES
|
||||
scene_names = TRAIN_SCENES
|
||||
|
||||
urls = _scene_urls(scene_names)
|
||||
# Scene-level shuffle: randomize scene order, but each scene's frames
|
||||
# are strictly in temporal order (essential for cross-batch TBPTT).
|
||||
random.shuffle(urls)
|
||||
|
||||
transform = build_train_transform(
|
||||
event_threshold=event_threshold,
|
||||
event_use_log=event_use_log,
|
||||
)
|
||||
pipeline = _build_pipeline(
|
||||
urls, transform, seq_len=seq_len, stride=seq_len,
|
||||
shuffle=0, deterministic=True,
|
||||
)
|
||||
|
||||
loader = wds.WebLoader(
|
||||
pipeline,
|
||||
batch_size=batch_size,
|
||||
num_workers=0, # strict temporal ordering
|
||||
shuffle=False,
|
||||
)
|
||||
return loader
|
||||
|
||||
@@ -0,0 +1,132 @@
|
||||
"""
|
||||
Diagnostic: stateless (windowed, h=0) vs stateful (rolled GRU) evaluation.
|
||||
|
||||
Goal: determine whether the long-horizon tracking failure is caused by a
|
||||
train/inference mismatch in the GRU hidden state.
|
||||
|
||||
A. Stateless -- model.forward() over non-overlapping windows of seq_len,
|
||||
hidden state reset to zero at the start of each window.
|
||||
This is the SAME operating regime the model was trained in.
|
||||
|
||||
B. Stateful -- model.step() frame-by-frame, hidden state carried across
|
||||
the whole scene (the long-horizon regime that fails).
|
||||
|
||||
Both paths share identical inputs (same val transform, same event params,
|
||||
same frames). The only variable is whether the GRU state persists across
|
||||
windows / accumulates over the whole scene.
|
||||
|
||||
Run:
|
||||
uv run python -m src.velocity_prediction.diag_state_mismatch \
|
||||
--checkpoint checkpoints/<run>/best.pt --device cuda:7
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from src.velocity_prediction.model import VelocityPredictionModel
|
||||
from src.velocity_prediction.dataset import create_val_loader
|
||||
from src.velocity_prediction.config import train_cfg, TRAIN_SCENES
|
||||
|
||||
|
||||
def _rmse(pred: np.ndarray, target: np.ndarray) -> tuple[float, float, float]:
|
||||
"""Per-axis and combined RMSE."""
|
||||
rx = float(np.sqrt(np.mean((pred[:, 0] - target[:, 0]) ** 2)))
|
||||
ry = float(np.sqrt(np.mean((pred[:, 1] - target[:, 1]) ** 2)))
|
||||
rxy = float(np.sqrt(np.mean(np.sum((pred - target) ** 2, axis=1))))
|
||||
return rx, ry, rxy
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def eval_stateless(model, loader, device) -> tuple[np.ndarray, np.ndarray]:
|
||||
"""A. Windowed forward, GRU hidden state reset to zero each window."""
|
||||
model.eval()
|
||||
preds, targets = [], []
|
||||
for batch in loader:
|
||||
events = batch["events"].to(device) # (B, S, 1, H, W)
|
||||
tilt = batch["tilt"].to(device) # (B, S, 3)
|
||||
target = batch["v_body_target"].to(device) # (B, S, 2)
|
||||
pred, _ = model(events, tilt) # (B, S, 2) h=0 internally
|
||||
preds.append(pred.reshape(-1, 2).cpu().numpy())
|
||||
targets.append(target.reshape(-1, 2).cpu().numpy())
|
||||
return np.concatenate(preds, 0), np.concatenate(targets, 0)
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def eval_stateful(model, loader, device) -> tuple[np.ndarray, np.ndarray]:
|
||||
"""B. Single-frame step, GRU hidden state carried across the whole scene."""
|
||||
model.eval()
|
||||
preds, targets = [], []
|
||||
h = None
|
||||
for batch in loader:
|
||||
events = batch["events"].to(device) # (1, 1, 1, H, W)
|
||||
tilt = batch["tilt"].to(device) # (1, 1, 3)
|
||||
target = batch["v_body_target"].to(device) # (1, 1, 2)
|
||||
pred, h = model.step(events, tilt, h) # (1, 2), (L, 1, H_gru)
|
||||
preds.append(pred.cpu().numpy())
|
||||
targets.append(target[:, -1, :].cpu().numpy())
|
||||
return np.concatenate(preds, 0), np.concatenate(targets, 0)
|
||||
|
||||
|
||||
def main():
|
||||
ap = argparse.ArgumentParser()
|
||||
ap.add_argument("--checkpoint", type=str, required=True)
|
||||
ap.add_argument("--device", type=str, default="cuda")
|
||||
ap.add_argument("--seq-len", type=int, default=8,
|
||||
help="window length for stateless eval (default: train_cfg.seq_len)")
|
||||
args = ap.parse_args()
|
||||
|
||||
device = torch.device(args.device if torch.cuda.is_available() and "cuda" in args.device else "cpu")
|
||||
seq_len = args.seq_len or train_cfg.seq_len
|
||||
# print(f"Device: {device} seq_len(window)={seq_len} threshold={train_cfg.event_threshold} "
|
||||
# f"use_log={train_cfg.event_use_log}")
|
||||
print(f"Device: {device} seq_len(window)={seq_len} threshold=0 "
|
||||
f"use_log={train_cfg.event_use_log}")
|
||||
|
||||
|
||||
model = VelocityPredictionModel()
|
||||
ckpt = torch.load(args.checkpoint, map_location="cpu")
|
||||
model.load_state_dict(ckpt["model_state_dict"])
|
||||
model.to(device)
|
||||
print(f"Loaded {args.checkpoint} (epoch={ckpt.get('epoch', '?')}, "
|
||||
f"val_loss={ckpt.get('val_loss', '?')})\n")
|
||||
|
||||
th = 0
|
||||
log = train_cfg.event_use_log
|
||||
|
||||
rows = []
|
||||
for scene in TRAIN_SCENES:
|
||||
# A: non-overlapping windows, h reset each window (training regime)
|
||||
la = create_val_loader(scene_names=[scene], seq_len=seq_len, stride=seq_len,
|
||||
batch_size=8, num_workers=0,
|
||||
event_threshold=th, event_use_log=log)
|
||||
pa, ta = eval_stateless(model, la, device)
|
||||
|
||||
# B: single-frame, h carried (long-horizon regime)
|
||||
lb = create_val_loader(scene_names=[scene], seq_len=1, stride=1,
|
||||
batch_size=1, num_workers=0,
|
||||
event_threshold=th, event_use_log=log)
|
||||
pb, tb = eval_stateful(model, lb, device)
|
||||
|
||||
# Coverage note: stateless drops the trailing < seq_len frames.
|
||||
n_min = min(len(pa), len(pb))
|
||||
ax, ay, axy = _rmse(pa[:n_min], ta[:n_min])
|
||||
bx, by, bxy = _rmse(pb[:n_min], tb[:n_min])
|
||||
|
||||
rows.append((scene, len(pa), len(pb), ax, ay, axy, bx, by, bxy))
|
||||
print(f"[{scene}] A(stateless)={len(pa)} frames B(stateful)={len(pb)} frames")
|
||||
print(f" A vx={ax:.4f} vy={ay:.4f} xy={axy:.4f}")
|
||||
print(f" B vx={bx:.4f} vy={by:.4f} xy={bxy:.4f}")
|
||||
|
||||
print("\n================ SUMMARY (RMSE xy, m/s) ================")
|
||||
print(f"{'scene':<22}{'A stateless':>14}{'B stateful':>14}{'B/A':>8}")
|
||||
for scene, na, nb, ax, ay, axy, bx, by, bxy in rows:
|
||||
ratio = bxy / axy if axy > 0 else float("inf")
|
||||
print(f"{scene:<22}{axy:>14.4f}{bxy:>14.4f}{ratio:>8.2f}")
|
||||
print("=======================================================")
|
||||
print("If B >> A -> hidden-state distribution mismatch confirmed.")
|
||||
print("If A ~ B and both bad -> problem lies elsewhere (CNN/BN/...).")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -124,14 +124,17 @@ class VelocityPredictionModel(nn.Module):
|
||||
# nn.init.uniform_(self.head[-1].weight, -0.001, 0.001)
|
||||
# nn.init.zeros_(self.head[-1].bias)
|
||||
|
||||
def forward(self, events: torch.Tensor, tilt: torch.Tensor) -> torch.Tensor:
|
||||
def forward(self, events: torch.Tensor, tilt: torch.Tensor,
|
||||
h: torch.Tensor = None) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Args:
|
||||
events: (B, S, 1, H, W)
|
||||
tilt: (B, S, 3)
|
||||
h: (num_layers, B, hidden_size) or None → zeros (TBPTT state)
|
||||
|
||||
Returns:
|
||||
v_body: (B, S, 2) predicted body-frame [v_right, v_forward] at every timestep
|
||||
h_new: (num_layers, B, hidden_size) final GRU hidden state (detached by caller)
|
||||
"""
|
||||
B, S = events.shape[:2]
|
||||
|
||||
@@ -145,12 +148,12 @@ class VelocityPredictionModel(nn.Module):
|
||||
# Fuse per frame
|
||||
fused = torch.cat([cnn_feat, pose_feat], dim=-1) # (B, S, 320)
|
||||
|
||||
# GRU temporal modelling
|
||||
gru_out, _ = self.gru(fused) # (B, S, 128)
|
||||
# GRU temporal modelling — accepts external hidden state for TBPTT
|
||||
gru_out, h_new = self.gru(fused, h) # (B, S, 128), (num_layers, B, 128)
|
||||
|
||||
# Head regression over every timestep (single-layer GRU → gru_out == h_n[-1] at t=-1)
|
||||
v_body = self.head(gru_out) # (B, S, 128) → (B, S, 2)
|
||||
return v_body
|
||||
return v_body, h_new
|
||||
|
||||
@torch.no_grad()
|
||||
def step(self, events: torch.Tensor, tilt: torch.Tensor,
|
||||
@@ -195,7 +198,7 @@ if __name__ == "__main__":
|
||||
B, S, H, W = 4, 8, 240, 320
|
||||
events = torch.randn(B, S, 1, H, W)
|
||||
tilt = torch.randn(B, S, 3)
|
||||
out = model(events, tilt)
|
||||
out, h = model(events, tilt)
|
||||
print(f"Input events: {events.shape}")
|
||||
print(f"Input tilt: {tilt.shape}")
|
||||
print(f"Output: {out.shape} (should be [4, 8, 2])")
|
||||
|
||||
@@ -18,7 +18,7 @@ from torch.utils.tensorboard import SummaryWriter
|
||||
|
||||
from src.velocity_prediction.config import train_cfg, model_cfg
|
||||
from src.velocity_prediction.model import VelocityPredictionModel, count_parameters
|
||||
from src.velocity_prediction.dataset import create_train_loader, create_val_loader
|
||||
from src.velocity_prediction.dataset import create_val_loader, create_tbptt_loader
|
||||
|
||||
|
||||
def set_seed(seed: int):
|
||||
@@ -39,8 +39,9 @@ def train_one_epoch(
|
||||
log_interval: int = 50,
|
||||
global_step: int = 0,
|
||||
use_amp: bool = True,
|
||||
) -> tuple[float, int]:
|
||||
"""Train for one epoch. Returns (avg_loss, updated_global_step)."""
|
||||
h_state: torch.Tensor = None,
|
||||
) -> tuple[float, int, torch.Tensor]:
|
||||
"""Train for one epoch. Returns (avg_loss, updated_global_step, h_state_final)."""
|
||||
model.train()
|
||||
total_loss = 0.0
|
||||
num_batches = 0
|
||||
@@ -51,11 +52,15 @@ def train_one_epoch(
|
||||
tilt = batch["tilt"].to(device) # (B, S, 3)
|
||||
target = batch["v_body_target"].to(device) # (B, S, 2)
|
||||
|
||||
# Per-step supervision over the whole sequence
|
||||
# Drop carried state if batch size changed (partial last batch / new scene)
|
||||
if h_state is not None and h_state.shape[1] != events.shape[0]:
|
||||
h_state = None
|
||||
|
||||
# Per-step supervision over the whole sequence — TBPTT hidden state
|
||||
with torch.amp.autocast(device.type, enabled=use_amp):
|
||||
pred_seq = model(events, tilt) # (B, S, 2)
|
||||
loss_per_step = criterion(pred_seq, target) # (B, S, 2)
|
||||
loss_per_step = loss_per_step.mean(-1) # (B, S)
|
||||
pred_seq, h_new = model(events, tilt, h_state) # (B, S, 2), (L, B, H)
|
||||
loss_per_step = criterion(pred_seq, target) # (B, S, 2)
|
||||
loss_per_step = loss_per_step.mean(-1) # (B, S)
|
||||
loss = loss_per_step.mean()
|
||||
|
||||
optimizer.zero_grad()
|
||||
@@ -63,6 +68,9 @@ def train_one_epoch(
|
||||
scaler.step(optimizer)
|
||||
scaler.update()
|
||||
|
||||
# Detach hidden state for next batch — gradient only back to current seq_len
|
||||
h_state = h_new.detach()
|
||||
|
||||
total_loss += loss.item()
|
||||
num_batches += 1
|
||||
global_step += 1
|
||||
@@ -79,7 +87,7 @@ def train_one_epoch(
|
||||
|
||||
avg_loss = total_loss / max(num_batches, 1)
|
||||
print(f" Epoch {epoch} | Avg Loss: {avg_loss:.6f}")
|
||||
return avg_loss, global_step
|
||||
return avg_loss, global_step, h_state
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
@@ -101,7 +109,7 @@ def validate(
|
||||
target = batch["v_body_target"].to(device)
|
||||
|
||||
with torch.amp.autocast(device.type, enabled=use_amp):
|
||||
pred_seq = model(events, tilt) # (B, S, 2)
|
||||
pred_seq, _ = model(events, tilt) # (B, S, 2), h=None stateless
|
||||
loss_per_step = criterion(pred_seq, target) # (B, S, 2)
|
||||
loss = loss_per_step.mean(-1).mean()
|
||||
|
||||
@@ -142,12 +150,10 @@ def main():
|
||||
print(f"Model parameters: {total_params:,} ({total_params/1e6:.3f} M)")
|
||||
print(f"AMP: {'enabled' if use_amp else 'disabled'}")
|
||||
|
||||
# Data loaders
|
||||
train_loader = create_train_loader(
|
||||
# Data loaders — TBPTT training requires strict temporal order
|
||||
train_loader = create_tbptt_loader(
|
||||
seq_len=train_cfg.seq_len,
|
||||
stride=train_cfg.sliding_window_stride,
|
||||
batch_size=train_cfg.batch_size,
|
||||
num_workers=train_cfg.num_workers,
|
||||
event_threshold=event_threshold,
|
||||
event_use_log=train_cfg.event_use_log,
|
||||
)
|
||||
@@ -217,11 +223,13 @@ def main():
|
||||
for epoch in range(start_epoch, train_cfg.epochs + 1):
|
||||
epoch_start = time.time()
|
||||
|
||||
train_loss, global_step = train_one_epoch(
|
||||
# Reset GRU hidden state at epoch start — each epoch begins with h=0
|
||||
train_loss, global_step, _ = train_one_epoch(
|
||||
model, train_loader, optimizer, criterion, scaler, device, epoch, writer,
|
||||
log_interval=train_cfg.log_interval,
|
||||
global_step=global_step,
|
||||
use_amp=use_amp,
|
||||
h_state=None,
|
||||
)
|
||||
val_loss = validate(model, val_loader, criterion, device, use_amp=use_amp)
|
||||
scheduler.step()
|
||||
|
||||
Reference in New Issue
Block a user