From d1d3310543fd92d94d5d6d9ff649923ca05bb20e Mon Sep 17 00:00:00 2001 From: CaoWangrenbo Date: Wed, 29 Jul 2026 16:32:41 +0800 Subject: [PATCH] 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 --- analyse.md | 210 ++++++++++++++++++ benchmark/evaluate.py | 2 +- src/velocity_prediction/config.py | 7 +- src/velocity_prediction/dataset.py | 45 ++++ .../diag_state_mismatch.py | 132 +++++++++++ src/velocity_prediction/model.py | 13 +- src/velocity_prediction/train.py | 36 +-- 7 files changed, 422 insertions(+), 23 deletions(-) create mode 100644 analyse.md create mode 100644 src/velocity_prediction/diag_state_mismatch.py diff --git a/analyse.md b/analyse.md new file mode 100644 index 0000000..df8104c --- /dev/null +++ b/analyse.md @@ -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 | 缺少时间编码 | 忽略帧间间隔变化 | diff --git a/benchmark/evaluate.py b/benchmark/evaluate.py index 23582cf..bb641be 100644 --- a/benchmark/evaluate.py +++ b/benchmark/evaluate.py @@ -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 diff --git a/src/velocity_prediction/config.py b/src/velocity_prediction/config.py index 2ba2000..e6a7429 100644 --- a/src/velocity_prediction/config.py +++ b/src/velocity_prediction/config.py @@ -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 diff --git a/src/velocity_prediction/dataset.py b/src/velocity_prediction/dataset.py index 364a189..a2ea6c4 100644 --- a/src/velocity_prediction/dataset.py +++ b/src/velocity_prediction/dataset.py @@ -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 diff --git a/src/velocity_prediction/diag_state_mismatch.py b/src/velocity_prediction/diag_state_mismatch.py new file mode 100644 index 0000000..0a2b807 --- /dev/null +++ b/src/velocity_prediction/diag_state_mismatch.py @@ -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//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() diff --git a/src/velocity_prediction/model.py b/src/velocity_prediction/model.py index d127b6e..38c6cc3 100644 --- a/src/velocity_prediction/model.py +++ b/src/velocity_prediction/model.py @@ -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])") diff --git a/src/velocity_prediction/train.py b/src/velocity_prediction/train.py index e70277c..0864526 100644 --- a/src/velocity_prediction/train.py +++ b/src/velocity_prediction/train.py @@ -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()