diff --git a/src/camera-based-e2e/create_vocab.py b/src/camera-based-e2e/create_vocab.py index 94fca63..b355638 100644 --- a/src/camera-based-e2e/create_vocab.py +++ b/src/camera-based-e2e/create_vocab.py @@ -13,8 +13,8 @@ def get_all_trajectories(): # Instantiate dataloader w/ n_items None to get everything dataset = WaymoE2E(indexFile='index_train.pkl', - data_dir='/anvil/scratch/x-mgagvani/wod/waymo_end_to_end_camera_v1_0_0/waymo_open_dataset_end_to_end_camera_v_1_0_0', - n_items=None # set to None eventually, + data_dir='/anvil/scratch/x-mgagvani/wod/waymo_end_to_end_camera_v1_0_0/waymo_open_dataset_end_to_end_camera_v_1_0_0', + n_items=None # set to None eventually ) dataloader = DataLoader(dataset, batch_size=512, num_workers=0, collate_fn=collate_with_images, pin_memory=True) diff --git a/src/camera-based-e2e/loader.py b/src/camera-based-e2e/loader.py index 45984a4..6401750 100644 --- a/src/camera-based-e2e/loader.py +++ b/src/camera-based-e2e/loader.py @@ -21,6 +21,7 @@ def __init__( ): self.data_dir = data_dir self.seed = seed + self._index_file = indexFile self.filename = "" self.file = None @@ -101,7 +102,7 @@ def collate_with_images(batch): import time from tqdm import tqdm # NOTE: Replace with your path - DATA_DIR = '/anvil/scratch/x-mgagvani/wod/waymo_end_to_end_camera_v1_0_0/waymo_open_dataset_end_to_end_camera_v_1_0_0' + DATA_DIR = '/scratch/gilbreth/mathur91/waymo/waymo_open_dataset_end_to_end_camera_v_1_0_0' BATCH_SIZE = 256 dataset = WaymoE2E(indexFile="index_train.pkl", data_dir = DATA_DIR) loader = DataLoader( diff --git a/src/camera-based-e2e/models/base_model.py b/src/camera-based-e2e/models/base_model.py index d6cf3fc..6f29f1c 100644 --- a/src/camera-based-e2e/models/base_model.py +++ b/src/camera-based-e2e/models/base_model.py @@ -1,3 +1,5 @@ +from dataclasses import asdict +from typing import List, Optional import torch import torch.nn as nn import torch.nn.functional as F @@ -6,6 +8,7 @@ from dataclasses import asdict, is_dataclass from .losses.depth_loss import DepthLoss +from .proposal_planner import IPadConfig class BaseModel(nn.Module): def __init__(self, in_dim, out_dim): @@ -22,31 +25,34 @@ def forward(self, x: dict) -> torch.Tensor: return self.nn(x) class LitModel(pl.LightningModule): - def __init__(self, model: nn.Module, lr: float, lr_vision: float | None = None, rfs_weight: float = 0.0): + def __init__( + self, + model: nn.Module, + lr: float, + lr_vision: Optional[float] = None, + ipad_config: Optional[IPadConfig] = None, + ): super(LitModel, self).__init__() self.model = model - - # NVJPEG fall back if we are running on Negishi AMD GPU self.has_nvjpeg = True - # If we are using ScorerModel, which has a cfg, then save the attributes of the cfg as hparams, so they go into wandb - cfg = getattr(model, "cfg", None) - if cfg is None: + # Collect model's own config (if any) for logging/hparams + _model_cfg = getattr(model, "cfg", None) + if _model_cfg is None: cfg_dict = {} - elif is_dataclass(cfg): - cfg_dict = asdict(cfg) - elif isinstance(cfg, dict): - cfg_dict = dict(cfg) + elif is_dataclass(_model_cfg): + cfg_dict = asdict(_model_cfg) + elif isinstance(_model_cfg, dict): + cfg_dict = dict(_model_cfg) else: try: - cfg_dict = dict(vars(cfg)) + cfg_dict = dict(vars(_model_cfg)) except TypeError: - cfg_dict = {"repr": repr(cfg)} + cfg_dict = {"repr": repr(_model_cfg)} hparams = { "lr": lr, "lr_vision": lr_vision, - "rfs_weight": rfs_weight, "model_name": model.__class__.__name__, "model_cfg": cfg_dict, } @@ -54,6 +60,11 @@ def __init__(self, model: nn.Module, lr: float, lr_vision: float | None = None, if isinstance(v, (int, float, str, bool)) or v is None: hparams[f"model_cfg_{k}"] = v + # iPad-style training hyperparameters + ipad_cfg = ipad_config if ipad_config is not None else IPadConfig() + for field, value in asdict(ipad_cfg).items(): + hparams[field] = value + self.example_input_array = ({ 'PAST': torch.zeros((1, 16, 6)), # PAST 'IMAGES': [torch.zeros((1, 3, 1280, 1920)) for _ in range(6)], # IMAGES @@ -92,10 +103,10 @@ def transfer_batch_to_device(self, batch, device, dataloader_idx): def decode_batch_jpeg( self, - images_jpeg: list[list[torch.Tensor]], - device: torch.device | None = None, - ) -> list[torch.Tensor | None]: - cam_idxs_used = tuple(getattr(self.model.cfg, "cam_idxs_used", range(len(images_jpeg)))) + images_jpeg: List[List[torch.Tensor]], + device: torch.device = None, + ) -> List[torch.Tensor]: + cam_idxs_used = tuple(getattr(getattr(self.model, "cfg", None), "cam_idxs_used", range(len(images_jpeg)))) decode_device = self.device if device is None else device selected = [(cam_idx, images_jpeg[cam_idx]) for cam_idx in cam_idxs_used] @@ -215,6 +226,145 @@ def _prepare_rfs_inputs(self, past, future, pred_future): t_idx = torch.tensor([3.0, 5.0], device=future.device).unsqueeze(0).expand(future.size(0), -1) return pred_slice, gt_slice, lng_dir_slice, lat_dir_slice, speed, t_idx + # ---- NAVSIM-style quality target ---- + @torch.no_grad() + def _compute_navsim_score( + self, + proposals: torch.Tensor, + future: torch.Tensor, + ) -> torch.Tensor: + """ + Approximate NAVSIM Eq. 5: S = NC * DAC * (5*EP + 5*TTC + 2*Comf) / 12 + + Without agent / map data we set NC=1, DAC=1, TTC=1 and compute EP and + Comf from trajectory geometry alone. + + Args: + proposals: (B, K, T, 2) predicted trajectories + future: (B, T, 2) ground-truth trajectory + Returns: + quality: (B, K) in [0, 1] + """ + B, K, T, _ = proposals.shape + gt = future[:, None, :, :] # (B, 1, T, 2) + + # --- Ego Progress (EP) --- + gt_disp = future[:, -1] - future[:, 0] # (B, 2) + gt_dist = gt_disp.norm(dim=-1, keepdim=True).clamp(min=1e-3) # (B, 1) + gt_dir = gt_disp / gt_dist # (B, 2) + + prop_disp = proposals[:, :, -1] - proposals[:, :, 0] # (B, K, 2) + progress = (prop_disp * gt_dir.unsqueeze(1)).sum(dim=-1) # (B, K) + ep = (progress / gt_dist).clamp(0.0, 1.0) # (B, K) + + # --- Comfort (Comf) --- + dt = 0.25 # 4 Hz + vel = (proposals[:, :, 1:] - proposals[:, :, :-1]) / dt # (B,K,T-1,2) + acc = (vel[:, :, 1:] - vel[:, :, :-1]) / dt # (B,K,T-2,2) + jerk = (acc[:, :, 1:] - acc[:, :, :-1]) / dt # (B,K,T-3,2) + jerk_mag = jerk.norm(dim=-1) # (B,K,T-3) + + jerk_thresh = getattr(self.hparams, "comfort_jerk_threshold", 5.0) + comf = (jerk_mag < jerk_thresh).float().mean(dim=-1) # (B,K) + + nc = 1.0 + dac = 1.0 + ttc = 1.0 + quality = nc * dac * (5.0 * ep + 5.0 * ttc + 2.0 * comf) / 12.0 # (B,K) + + return quality.clamp(0.0, 1.0) + + @torch.no_grad() + def _compute_rfs_quality( + self, + proposals: torch.Tensor, + reference: torch.Tensor, + past: torch.Tensor, + ) -> torch.Tensor: + """ + Per-proposal RFS-style quality in [0, 1] for BCE scorer targets (iPad Eq. 4). + + Same longitudinal/lateral deviation + speed scaling as ``rfs_loss``, evaluated at + 3 s and 5 s (indices 11, 19 at 4 Hz). ``reference`` is typically batch ``FUTURE`` + (expert trajectory); it can be swapped for a route or rollout proxy when GT is absent. + + Optionally multiplies by a jerk comfort factor (same spirit as NAVSIM Comf in Eq. 5). + """ + device = proposals.device + indices = [11, 19] + if reference.shape[1] <= max(indices): + raise ValueError( + f"reference horizon {reference.shape[1]} must exceed max RFS index {max(indices)}" + ) + + speed = torch.norm(past[..., 2:4], dim=-1)[:, -1] + full_lng_dir, full_lat_dir = self.compute_direction(reference) + + ref_slice = reference[:, indices, :] + lng_slice = full_lng_dir[:, indices, :] + lat_slice = full_lat_dir[:, indices, :] + + prop_slice = proposals[:, :, indices, :] + delta = prop_slice - ref_slice.unsqueeze(1) + + delta_lng = (delta * lng_slice.unsqueeze(1)).sum(dim=-1).abs() + delta_lat = (delta * lat_slice.unsqueeze(1)).sum(dim=-1).abs() + + t_idx = torch.tensor([3.0, 5.0], device=device).unsqueeze(0).expand(reference.size(0), -1) + tau_lat_raw, tau_lng_raw = self.time_thresholds(t_idx) + scale = self.speed_scale(speed) + if scale.dim() == 1: + scale = scale.unsqueeze(1) + tau_lat = tau_lat_raw * scale + tau_lng = tau_lng_raw * scale + + deviation = torch.max( + delta_lat / tau_lat[:, None, :], + delta_lng / tau_lng[:, None, :], + ) + score = torch.where( + deviation <= 1, + torch.ones_like(deviation), + torch.pow(0.1, deviation - 1), + ) + rfs_quality = score.mean(dim=-1) + + use_comf = getattr(self.hparams, "rfs_target_use_comfort", True) + if use_comf: + dt = 0.25 + vel = (proposals[:, :, 1:] - proposals[:, :, :-1]) / dt + acc = (vel[:, :, 1:] - vel[:, :, :-1]) / dt + jerk = (acc[:, :, 1:] - acc[:, :, :-1]) / dt + jerk_mag = jerk.norm(dim=-1) + jerk_thresh = getattr(self.hparams, "comfort_jerk_threshold", 5.0) + comf = (jerk_mag < jerk_thresh).float().mean(dim=-1) + rfs_quality = rfs_quality * comf + + return rfs_quality.clamp(0.0, 1.0) + + @torch.no_grad() + def _compute_quality_target( + self, + pred: torch.Tensor, + gt_expanded: torch.Tensor, + future: torch.Tensor, + tau: float, + past: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + """ + Returns (B, K) quality target in [0, 1] based on score_target_type. + """ + target_type = getattr(self.hparams, "score_target_type", "l1") + if target_type == "navsim": + return self._compute_navsim_score(pred, future) + if target_type == "rfs": + if past is None: + raise ValueError("score_target_type='rfs' requires past trajectory") + return self._compute_rfs_quality(pred, future, past) + + l1_target = (pred - gt_expanded).abs().sum(dim=-1).mean(dim=-1) # (B, K) + return torch.exp(-l1_target / max(tau, 1e-6)) + # ---- optimizers ---- def configure_optimizers(self): # NOTE: This can be extended and tuned, LR especially will differ and have an impact. @@ -253,25 +403,25 @@ def _shared_step(self, batch: torch.Tensor, stage: str) -> torch.Tensor: images = self.decode_batch_jpeg(images_jpeg) else: raise KeyError("Batch must contain either 'IMAGES_JPEG' or 'IMAGES' key.") - - # `past` is our input (B, 16, 6) e.g. Batch x Time x (x, y, v_x, v_y, a_x, a_y) - # and `future` is our output (B, 20, 2) e.g. Batch x Time x (x, y) - # create all input data that we are allowed to give to a model model_inputs = {'PAST': past, 'IMAGES': images, 'INTENT': intent} - pred_future = self.forward(model_inputs) # (B, T*2) + raw_output = self.forward(model_inputs) pred_depth = None pred_scores: torch.Tensor = None pred_traj_flat: torch.Tensor = None query_for_score: torch.Tensor = None - if isinstance(pred_future, dict): - outputs = pred_future - pred_future = outputs["trajectory"] - pred_depth = outputs.get("depth", None) + proposal_list = None + if isinstance(raw_output, dict): + outputs = raw_output + proposal_list = outputs.get("proposal_list", None) pred_scores = outputs.get("scores", None) + pred_depth = outputs.get("depth", None) pred_traj_flat = outputs.get("trajectory_flat", None) query_for_score = outputs.get("query_for_score", None) + pred_future = outputs["trajectory"] + else: + pred_future = raw_output pred = pred_future t_steps = future.shape[1] @@ -285,72 +435,128 @@ def _shared_step(self, batch: torch.Tensor, stage: str) -> torch.Tensor: k_modes = pred.shape[1] // t2 pred = pred.view(pred.size(0), k_modes, t_steps, 2) + if not torch.isfinite(pred).all(): + pred = torch.nan_to_num(pred, nan=0.0, posinf=1e3, neginf=-1e3) + + # ---- MoN L1 trajectory loss (iPad-style) ---- + # min over N proposals of mean L1 displacement per timestep + gt = future[:, None, :, :] # (B, 1, T, 2) + + if proposal_list is not None and len(proposal_list) > 0: + # Discounted intermediate supervision: L = sum_k λ^(K-1-k) * MoN_L1(P_k) + prev_w = getattr(self.hparams, "prev_weight", 0.1) + loss_traj = torch.tensor(0.0, device=self.device) + for proposals_k in proposal_list: + if not torch.isfinite(proposals_k).all(): + proposals_k = torch.nan_to_num(proposals_k, nan=0.0, posinf=1e3, neginf=-1e3) + l1_per_mode = (proposals_k - gt).abs().sum(dim=-1).mean(dim=-1) # (B, K) + mon_l1 = l1_per_mode.amin(dim=1).mean() + loss_traj = prev_w * loss_traj + mon_l1 + else: + l1_per_mode = (pred - gt).abs().sum(dim=-1).mean(dim=-1) # (B, K) + loss_traj = l1_per_mode.amin(dim=1).mean() + + # ---- Metrics (for logging, not loss) ---- + refine_oracle_ades: List[torch.Tensor] = [] + with torch.no_grad(): + dist = torch.norm(pred - gt, dim=-1) # (B, K, T) L2 + ade_per_mode = dist.mean(dim=-1) # (B, K) + oracle_ade = ade_per_mode.min(dim=1).values.mean() + ade_pred = None + if pred_scores is not None and k_modes > 1: + pred_idx = pred_scores.detach().argmax(dim=1) + ade_pred = ade_per_mode[torch.arange(pred.size(0), device=pred.device), pred_idx].mean() + elif k_modes == 1: + ade_pred = ade_per_mode.squeeze(1).mean() + regret = (ade_pred - oracle_ade) if ade_pred is not None else None + + if proposal_list is not None and len(proposal_list) > 0: + for proposals_k in proposal_list: + pk = proposals_k + if not torch.isfinite(pk).all(): + pk = torch.nan_to_num(pk, nan=0.0, posinf=1e3, neginf=-1e3) + dist_r = torch.norm(pk - gt, dim=-1) + ade_pm = dist_r.mean(dim=-1) + refine_oracle_ades.append(ade_pm.min(dim=1).values.mean()) + + # ---- RFS loss ---- if pred_scores is not None and pred.size(1) > 1: - rfs_pred_idx = pred_scores.argmin(dim=1) + rfs_pred_idx = pred_scores.detach().argmax(dim=1) else: rfs_pred_idx = torch.zeros(pred.size(0), dtype=torch.long, device=pred.device) pred_for_rfs = pred[torch.arange(pred.size(0), device=pred.device), rfs_pred_idx] pred_slice, gt_slice, lng_dir_slice, lat_dir_slice, speed, t_idx = self._prepare_rfs_inputs( - past, - future, - pred_for_rfs, + past, future, pred_for_rfs, ) rfs_unweighted = self.rfs_loss(pred_slice, gt_slice, lng_dir_slice, lat_dir_slice, speed, t_idx) rfs_weight = getattr(self.hparams, "rfs_weight", 0.0) loss_rfs = rfs_weight * rfs_unweighted - loss_type = getattr(self.hparams, "model_cfg_loss_type", "mse") - - # ADE per mode: (B, K) - dist = torch.norm(pred - future[:, torch.newaxis, :, :], dim=-1) # (B, K, T) - ade_per_mode = dist.mean(dim=-1) - - # Top-M WTA for trajectory loss. Here, we have an "oracle" that picks the best mode - # so, our loss is calculated on the mean of the top n trajectories. - top_m = min(getattr(self.hparams, "model_cfg_loss_top_n", 5), ade_per_mode.size(1)) - loss_ade = ade_per_mode.topk(top_m, largest=False, dim=1).values.mean() - - # oracle ade is best of all proposals, since we have the GT data during training - oracle_ade = ade_per_mode.min(dim=1).values.mean() - ade_pred = None - # pred_scores is now the predicted ADE of each trajectory / expectation loss - if pred_scores is not None and k_modes > 1: - pred_idx = pred_scores.argmin(dim=1) - ade_pred = ade_per_mode[torch.arange(pred.size(0), device=pred.device), pred_idx].mean() - elif k_modes == 1: - ade_pred = ade_per_mode.squeeze(1).mean() - regret = (ade_pred - oracle_ade) if ade_pred is not None else None - - # Scorer Losses -> encourage ranking of predicted scores to match true ranking of ades that are generated + # ---- Score loss ---- if k_modes > 1 and pred_scores is not None: - ade = ade_per_mode.detach() # (B, K) - if loss_type == "mse": - loss_score = F.mse_loss(pred_scores, ade) - elif loss_type == "reinforce": - tau_base = getattr(self.hparams, "model_cfg_loss_tau_base", 1.0) - decay_factor = getattr(self.hparams, "model_cfg_loss_tau_decay", 0.95) - entropy_lambda = getattr(self.hparams, "model_cfg_loss_entropy_lambda", 0.01) - # p_k = softmax(s_k / tau) - logits = -pred_scores / max(0.1, tau_base * (decay_factor ** self.current_epoch)) - p = F.softmax(logits, dim=1) - # Loss = sum_k {p_k * e_k} (e.g., expectation of e_k) - loss_selection = (p * ade).sum(dim=1).mean() - # H(p) = -sum(p) * log(p) - entropy = -(p * (p + 1e-8).log()).sum(dim=1).mean() - loss_score = loss_selection - entropy_lambda * entropy - elif loss_type == "margin": - s_diff = pred_scores.unsqueeze(2) - pred_scores.unsqueeze(1) # (B, K, K): s_i - s_j - ade_diff = ade.unsqueeze(2) - ade.unsqueeze(1) # (B, K, K): a_i - a_j - # target=1 when i is worse (ade_i > ade_j), want s_i > s_j - target = (ade_diff > 0).float() - weight = ade_diff.abs() - per_pair = F.binary_cross_entropy_with_logits(s_diff, target, reduction='none') - loss_score = (per_pair * weight).sum() / (weight.sum() + 1e-8) + if not torch.isfinite(pred_scores).all(): + pred_scores = torch.nan_to_num(pred_scores, nan=0.0, posinf=1e3, neginf=-1e3) + + # Legacy ScorerModel path (mse/reinforce/margin) when model.cfg specifies a loss_type + model_loss_type = getattr(self.hparams, "model_cfg_loss_type", None) + if model_loss_type is not None: + ade = ade_per_mode.detach() + if model_loss_type == "mse": + loss_score = F.mse_loss(pred_scores, ade) + elif model_loss_type == "reinforce": + tau_base = getattr(self.hparams, "model_cfg_loss_tau_base", 1.0) + decay_factor = getattr(self.hparams, "model_cfg_loss_tau_decay", 0.95) + entropy_lambda = getattr(self.hparams, "model_cfg_loss_entropy_lambda", 0.01) + logits = -pred_scores / max(0.1, tau_base * (decay_factor ** self.current_epoch)) + p = F.softmax(logits, dim=1) + loss_selection = (p * ade).sum(dim=1).mean() + entropy = -(p * (p + 1e-8).log()).sum(dim=1).mean() + loss_score = loss_selection - entropy_lambda * entropy + elif model_loss_type == "margin": + s_diff = pred_scores.unsqueeze(2) - pred_scores.unsqueeze(1) + ade_diff = ade.unsqueeze(2) - ade.unsqueeze(1) + target = (ade_diff > 0).float() + weight = ade_diff.abs() + per_pair = F.binary_cross_entropy_with_logits(s_diff, target, reduction='none') + loss_score = (per_pair * weight).sum() / (weight.sum() + 1e-8) + else: + raise NotImplementedError(f"Loss {model_loss_type} is not implemented") + pred_idx = pred_scores.argmin(dim=1) + ade_pred = ade_per_mode[torch.arange(pred.size(0), device=pred.device), pred_idx].mean() else: - raise NotImplementedError(f"Loss {loss_type} is not implemented") - pred_idx = pred_scores.argmin(dim=1) - ade_pred = ade_per_mode[torch.arange(pred.size(0), device=pred.device), pred_idx].mean() + # iPad-style BCE / CE / bce_pairwise / listnet scorer + score_loss_type = getattr(self.hparams, "score_loss_type", "bce") + tau = getattr(self.hparams, "score_temperature", 5.0) + with torch.no_grad(): + quality_target = self._compute_quality_target(pred, gt, future, tau, past=past) + best_idx = quality_target.argmax(dim=1) + if score_loss_type == "ce": + loss_score = F.cross_entropy(pred_scores, best_idx) + elif score_loss_type == "listnet": + target_probs = F.softmax(quality_target / max(tau * 0.1, 1e-6), dim=1) + loss_score = F.kl_div( + F.log_softmax(pred_scores, dim=1), + target_probs, + reduction="batchmean", + ) + else: + bce_loss = F.binary_cross_entropy_with_logits(pred_scores, quality_target) + loss_score = bce_loss + if score_loss_type == "bce_pairwise": + margin = getattr(self.hparams, "score_margin", 0.2) + rank_weight = getattr(self.hparams, "score_rank_weight", 0.2) + topk = int(getattr(self.hparams, "score_topk", 0)) + best_scores = pred_scores.gather(1, best_idx.unsqueeze(1)) + pairwise_margin = margin - (best_scores - pred_scores) + pairwise_margin.scatter_(1, best_idx.unsqueeze(1), 0.0) + if topk > 0: + k_eff = min(topk, pairwise_margin.size(1) - 1) + hardest = pairwise_margin.topk(k_eff, dim=1).values + rank_loss = F.relu(hardest).mean() + else: + rank_loss = F.relu(pairwise_margin).mean() + loss_score = bce_loss + rank_weight * rank_loss else: loss_score = torch.tensor(0.0, device=self.device) @@ -410,20 +616,51 @@ def _shared_step(self, batch: torch.Tensor, stage: str) -> torch.Tensor: # Depth Loss if pred_depth is not None: - front_img = images[1] # front camera + front_img = images[1] depth_in = F.interpolate(front_img, size=(128, 128), mode='nearest') loss_depth = self.depth_loss(depth_in, pred_depth, loss_fn=F.l1_loss) else: loss_depth = torch.tensor(0.0, device=self.device) + loss_depth *= 0.1 + + # Score loss warmup + sw_score = getattr(self.hparams, "score_weight", 1.0) + warmup_epochs = getattr(self.hparams, "score_warmup_epochs", 2) + current_epoch = self.current_epoch if hasattr(self, "current_epoch") else 0 + if current_epoch < warmup_epochs: + effective_score_weight = 0.0 + elif current_epoch < warmup_epochs + 1: + effective_score_weight = sw_score * (current_epoch - warmup_epochs) + else: + effective_score_weight = sw_score + + loss_score *= effective_score_weight + total_loss = loss_traj + loss_depth + loss_score + loss_rfs + + # Smoothness / comfort losses on best proposal + loss_smooth = torch.tensor(0.0, device=self.device) + loss_collision = torch.tensor(0.0, device=self.device) + loss_comfort = torch.tensor(0.0, device=self.device) + sw = getattr(self.hparams, "smoothness_weight", 0.0) + cw = getattr(self.hparams, "collision_weight", 0.0) + cfw = getattr(self.hparams, "comfort_weight", 0.0) + if (sw > 0 or cw > 0 or cfw > 0) and pred_for_rfs.shape[1] >= 3: + vel = pred_for_rfs[:, 1:] - pred_for_rfs[:, :-1] + acc = vel[:, 1:] - vel[:, :-1] + jerk = acc[:, 1:] - acc[:, :-1] + if sw > 0: + loss_smooth = (jerk ** 2).mean() + if cfw > 0: + v_mid = vel[:, :-1] + a_mid = acc + v_speed = (v_mid ** 2).sum(dim=-1).clamp(min=0.25).sqrt() + a_mag = (a_mid ** 2).sum(dim=-1).sqrt() + curv = a_mag / (v_speed ** 2) + loss_comfort = curv.mean() + total_loss = total_loss + sw * loss_smooth + cw * loss_collision + cfw * loss_comfort - loss_depth *= 0.1 # slightly enabled - loss_ade *= 1.0 # TODO: tune loss terms - loss_score *= 1.0 - adv_lambda = float(getattr(self.hparams, "model_cfg_adv_lambda", 0.1)) - total_loss = loss_ade + loss_depth + loss_score + loss_rfs + (adv_lambda * loss_adv) - # TODO: improve logging both to disk and to console log_payload = { - f"{stage}_loss_ade": loss_ade, + f"{stage}_loss_traj": loss_traj, f"{stage}_loss_score": loss_score, f"{stage}_loss_depth": loss_depth, f"{stage}_loss_rfs": loss_rfs, @@ -431,10 +668,16 @@ def _shared_step(self, batch: torch.Tensor, stage: str) -> torch.Tensor: f"{stage}_rfs_unweighted": rfs_unweighted, f"{stage}_loss": total_loss, } + if sw > 0 or cw > 0 or cfw > 0: + log_payload[f"{stage}_loss_smooth"] = loss_smooth + log_payload[f"{stage}_loss_collision"] = loss_collision + log_payload[f"{stage}_loss_comfort"] = loss_comfort if ade_pred is not None: log_payload[f"{stage}_ade_pred"] = ade_pred log_payload[f"{stage}_ade_oracle"] = oracle_ade log_payload[f"{stage}_ade_regret"] = regret + for ri, oade in enumerate(refine_oracle_ades): + log_payload[f"{stage}_ade_oracle_refine_{ri}"] = oade log_payload.update(scorer_metrics) self.log_dict(log_payload, prog_bar=True, logger=True, batch_size=past.size(0), diff --git a/src/camera-based-e2e/models/feature_extractors.py b/src/camera-based-e2e/models/feature_extractors.py index 0864b85..9096d09 100644 --- a/src/camera-based-e2e/models/feature_extractors.py +++ b/src/camera-based-e2e/models/feature_extractors.py @@ -3,6 +3,31 @@ import torch.nn as nn import timm + +class ResNetFeatures(nn.Module): + """ResNet backbone (resnet50) - widely available in timm, fallback when DINO/SAM unavailable.""" + + def __init__(self, model_name: str = "resnet50", frozen: bool = True, feature_stage: int = -1): + super().__init__() + self.backbone = timm.create_model(model_name, pretrained=True, features_only=True) + self.data_config = timm.data.resolve_data_config(model=self.backbone) + self.transforms = timm.data.create_transform(**self.data_config, is_training=False) + if frozen: + for p in self.backbone.parameters(): + p.requires_grad = False + self.backbone.eval() + channels = self.backbone.feature_info.channels() + reductions = self.backbone.feature_info.reduction() + self.feature_stage = feature_stage + self.dims = [channels[feature_stage]] + self.patch_size = reductions[feature_stage] + + def forward(self, x: torch.Tensor) -> List[torch.Tensor]: + x_t = self.transforms(x.float().div(255.0)) + feats = self.backbone(x_t) + return [feats[self.feature_stage]] + + class DINOFeatures(nn.Module): def __init__(self, model_name: str = "vit_small_plus_patch16_dinov3.lvd1689m", frozen: bool = True): super(DINOFeatures, self).__init__() diff --git a/src/camera-based-e2e/models/proposal_init.py b/src/camera-based-e2e/models/proposal_init.py new file mode 100644 index 0000000..a0f141d --- /dev/null +++ b/src/camera-based-e2e/models/proposal_init.py @@ -0,0 +1,63 @@ +""" +Proposal initialization (iPad-style). + +Per-timestep learnable embeddings for N proposals × T timesteps, +conditioned on ego status (past trajectory + intent). + +Output: bev_feature (B, N*T, C) — the initial BEV proposal queries. +""" +import torch +import torch.nn as nn +import torch.nn.functional as F + + +class ProposalInit(nn.Module): + """ + Initialize per-timestep BEV proposal features from learnable embeddings + plus ego status encoding. Matches iPad's init_feature + hist_encoding. + """ + + def __init__( + self, + d_model: int, + num_proposals: int = 16, + horizon: int = 20, + past_dim: int = 16 * 6, + intent_classes: int = 3, + ): + super().__init__() + self.d_model = d_model + self.num_proposals = num_proposals + self.horizon = horizon + + # Per-(proposal, timestep) learnable embeddings + self.proposal_embed = nn.Parameter( + torch.zeros(1, num_proposals * horizon, d_model) + ) + nn.init.trunc_normal_(self.proposal_embed, std=0.1) + + ego_dim = past_dim + intent_classes + self.ego_enc = nn.Sequential( + nn.Linear(ego_dim, d_model), + nn.GELU(), + nn.Linear(d_model, d_model), + ) + + def forward(self, past: torch.Tensor, intent: torch.Tensor) -> torch.Tensor: + """ + Args: + past: (B, 16, 6) + intent: (B,) integer 1/2/3 + Returns: + bev_feature: (B, N*T, d_model) + """ + B = past.size(0) + past_flat = past.view(B, -1) + intent_onehot = F.one_hot( + (intent - 1).long().clamp(0, 2), num_classes=3 + ).float() + ego = torch.cat([intent_onehot, past_flat], dim=1) + ego_feat = self.ego_enc(ego) # (B, d_model) + + bev_feature = self.proposal_embed.expand(B, -1, -1) + ego_feat[:, None, :] + return bev_feature diff --git a/src/camera-based-e2e/models/proposal_planner.py b/src/camera-based-e2e/models/proposal_planner.py new file mode 100644 index 0000000..56e88d4 --- /dev/null +++ b/src/camera-based-e2e/models/proposal_planner.py @@ -0,0 +1,107 @@ +""" +Proposal-centric E2E planner (iPad-style). + +Pipeline: + scene encoder → proposal init → iterative refinement → scorer + (with intermediate proposal supervision) +""" +from dataclasses import dataclass +from typing import Dict, List + +import torch +import torch.nn as nn + +from .scene_encoder import SceneEncoder +from .proposal_init import ProposalInit +from .refinement import Refinement +from .scorer import Scorer + + +@dataclass +class IPadConfig: + """Loss / scorer hyperparameters for the iPad-style ProposalPlanner. + + Bundled here so the LitModel only needs a single config argument instead of + a long list of keyword args. Defaults match the original LitModel signature. + """ + + # Loss weights + rfs_weight: float = 0.0 + smoothness_weight: float = 0.0 + collision_weight: float = 0.0 + comfort_weight: float = 0.0 + diversity_weight: float = 0.0 + + # Scorer + score_weight: float = 1.0 + score_warmup_epochs: int = 2 + score_temperature: float = 5.0 + score_loss_type: str = "bce" # bce | ce | bce_pairwise | listnet + score_target_type: str = "l1" # l1 | navsim | rfs + score_rank_weight: float = 0.0 + score_margin: float = 0.2 + score_topk: int = 0 + + # Misc + comfort_jerk_threshold: float = 5.0 + prev_weight: float = 0.1 + rfs_target_use_comfort: bool = True + + +class ProposalPlanner(nn.Module): + + def __init__( + self, + backbone: nn.Module, + d_model: int = 256, + num_proposals: int = 16, + num_refinement_steps: int = 4, + horizon: int = 20, + num_cams: int = 8, + ): + super().__init__() + self.horizon = horizon + self.n_proposals = num_proposals + + self.scene_encoder = SceneEncoder(backbone, d_model=d_model, num_cams=num_cams) + self.proposal_init = ProposalInit( + d_model=d_model, + num_proposals=num_proposals, + horizon=horizon, + ) + self.refinement = Refinement( + d_model=d_model, + num_steps=num_refinement_steps, + num_heads=8, + num_proposals=num_proposals, + horizon=horizon, + ) + self.scorer = Scorer( + d_model=d_model, + num_proposals=num_proposals, + horizon=horizon, + ) + + def forward(self, x: Dict[str, torch.Tensor]) -> Dict[str, torch.Tensor]: + past = x["PAST"] + images: List[torch.Tensor] = x["IMAGES"] + intent = x["INTENT"] + + scene_feat = self.scene_encoder(images) + bev_feature = self.proposal_init(past, intent) + proposals, bev_feature, proposal_list = self.refinement(bev_feature, scene_feat) + scores = self.scorer(bev_feature) + + if getattr(self, "_debug", False): + self._debug_proposals = proposals.detach() + self._debug_scores = scores.detach() + self._debug_proposal_list = [p.detach() for p in proposal_list] + + B, K, T, _ = proposals.shape + trajectory_flat = proposals.view(B, K * T * 2) + + return { + "trajectory": trajectory_flat, + "scores": scores, + "proposal_list": proposal_list, + } diff --git a/src/camera-based-e2e/models/refinement.py b/src/camera-based-e2e/models/refinement.py new file mode 100644 index 0000000..2bd7561 --- /dev/null +++ b/src/camera-based-e2e/models/refinement.py @@ -0,0 +1,122 @@ +""" +Iterative refinement (iPad-style predict-attend-refine loop). + +A single shared RefinementBlock is applied K times. At each iteration: + 1. Decode per-timestep features into full trajectory proposals + 2. Encode proposals back into feature space (proposal-anchored) + 3. Cross-attend proposal features to scene tokens + 4. FFN update +Returns all intermediate proposals for discounted supervision. +""" +from typing import List, Tuple + +import torch +import torch.nn as nn + +from .blocks import MHA + + +class RefinementBlock(nn.Module): + """One iteration: decode proposals, encode them back, cross-attend to scene.""" + + def __init__( + self, + d_model: int, + num_heads: int = 8, + num_proposals: int = 16, + horizon: int = 20, + mlp_ratio: int = 4, + ): + super().__init__() + self.num_proposals = num_proposals + self.horizon = horizon + + # Per-timestep feature → (x, y) + self.traj_decoder = nn.Sequential( + nn.Linear(d_model, d_model * mlp_ratio), + nn.GELU(), + nn.Linear(d_model * mlp_ratio, 2), + ) + + # Encode (x, y) back into feature space + self.traj_enc = nn.Sequential( + nn.Linear(2, d_model), + nn.GELU(), + nn.Linear(d_model, d_model), + ) + + self.ln1 = nn.LayerNorm(d_model) + self.cross_attn = MHA(d_model, num_heads) + self.ln2 = nn.LayerNorm(d_model) + self.mlp = nn.Sequential( + nn.Linear(d_model, d_model * mlp_ratio), + nn.GELU(), + nn.Linear(d_model * mlp_ratio, d_model), + ) + + def forward( + self, + bev_feature: torch.Tensor, + scene_feat: torch.Tensor, + ) -> Tuple[torch.Tensor, torch.Tensor]: + """ + Args: + bev_feature: (B, N*T, C) per-timestep proposal features + scene_feat: (B, S, C) visual tokens from scene encoder + Returns: + proposals: (B, N, T, 2) fully predicted trajectories + bev_feature: (B, N*T, C) refined features + """ + B = bev_feature.size(0) + N, T = self.num_proposals, self.horizon + + proposals = self.traj_decoder(bev_feature).view(B, N, T, 2) + + prop_enc = self.traj_enc(proposals.view(B, N * T, 2)) + bev_feature = bev_feature + prop_enc + + bev_feature = bev_feature + self.cross_attn( + self.ln1(bev_feature), context=scene_feat + ) + bev_feature = bev_feature + self.mlp(self.ln2(bev_feature)) + + return proposals, bev_feature + + +class Refinement(nn.Module): + """Weight-shared iterative refinement applied num_steps times.""" + + def __init__( + self, + d_model: int, + num_steps: int = 4, + num_heads: int = 8, + num_proposals: int = 16, + horizon: int = 20, + ): + super().__init__() + self.num_steps = num_steps + self.block = RefinementBlock( + d_model=d_model, + num_heads=num_heads, + num_proposals=num_proposals, + horizon=horizon, + ) + + def forward( + self, + bev_feature: torch.Tensor, + scene_feat: torch.Tensor, + ) -> Tuple[torch.Tensor, torch.Tensor, List[torch.Tensor]]: + """ + Returns: + proposals: (B, N, T, 2) final iteration proposals + bev_feature: (B, N*T, C) final features + proposal_list: list of (B, N, T, 2) from each iteration + """ + proposal_list: List[torch.Tensor] = [] + proposals = None + for _ in range(self.num_steps): + proposals, bev_feature = self.block(bev_feature, scene_feat) + proposal_list.append(proposals) + return proposals, bev_feature, proposal_list diff --git a/src/camera-based-e2e/models/scene_encoder.py b/src/camera-based-e2e/models/scene_encoder.py new file mode 100644 index 0000000..144d2f5 --- /dev/null +++ b/src/camera-based-e2e/models/scene_encoder.py @@ -0,0 +1,60 @@ +""" +Multi-camera scene encoder for proposal-centric E2E planner. +Encodes all 8 Waymo cameras with a shared backbone and fuses via concatenation + camera embeddings. +""" +from typing import List + +import torch +import torch.nn as nn +import torch.nn.functional as F + + +class SceneEncoder(nn.Module): + """ + Encode multi-camera images into a single scene feature sequence. + Uses a shared backbone per camera, then concatenates tokens and adds camera embeddings. + """ + + def __init__( + self, + backbone: nn.Module, + d_model: int = 256, + num_cams: int = 8, + ): + super().__init__() + self.backbone = backbone + self.num_cams = num_cams + # Backbone output dim (from feature extractor .dims or .feature_dim) + if hasattr(backbone, "dims"): + self.backbone_dim = sum(backbone.dims) + else: + self.backbone_dim = getattr(backbone, "feature_dim", 384) + + self.proj = nn.Linear(self.backbone_dim, d_model) + # Per-camera embedding (1, num_cams, d_model) added to all tokens of that camera + self.cam_embed = nn.Parameter(torch.zeros(1, num_cams, d_model)) + nn.init.trunc_normal_(self.cam_embed, std=0.02) + self.d_model = d_model + self.ln = nn.LayerNorm(d_model) + + def forward(self, images: List[torch.Tensor]) -> torch.Tensor: + """ + Args: + images: List of (B, C, H, W) tensors, one per camera. len(images) == num_cams. + Returns: + scene_feat: (B, N, d_model) with N = num_cams * n_tokens_per_cam + """ + B = images[0].size(0) + all_tokens = [] + for c, img in enumerate(images): + with torch.no_grad(): + feats = self.backbone(img) + if isinstance(feats, (list, tuple)): + feats = feats[0] # (B, C, H, W) + # (B, C, H, W) -> (B, C, H*W) -> (B, H*W, C) + tokens = feats.flatten(2).permute(0, 2, 1) # (B, n_tokens, backbone_dim) + tokens = self.proj(tokens) + self.cam_embed[:, c : c + 1, :] # (B, n_tokens, d_model) + all_tokens.append(tokens) + # (B, num_cams * n_tokens_per_cam, d_model) + scene = torch.cat(all_tokens, dim=1) + return self.ln(scene) diff --git a/src/camera-based-e2e/models/scorer.py b/src/camera-based-e2e/models/scorer.py new file mode 100644 index 0000000..2ee9cbd --- /dev/null +++ b/src/camera-based-e2e/models/scorer.py @@ -0,0 +1,41 @@ +""" +Scorer (paper-faithful): max-pool per-timestep BEV features -> MLP -> score. +""" +import torch +import torch.nn as nn + + +class Scorer(nn.Module): + """ + Score each proposal from BEV proposal features. + Output (B, K) raw logits — higher = better. + """ + + def __init__( + self, + d_model: int, + num_proposals: int = 16, + horizon: int = 20, + hidden_dim: int = 1024, + ): + super().__init__() + self.num_proposals = num_proposals + self.horizon = horizon + self.score_mlp = nn.Sequential( + nn.Linear(d_model, hidden_dim), + nn.GELU(), + nn.Linear(hidden_dim, 1), + ) + + def forward(self, bev_feature: torch.Tensor) -> torch.Tensor: + """ + Args: + bev_feature: (B, N*T, C) per-timestep features + Returns: + scores: (B, N) raw logits, higher is better + """ + B = bev_feature.size(0) + N, T = self.num_proposals, self.horizon + feat = bev_feature.view(B, N, T, -1).amax(dim=2) # (B, N, C) + scores = self.score_mlp(feat).squeeze(-1) # (B, N) + return scores diff --git a/src/camera-based-e2e/score_experiment_viz.py b/src/camera-based-e2e/score_experiment_viz.py new file mode 100644 index 0000000..d0139e5 --- /dev/null +++ b/src/camera-based-e2e/score_experiment_viz.py @@ -0,0 +1,298 @@ +""" +Compare scorer-loss experiments across runs and auto-generate rankings. + +This script scans Lightning CSV logs under a logs root, extracts final validation +metrics per run, groups by scorer-loss configuration, and writes: + - summary CSV + - ranking bar chart + - oracle-vs-pred scatter (diagnose ranking bottleneck) + - regret-vs-top1 scatter (diagnose scorer discrimination) + - markdown report with "what went right/wrong" + +Usage: + python score_experiment_viz.py --logs_root /scratch/.../waymo/logs + python score_experiment_viz.py --logs_root /scratch/.../waymo/logs --latest_n 8 +""" + +import argparse +from pathlib import Path +from typing import Dict, List, Optional + +import matplotlib.pyplot as plt +import numpy as np +import pandas as pd + + +def _read_hparams_kv(hparams_path: Path) -> Dict[str, str]: + """ + Parse a few scalar keys from hparams.yaml without requiring strict YAML parse. + This is intentionally tolerant because some generated hparams files can contain + very large content blocks. + """ + keys = { + "score_loss_type", + "score_target_type", + "score_weight", + "score_temperature", + "score_rank_weight", + "score_margin", + "score_topk", + "model_type", + } + out: Dict[str, str] = {} + if not hparams_path.exists(): + return out + + try: + with hparams_path.open("r", errors="replace") as f: + for line in f: + line = line.strip() + if ":" not in line: + continue + k, v = line.split(":", 1) + k = k.strip() + if k in keys: + out[k] = v.strip().strip("'\"") + except Exception: + return out + return out + + +def _to_float(v: Optional[str], default: float) -> float: + if v is None or v == "": + return default + try: + return float(v) + except Exception: + return default + + +def _to_int(v: Optional[str], default: int) -> int: + if v is None or v == "": + return default + try: + return int(float(v)) + except Exception: + return default + + +def _find_runs(logs_root: Path) -> List[Path]: + runs = sorted(logs_root.glob("camera_e2e_*/version_0/metrics.csv")) + return runs + + +def _final_metric(df: pd.DataFrame, col: str) -> float: + if col not in df.columns: + return np.nan + s = df[col].dropna() + if s.empty: + return np.nan + return float(s.iloc[-1]) + + +def collect_run_table(logs_root: Path, latest_n: int = 0) -> pd.DataFrame: + rows = [] + metrics_files = _find_runs(logs_root) + if latest_n > 0: + metrics_files = metrics_files[-latest_n:] + + for mpath in metrics_files: + run_dir = mpath.parent + run_name = run_dir.parent.name + try: + df = pd.read_csv(mpath) + except Exception: + continue + + hp = _read_hparams_kv(run_dir / "hparams.yaml") + score_loss_type = hp.get("score_loss_type", "bce") + + row = { + "run_name": run_name, + "run_dir": str(run_dir), + "score_loss_type": score_loss_type, + "score_target_type": hp.get("score_target_type", "l1"), + "score_weight": _to_float(hp.get("score_weight"), 1.0), + "score_temperature": _to_float(hp.get("score_temperature"), 5.0), + "score_rank_weight": _to_float(hp.get("score_rank_weight"), 0.0), + "score_margin": _to_float(hp.get("score_margin"), 0.2), + "score_topk": _to_int(hp.get("score_topk"), 0), + "val_ade_pred": _final_metric(df, "val_ade_pred"), + "val_ade_oracle": _final_metric(df, "val_ade_oracle"), + "val_ade_regret": _final_metric(df, "val_ade_regret"), + "val_loss": _final_metric(df, "val_loss"), + "val_loss_score": _final_metric(df, "val_loss_score"), + "val_score_top1_acc": _final_metric(df, "val_score_top1_acc"), + "val_score_gap_best_second": _final_metric(df, "val_score_gap_best_second"), + } + rows.append(row) + + if not rows: + return pd.DataFrame() + out = pd.DataFrame(rows) + out = out.sort_values("run_name").reset_index(drop=True) + return out + + +def _approach_label(r: pd.Series) -> str: + t = r["score_loss_type"] + tgt = r.get("score_target_type", "l1") + prefix = f"{t}+{tgt}" if tgt != "l1" else t + if t == "bce_pairwise": + return f"{prefix}(w={r['score_rank_weight']:.2f},m={r['score_margin']:.2f},k={int(r['score_topk'])})" + if t in ("bce", "listnet"): + return f"{prefix}(tau={r['score_temperature']:.1f})" + return prefix + + +def _diagnose_row(r: pd.Series) -> str: + pred = r["val_ade_pred"] + oracle = r["val_ade_oracle"] + regret = r["val_ade_regret"] + top1 = r.get("val_score_top1_acc", np.nan) + + if np.isnan(pred) or np.isnan(oracle): + return "incomplete metrics" + if oracle < 0.8 and regret > 1.0: + return "strong proposals, weak ranking (scorer bottleneck)" + if oracle >= 0.8: + return "trajectory quality bottleneck (oracle not strong enough)" + if not np.isnan(top1) and top1 < 0.3: + return "low top1 match; scorer ordering not learning" + if regret < 0.7: + return "ranking improved" + return "mixed behavior" + + +def _plot_rank_bar(df: pd.DataFrame, out_dir: Path): + g = df.groupby("approach", as_index=False)["val_ade_pred"].min().sort_values("val_ade_pred") + if g.empty: + return + fig, ax = plt.subplots(figsize=(11, 5)) + x = np.arange(len(g)) + ax.bar(x, g["val_ade_pred"], color="#2f6db0") + ax.set_ylabel("Best val_ade_pred (m)") + ax.set_title("Approach Ranking (lower is better)") + ax.set_xticks(x) + ax.set_xticklabels(g["approach"], rotation=25, ha="right") + ax.grid(True, axis="y", alpha=0.3) + fig.tight_layout() + fig.savefig(out_dir / "ranking_val_ade_pred.png", dpi=180) + plt.close(fig) + + +def _plot_oracle_vs_pred(df: pd.DataFrame, out_dir: Path): + fig, ax = plt.subplots(figsize=(7, 6)) + for name, sub in df.groupby("approach"): + ax.scatter(sub["val_ade_oracle"], sub["val_ade_pred"], label=name, alpha=0.85) + lim_lo = np.nanmin([df["val_ade_oracle"].min(), df["val_ade_pred"].min()]) * 0.9 + lim_hi = np.nanmax([df["val_ade_oracle"].max(), df["val_ade_pred"].max()]) * 1.1 + ax.plot([lim_lo, lim_hi], [lim_lo, lim_hi], "--", color="gray", alpha=0.7, label="pred=oracle") + ax.set_xlabel("val_ade_oracle (m)") + ax.set_ylabel("val_ade_pred (m)") + ax.set_title("Oracle vs Selected ADE") + ax.grid(True, alpha=0.3) + ax.legend(fontsize=8) + fig.tight_layout() + fig.savefig(out_dir / "oracle_vs_pred_scatter.png", dpi=180) + plt.close(fig) + + +def _plot_regret_vs_top1(df: pd.DataFrame, out_dir: Path): + if "val_score_top1_acc" not in df.columns or df["val_score_top1_acc"].isna().all(): + return + fig, ax = plt.subplots(figsize=(7, 6)) + for name, sub in df.groupby("approach"): + ax.scatter(sub["val_score_top1_acc"], sub["val_ade_regret"], label=name, alpha=0.85) + ax.set_xlabel("val_score_top1_acc") + ax.set_ylabel("val_ade_regret (m)") + ax.set_title("Ranking Accuracy vs Regret") + ax.grid(True, alpha=0.3) + ax.legend(fontsize=8) + fig.tight_layout() + fig.savefig(out_dir / "regret_vs_top1_scatter.png", dpi=180) + plt.close(fig) + + +def _write_report(df: pd.DataFrame, out_dir: Path): + ranked = df.sort_values("val_ade_pred").reset_index(drop=True) + by_approach = ( + df.groupby("approach", as_index=False)[ + ["val_ade_pred", "val_ade_oracle", "val_ade_regret", "val_loss_score", "val_score_top1_acc"] + ] + .agg("mean") + .sort_values("val_ade_pred") + ) + + lines: List[str] = [] + lines.append("# Scorer Experiment Report") + lines.append("") + lines.append("## Overall ranking (by val_ade_pred)") + lines.append("") + for i, r in ranked.head(10).iterrows(): + lines.append( + f"{i+1}. `{r['run_name']}` | approach `{r['approach']}` | " + f"val_ade_pred={r['val_ade_pred']:.3f}, oracle={r['val_ade_oracle']:.3f}, " + f"regret={r['val_ade_regret']:.3f} | diagnosis: {r['diagnosis']}" + ) + + lines.append("") + lines.append("## Approach-level means") + lines.append("") + lines.append("| approach | val_ade_pred | val_ade_oracle | val_ade_regret | val_loss_score | val_score_top1_acc |") + lines.append("|---|---:|---:|---:|---:|---:|") + for _, r in by_approach.iterrows(): + lines.append( + f"| {r['approach']} | {r['val_ade_pred']:.3f} | {r['val_ade_oracle']:.3f} | " + f"{r['val_ade_regret']:.3f} | {r['val_loss_score']:.3f} | {r['val_score_top1_acc']:.3f} |" + ) + + lines.append("") + lines.append("## What went wrong / right") + lines.append("") + scorer_bad = ranked[(ranked["val_ade_oracle"] < 0.8) & (ranked["val_ade_regret"] > 1.0)] + if not scorer_bad.empty: + lines.append("- **Scorer bottleneck persists** in runs where oracle is strong but regret stays high.") + trajectory_bad = ranked[ranked["val_ade_oracle"] >= 0.8] + if not trajectory_bad.empty: + lines.append("- **Trajectory generation bottleneck** appears in some runs (high oracle ADE).") + improved = ranked[ranked["val_ade_regret"] < 0.7] + if not improved.empty: + lines.append("- **Ranking improved** for runs with regret below 0.7.") + if scorer_bad.empty and trajectory_bad.empty and improved.empty: + lines.append("- Mixed outcomes; inspect scatter plots for separation patterns.") + + (out_dir / "score_experiment_report.md").write_text("\n".join(lines)) + + +def main(): + parser = argparse.ArgumentParser(description="Rank scorer-loss experiments and generate diagnostics") + parser.add_argument("--logs_root", type=str, required=True, help="Root containing camera_e2e_*/version_0") + parser.add_argument("--out_dir", type=str, default=None, help="Output dir (default: /score_experiments)") + parser.add_argument("--latest_n", type=int, default=0, help="Only process latest N runs (0 = all)") + args = parser.parse_args() + + logs_root = Path(args.logs_root) + out_dir = Path(args.out_dir) if args.out_dir else logs_root / "score_experiments" + out_dir.mkdir(parents=True, exist_ok=True) + + df = collect_run_table(logs_root, latest_n=args.latest_n) + if df.empty: + raise RuntimeError(f"No runs with metrics found under {logs_root}") + + df["approach"] = df.apply(_approach_label, axis=1) + df["diagnosis"] = df.apply(_diagnose_row, axis=1) + df.to_csv(out_dir / "score_experiment_summary.csv", index=False) + + _plot_rank_bar(df, out_dir) + _plot_oracle_vs_pred(df, out_dir) + _plot_regret_vs_top1(df, out_dir) + _write_report(df, out_dir) + + print(f"Wrote summary CSV and plots to: {out_dir}") + print(f"Top run by val_ade_pred: {df.sort_values('val_ade_pred').iloc[0]['run_name']}") + + +if __name__ == "__main__": + main() + diff --git a/src/camera-based-e2e/train.py b/src/camera-based-e2e/train.py index c254836..29f2fd7 100644 --- a/src/camera-based-e2e/train.py +++ b/src/camera-based-e2e/train.py @@ -20,27 +20,27 @@ # Replace with your model defined in models/ from models.base_model import LitModel, collate_with_images -from models.monocular import DeepMonocularModel -from models.feature_extractors import SAMFeatures +from models.proposal_planner import ProposalPlanner, IPadConfig +from models.feature_extractors import SAMFeatures, DINOFeatures, ResNetFeatures class HomogeneousConcatBatchSampler(BatchSampler): """Emit batches from one ConcatDataset source at a time. Designed for ConcatDataset([waymo, nuscenes]) so every batch is - from one source. Works with DDP + from one source. Works with DDP. """ def __init__( self, - dataset_lengths: tuple[int, int], + dataset_lengths: tuple, batch_size: int, - rank: int | None = None, - world_size: int | None = None, + rank: int = None, + world_size: int = None, drop_last: bool = False, shuffle: bool = True, seed: int = 42, - source_ratio: tuple[int, int] = (1, 1), + source_ratio: tuple = (1, 1), **kwargs, ): if len(dataset_lengths) != 2: @@ -51,13 +51,10 @@ def __init__( rank = int(os.environ.get("RANK", "0")) if world_size is None: world_size = int(os.environ.get("WORLD_SIZE", "1")) - if world_size <= 0: raise ValueError(f"world_size must be > 0, got {world_size}") if rank < 0 or rank >= world_size: - raise ValueError( - f"Invalid rank/world_size pair: rank={rank}, world_size={world_size}" - ) + raise ValueError(f"Invalid rank/world_size pair: rank={rank}, world_size={world_size}") self.lengths = dataset_lengths self.batch_size = batch_size @@ -68,7 +65,6 @@ def __init__( self.seed = seed self.epoch = 0 self.source_ratio = source_ratio - self.offset0 = 0 self.offset1 = dataset_lengths[0] @@ -82,13 +78,12 @@ def _num_samples_per_rank(self, length: int) -> int: return math.ceil((length - self.world_size) / self.world_size) return math.ceil(length / self.world_size) - def _make_rank_indices(self, start: int, length: int, rng: random.Random): + def _make_rank_indices(self, start, length, rng): idx = list(range(start, start + length)) if self.shuffle: rng.shuffle(idx) num_samples = self._num_samples_per_rank(length) total_size = num_samples * self.world_size - if self.drop_last: idx = idx[:total_size] elif total_size > len(idx): @@ -97,17 +92,12 @@ def _make_rank_indices(self, start: int, length: int, rng: random.Random): padding_size = total_size - len(idx) repeats = math.ceil(padding_size / len(idx)) idx += (idx * repeats)[:padding_size] - return idx[self.rank : total_size : self.world_size] - def _chunk(self, indices: list[int]) -> list[list[int]]: + def _chunk(self, indices): if self.drop_last: n_full = len(indices) // self.batch_size - return [ - indices[i * self.batch_size : (i + 1) * self.batch_size] - for i in range(n_full) - ] - + return [indices[i * self.batch_size : (i + 1) * self.batch_size] for i in range(n_full)] out = [] for i in range(0, len(indices), self.batch_size): out.append(indices[i : i + self.batch_size]) @@ -115,72 +105,72 @@ def _chunk(self, indices: list[int]) -> list[list[int]]: def __iter__(self): rng = random.Random(self.seed + self.epoch) - idx0 = self._make_rank_indices(self.offset0, self.lengths[0], rng) idx1 = self._make_rank_indices(self.offset1, self.lengths[1], rng) - - b0 = self._chunk(idx0) - b1 = self._chunk(idx1) - + b0, b1 = self._chunk(idx0), self._chunk(idx1) w0 = max(0, int(self.source_ratio[0])) w1 = max(0, int(self.source_ratio[1])) if w0 == 0 and w1 == 0: w0, w1 = 1, 1 - - pattern = [0] * w0 + [1] * w1 - if not pattern: - pattern = [0, 1] - + pattern = [0] * w0 + [1] * w1 or [0, 1] i0, i1 = 0, 0 for source in itertools.cycle(pattern): if i0 >= len(b0) and i1 >= len(b1): break if source == 0: if i0 < len(b0): - yield b0[i0] - i0 += 1 + yield b0[i0]; i0 += 1 else: if i1 < len(b1): - yield b1[i1] - i1 += 1 + yield b1[i1]; i1 += 1 def __len__(self): - def n_batches(length: int): + def n_batches(length): per_rank = self._num_samples_per_rank(length) - if self.drop_last: - return per_rank // self.batch_size - return math.ceil(per_rank / self.batch_size) - + return per_rank // self.batch_size if self.drop_last else math.ceil(per_rank / self.batch_size) return n_batches(self.lengths[0]) + n_batches(self.lengths[1]) if __name__ == "__main__": parser = argparse.ArgumentParser() - parser.add_argument( - "--data_dir", type=str, required=True, help="Path to data directory" - ) - parser.add_argument( - "--batch_size", type=int, default=16, help="Batch size for training" - ) - parser.add_argument("--lr", type=float, default=1e-4, help="Learning rate") - parser.add_argument( - "--max_epochs", type=int, default=10, help="Number of epochs to train" - ) - parser.add_argument( - "--compile", - action="store_true", - help="Whether to compile the model with torch.compile", - ) - parser.add_argument( - "--profile", action="store_true", help="Whether to run the profiler" - ) - parser.add_argument( - "--dataset", - type=str, - default="waymo", - choices=["waymo", "nuscenes", "all"], - help="Which dataset to train on", - ) + parser.add_argument('--data_dir', type=str, required=True, help='Path to data directory') + parser.add_argument('--dataset', type=str, default='waymo', choices=['waymo', 'nuscenes', 'all'], + help='Which dataset to train on') + parser.add_argument('--model_type', type=str, default='deep_monocular', choices=['deep_monocular', 'proposal'], + help='Model type: deep_monocular or proposal (proposal-centric planner)') + parser.add_argument('--batch_size', type=int, default=16, help='Batch size for training') + parser.add_argument('--lr', type=float, default=1e-4, help='Learning rate') + parser.add_argument('--max_epochs', type=int, default=10, help='Number of epochs to train') + parser.add_argument('--num_proposals', type=int, default=16, help='Number of proposals (proposal model only)') + parser.add_argument('--num_refinement_steps', type=int, default=4, help='Refinement iterations (weight-shared, iPad default=4)') + parser.add_argument('--smoothness_weight', type=float, default=0.01, help='Smoothness (jerk) loss weight (proposal model)') + parser.add_argument('--collision_weight', type=float, default=0.0, help='Collision penalty weight (proposal model)') + parser.add_argument('--comfort_weight', type=float, default=0.01, help='Comfort (curvature) loss weight (proposal model)') + parser.add_argument('--rfs_weight', type=float, default=0.0, help='RFS loss weight') + parser.add_argument('--diversity_weight', type=float, default=0.0, help='Diversity weight') + parser.add_argument('--prev_weight', type=float, default=0.1, help='Discount λ for intermediate proposal losses') + parser.add_argument('--score_weight', type=float, default=1.0, help='Score loss weight') + parser.add_argument('--score_warmup_epochs', type=int, default=2, help='Epochs before score loss activates') + parser.add_argument('--score_temperature', type=float, default=5.0, help='Temperature τ for quality target exp(-ADE/τ)') + parser.add_argument('--score_loss_type', type=str, default='bce', + choices=['bce', 'ce', 'bce_pairwise', 'listnet'], + help='Scorer objective: bce (iPad-faithful), ce, bce_pairwise, or listnet') + parser.add_argument('--score_target_type', type=str, default='l1', + choices=['l1', 'navsim', 'rfs'], + help='Quality target: l1, navsim, or rfs') + parser.add_argument('--no_rfs_target_comfort', action='store_true', + help='For score_target_type=rfs: use pure RFS mean only (no jerk comfort multiplier)') + parser.add_argument('--score_rank_weight', type=float, default=0.2, help='Aux weight for pairwise ranking term') + parser.add_argument('--score_margin', type=float, default=0.2, help='Pairwise ranking margin') + parser.add_argument('--score_topk', type=int, default=0, help='Hard negative top-k for pairwise ranking (0=all)') + parser.add_argument('--comfort_jerk_threshold', type=float, default=5.0, help='Jerk threshold (m/s^3) for comfort metric') + parser.add_argument('--grad_clip', type=float, default=1.0, help='Gradient clipping max norm (0 to disable)') + parser.add_argument('--log_every_n_steps', type=int, default=100, help='How often (steps) to emit trainer logs') + parser.add_argument('--backbone', type=str, default='resnet', choices=['resnet', 'dino', 'sam'], + help='Backbone: resnet (default), dino, or sam') + parser.add_argument('--no_wandb', action='store_true', help='Disable wandb logging (use CSV only)') + parser.add_argument('--compile', action='store_true', help='Whether to compile the model with torch.compile') + parser.add_argument('--profile', action='store_true', help='Whether to run the profiler') args = parser.parse_args() pl.seed_everything(42, workers=True) @@ -295,41 +285,84 @@ def n_batches(length: int): ) # Model - in_dim = 16 * 6 # Past: (B, 16, 6) out_dim = 20 * 2 # Future: (B, 20, 2) - model = DeepMonocularModel( - feature_extractor=SAMFeatures( - model_name="timm/vit_pe_spatial_small_patch16_512.fb", frozen=True - ), - out_dim=out_dim, - ) + if args.model_type == 'proposal': + backbone = ( + ResNetFeatures(frozen=True) if args.backbone == 'resnet' + else DINOFeatures(frozen=True) if args.backbone == 'dino' + else SAMFeatures(frozen=True) + ) + model = ProposalPlanner( + backbone=backbone, + d_model=256, + num_proposals=args.num_proposals, + num_refinement_steps=args.num_refinement_steps, + horizon=20, + num_cams=8, + ) + else: + backbone = ( + ResNetFeatures(frozen=True) if args.backbone == 'resnet' + else DINOFeatures(frozen=True) if args.backbone == 'dino' + else SAMFeatures(frozen=True) + ) + model = DeepMonocularModel(feature_extractor=backbone, out_dim=out_dim, n_blocks=8) name = str(model.__class__.__name__.replace("Model", "")).lower() + if args.compile: model = torch.compile(model, mode="max-autotune") - lit_model = LitModel(model=model, lr=args.lr) + + if args.model_type == 'proposal': + ipad_config = IPadConfig( + rfs_weight=args.rfs_weight, + smoothness_weight=args.smoothness_weight, + collision_weight=args.collision_weight, + comfort_weight=args.comfort_weight, + diversity_weight=args.diversity_weight, + score_weight=args.score_weight, + score_warmup_epochs=args.score_warmup_epochs, + score_temperature=args.score_temperature, + score_loss_type=args.score_loss_type, + score_target_type=args.score_target_type, + score_rank_weight=args.score_rank_weight, + score_margin=args.score_margin, + score_topk=args.score_topk, + comfort_jerk_threshold=args.comfort_jerk_threshold, + prev_weight=args.prev_weight, + rfs_target_use_comfort=not args.no_rfs_target_comfort, + ) + else: + ipad_config = None + + lit_model = LitModel( + model=model, + lr=args.lr, + ipad_config=ipad_config, + ) # We don't want to save logs or checkpoints in the home directory - it'll fill up fast base_path = Path(args.data_dir).parent.as_posix() timestamp = f"{name}_e2e_{args.dataset}_{datetime.now().strftime('%Y%m%d_%H%M')}" - wandb_logger = WandbLogger( - name=timestamp, - save_dir=base_path + "/logs", - project="robotvision", - log_model=True, - ) - wandb_logger.watch(lit_model, log="all") + loggers = [CSVLogger(base_path + "/logs", name=timestamp)] + if not args.no_wandb: + wandb_logger = WandbLogger(name=timestamp, save_dir=base_path + "/logs", project="robotvision", log_model=True) + wandb_logger.watch(lit_model, log="all") + loggers.append(wandb_logger) strategy = "ddp" if torch.cuda.device_count() > 1 else "auto" use_distributed_sampler = args.dataset != "all" torch.set_float32_matmul_precision("medium") + trainer = pl.Trainer( max_epochs=args.max_epochs, - logger=[CSVLogger(base_path + "/logs", name=timestamp), wandb_logger], + logger=loggers, strategy=strategy, use_distributed_sampler=use_distributed_sampler, precision="bf16-mixed" if torch.cuda.is_bf16_supported() else 16, - log_every_n_steps=10, + log_every_n_steps=args.log_every_n_steps, + gradient_clip_val=args.grad_clip if args.grad_clip > 0 else None, + gradient_clip_algorithm="norm", profiler=SimpleProfiler(extended=True) if args.profile else None, callbacks=[ ModelCheckpoint( @@ -360,6 +393,8 @@ def n_batches(length: int): plt.legend() plt.tight_layout() out = Path("./visualizations") + out.mkdir(parents=True, exist_ok=True) plt.savefig(out / "loss.png", dpi=200) except Exception as e: print(f"Could not save loss plot: {e}") +