diff --git a/benchmark/evaluate.py b/benchmark/evaluate.py index e886771..23582cf 100644 --- a/benchmark/evaluate.py +++ b/benchmark/evaluate.py @@ -114,7 +114,8 @@ def evaluate_scene( tilt = batch["tilt"].to(device) target = batch["v_body_target"].to(device) # (B, S, 2) normalized - pred = model(events, tilt) # (B, 2) normalized + pred = model(events, tilt) # (B, S, 2) + pred = pred[:, -1, :] # (B, 2) — last timestep target_last = target[:, -1, :] # (B, 2) normalized all_preds.append(pred.cpu().numpy()) diff --git a/src/velocity_prediction/model.py b/src/velocity_prediction/model.py index b21f514..d127b6e 100644 --- a/src/velocity_prediction/model.py +++ b/src/velocity_prediction/model.py @@ -86,13 +86,13 @@ class PoseMLP(nn.Module): class VelocityPredictionModel(nn.Module): """ - Full model: CNN + PoseMLP → concat → GRU → Head → [vx, vy]. + Full model: CNN + PoseMLP → concat → GRU → Head → [v_right, v_forward]. Input: events: (B, S, 1, H, W) tilt: (B, S, 3) Output: - v_body: (B, 2) — body-frame [v_forward, v_lateral] for the last frame in the sequence + v_body: (B, S, 2) — body-frame [v_right, v_forward] for each frame in the sequence """ def __init__(self, cnn_cfg=model_cfg.cnn, pose_cfg=model_cfg.pose_mlp, @@ -131,8 +131,10 @@ class VelocityPredictionModel(nn.Module): tilt: (B, S, 3) Returns: - v_body: (B, 2) predicted body-frame [v_forward, v_lateral] at the last timestep + v_body: (B, S, 2) predicted body-frame [v_right, v_forward] at every timestep """ + B, S = events.shape[:2] + # Per-frame encoding cnn_feat = self.cnn(events) # (B, S, 256) # B, S = events.shape[:2] @@ -144,13 +146,10 @@ class VelocityPredictionModel(nn.Module): fused = torch.cat([cnn_feat, pose_feat], dim=-1) # (B, S, 320) # GRU temporal modelling - gru_out, h_n = self.gru(fused) # gru_out: (B, S, 128), h_n: (1, B, 128) + gru_out, _ = self.gru(fused) # (B, S, 128) - # Use last hidden state - last_hidden = h_n[-1] # (B, 128) - - # Head regression - v_body = self.head(last_hidden) # (B, 2) + # 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 @torch.no_grad() @@ -199,4 +198,4 @@ if __name__ == "__main__": out = model(events, tilt) print(f"Input events: {events.shape}") print(f"Input tilt: {tilt.shape}") - print(f"Output: {out.shape} (should be [4, 2])") + 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 8e22fef..e70277c 100644 --- a/src/velocity_prediction/train.py +++ b/src/velocity_prediction/train.py @@ -51,11 +51,12 @@ def train_one_epoch( tilt = batch["tilt"].to(device) # (B, S, 3) target = batch["v_body_target"].to(device) # (B, S, 2) - # Predict velocity for the last frame in the sequence + # Per-step supervision over the whole sequence with torch.amp.autocast(device.type, enabled=use_amp): - pred = model(events, tilt) # (B, 2) - target_last = target[:, -1, :] # (B, 2) - loss = criterion(pred, target_last) + 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) + loss = loss_per_step.mean() optimizer.zero_grad() scaler.scale(loss).backward() @@ -70,6 +71,11 @@ def train_one_epoch( elapsed = time.time() - start_time print(f" Epoch {epoch} | Batch {batch_idx} | Loss: {loss.item():.6f} | {elapsed:.1f}s") writer.add_scalar("train/loss_batch", loss.item(), global_step) + # Per-step loss quantiles — observe GRU stabilization across timesteps + S = loss_per_step.shape[1] + for i in range(0, S, max(1, S // 8)): + writer.add_scalar(f"train/loss_step_{i:02d}", + loss_per_step[:, i].mean().item(), global_step) avg_loss = total_loss / max(num_batches, 1) print(f" Epoch {epoch} | Avg Loss: {avg_loss:.6f}") @@ -95,9 +101,9 @@ def validate( target = batch["v_body_target"].to(device) with torch.amp.autocast(device.type, enabled=use_amp): - pred = model(events, tilt) - target_last = target[:, -1, :] - loss = criterion(pred, target_last) + pred_seq = model(events, tilt) # (B, S, 2) + loss_per_step = criterion(pred_seq, target) # (B, S, 2) + loss = loss_per_step.mean(-1).mean() total_loss += loss.item() num_batches += 1 @@ -165,8 +171,8 @@ def main(): step_size=train_cfg.lr_scheduler_step, gamma=train_cfg.lr_scheduler_gamma, ) - # criterion = nn.SmoothL1Loss() - criterion = nn.MSELoss() + # criterion = nn.SmoothL1Loss(reduction='none') + criterion = nn.MSELoss(reduction='none') # ── Resume from checkpoint ──────────────────────────────────── start_epoch = 1