From 5a69e65d1e72455f56d9d1eae9f4099d040138b3 Mon Sep 17 00:00:00 2001 From: Manav Gagvani Date: Wed, 29 Jul 2026 02:24:57 -0400 Subject: [PATCH] sae: add scorer correction adapter, ported to the canonical SAE stack Trains a lightweight residual adapter over the frozen planner's proposal scorer, predicting the scorer's own error from SAE latents and correcting scores additively at inference. The base planner and proposal set are untouched - only the ranking changes. Ported off the legacy sparseAE L1 stack onto models/sae.py + sae_utils.py. sparseAE.py, its colliding 88-line sae_utils.py, and test_sparseAE.py are deliberately not merged: main now has exactly one SAE implementation and one checkpoint format. Port surface was narrow - SparseAutoencoder already exposes .encoder, so infer_residual_dims needed only a dim source, and Ben had already ported compute_hidden_activations and load_model_and_sae in new_sae_utils. The substantive change is activation capture: sae.internal_acts (state hung off the SAE) becomes an explicit ActivationCapture threaded through extract_batch_targets. Adds --latent_source {sae,raw}. The paper's control - the same adapter fed the dense activation the SAE reconstructs instead of its sparse latents - was previously run by hand and left no reproducible path. It is now a flag, persisted into the residual checkpoint and read back by the evaluator, so the control arm can be re-run exactly. sae_utils gains three things this needs: * ActivationCapture.register/remove - the hook handle was previously discarded, so hooks could not be detached. * freeze_module / set_eval_mode, carried over from the legacy utils. * normalize_compiled_state_dict, salvaged from sparseAE.py. A checkpoint saved from a compiled SAE stores encoder._orig_mod.weight; loading it into an eager module under strict=False silently dropped every tensor and yielded a randomly-initialised SAE that still appeared to load. Verified on CPU against the real block-3 checkpoint: hook fires, captured activation is (B,1,384), latents are (B,384), and the handle detaches. --- src/camera-based-e2e/sae_scorer_correction.py | 860 ++++++++++++++++++ src/camera-based-e2e/sae_utils.py | 57 +- .../test_sae_scorer_residual.py | 209 +++++ 3 files changed, 1122 insertions(+), 4 deletions(-) create mode 100644 src/camera-based-e2e/sae_scorer_correction.py create mode 100644 src/camera-based-e2e/test_sae_scorer_residual.py diff --git a/src/camera-based-e2e/sae_scorer_correction.py b/src/camera-based-e2e/sae_scorer_correction.py new file mode 100644 index 0000000..00bc043 --- /dev/null +++ b/src/camera-based-e2e/sae_scorer_correction.py @@ -0,0 +1,860 @@ +""" +Utilities for training and evaluating an SAE-informed scorer correction model. +""" + +from __future__ import annotations + +import argparse +import os +from datetime import datetime +from pathlib import Path +from typing import Any + +import pytorch_lightning as pl +import torch +import torch.nn as nn +import torch.nn.functional as F +from pytorch_lightning.callbacks import EarlyStopping, ModelCheckpoint +from pytorch_lightning.loggers import CSVLogger, WandbLogger +from torch.utils.data import DataLoader + +from loader import WaymoE2E +from models.base_model import LitModel, collate_with_images +from sae_utils import ( + ActivationCapture, + compute_hidden_activations, + freeze_module, + load_model_and_sae, + set_eval_mode, +) +from models.sae import SparseAutoencoder + +DEFAULT_TRAIN_SPLIT = "train" +DEFAULT_VAL_SPLIT = "val" +DEFAULT_PATIENCE, DEFAULT_MIN_DELTA = 2, 1e-3 +DEFAULT_LOG_EVERY_N_STEPS, DEFAULT_SEED = 10, 42 +VAL_MONITOR = "val_corrected_picked_ade" +INDEX_FILES = { + DEFAULT_TRAIN_SPLIT: "index_train.pkl", + DEFAULT_VAL_SPLIT: "index_val.pkl", + "test": "index_test.pkl", +} +SCORE_METRIC_NAMES = { + "baseline_score_mse", + "corrected_score_mse", + "baseline_score_mae", + "corrected_score_mae", +} + + +def default_num_workers() -> int: + candidates: list[int] = [] + for env_name in ("SLURM_CPUS_PER_TASK", "SLURM_CPUS_ON_NODE"): + env_value = os.environ.get(env_name) + if env_value is not None: + try: + candidates.append(int(env_value)) + except ValueError: + pass + if hasattr(os, "sched_getaffinity"): + candidates.append(len(os.sched_getaffinity(0))) + candidates.append(os.cpu_count() or 1) + return max(1, min(candidates)) + + +def default_index_file(split: str) -> str: + return INDEX_FILES.get(split, INDEX_FILES[DEFAULT_TRAIN_SPLIT]) + + +def checkpoint_value(checkpoint: dict[str, Any], key: str) -> Any: + if key in checkpoint: + return checkpoint[key] + args = checkpoint.get("args") + if isinstance(args, dict): + return args.get(key) + return None + + +def resolve_required_arg( + cli_value: Any, + checkpoint: dict[str, Any], + key: str, +) -> Any: + if cli_value is not None: + return cli_value + checkpoint_val = checkpoint_value(checkpoint, key) + if checkpoint_val is None: + raise ValueError( + f"Missing required argument '{key}'. Pass --{key} or use a residual checkpoint that stores it." + ) + return checkpoint_val + + +def build_train_val_loaders( + *, + data_dir: str, + train_items: int | None, + val_items: int | None, + batch_size: int, + num_workers: int, + seed: int, +) -> tuple[DataLoader, DataLoader]: + train_dataset = WaymoE2E( + indexFile=default_index_file(DEFAULT_TRAIN_SPLIT), + data_dir=data_dir, + n_items=train_items, + seed=seed, + ) + val_dataset = WaymoE2E( + indexFile=default_index_file(DEFAULT_VAL_SPLIT), + data_dir=data_dir, + n_items=val_items, + seed=seed + 1, + ) + return ( + DataLoader( + train_dataset, + batch_size=batch_size, + num_workers=num_workers, + collate_fn=collate_with_images, + persistent_workers=False, + pin_memory=False, + shuffle=True, + ), + DataLoader( + val_dataset, + batch_size=batch_size, + num_workers=num_workers, + collate_fn=collate_with_images, + persistent_workers=False, + pin_memory=False, + shuffle=False, + ), + ) + + +LATENT_SOURCES = ("sae", "raw") + + +def latent_features( + sae: SparseAutoencoder, + activations: torch.Tensor, + latent_source: str, +) -> torch.Tensor: + """Features the residual adapter consumes. + + "sae" encodes the captured activation into the sparse latent space; "raw" + is the paper's control arm, feeding the same adapter the dense activation + the SAE was trained to reconstruct. + """ + if latent_source == "sae": + return compute_hidden_activations(sae, activations) + if latent_source == "raw": + return activations.reshape(activations.size(0), -1) + raise ValueError(f"latent_source must be one of {LATENT_SOURCES}, got {latent_source!r}") + + +def extract_batch_targets( + model: LitModel, + sae: SparseAutoencoder, + capture: ActivationCapture, + batch: dict, + device: torch.device, + latent_source: str = "sae", +) -> dict[str, torch.Tensor]: + past = batch["PAST"].to(device) + future = batch["FUTURE"].to(device) + intent = batch["INTENT"].to(device) + + if "IMAGES" in batch: + images = [ + image.to(device) if isinstance(image, torch.Tensor) else image + for image in batch["IMAGES"] + ] + else: + images = model.decode_batch_jpeg(batch["IMAGES_JPEG"], device=device) + + model_inputs = {"PAST": past, "IMAGES": images, "INTENT": intent} + capture.clear() + output = model(model_inputs) + + if not isinstance(output, dict): + raise TypeError("Expected the backbone model to return a dict with trajectory and scores") + if output.get("scores") is None: + raise RuntimeError("Backbone model does not expose scorer predictions in output['scores']") + if capture.activations is None: + raise RuntimeError("No SAE activations were captured from the registered hook") + + pred_scores = output["scores"] + pred_future = output["trajectory"] + t_steps = future.size(1) + pred = pred_future.view(pred_future.size(0), -1, t_steps, 2) + ade_per_mode = torch.norm(pred - future[:, None, :, :], dim=-1).mean(dim=-1) + + sae_hidden = latent_features(sae, capture.activations, latent_source) + trajectories = pred.reshape(pred.size(0), pred.size(1), -1) + + return { + "sae_hidden": sae_hidden.detach(), + "trajectories": trajectories.detach(), + "score_pred": pred_scores.detach(), + "true_ade": ade_per_mode.detach(), + "residual_target": (ade_per_mode - pred_scores).detach(), + } + + +def flatten_proposal_batch(batch_targets: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]: + sae_hidden = batch_targets["sae_hidden"] + trajectories = batch_targets["trajectories"] + score_pred = batch_targets["score_pred"] + residual_target = batch_targets["residual_target"] + true_ade = batch_targets["true_ade"] + + batch_size, n_modes, traj_dim = trajectories.shape + sae_dim = sae_hidden.size(-1) + + return { + "sae_hidden": sae_hidden[:, None, :].expand(-1, n_modes, -1).reshape(batch_size * n_modes, sae_dim), + "trajectory": trajectories.reshape(batch_size * n_modes, traj_dim), + "score_pred": score_pred.reshape(batch_size * n_modes), + "residual_target": residual_target.reshape(batch_size * n_modes), + "true_ade": true_ade.reshape(batch_size * n_modes), + } + + +def score_comparison_metrics( + baseline_scores: torch.Tensor, + corrected_scores: torch.Tensor, + true_ade: torch.Tensor, +) -> dict[str, torch.Tensor]: + metrics = { + "baseline_score_mse": torch.mean((baseline_scores - true_ade) ** 2), + "corrected_score_mse": torch.mean((corrected_scores - true_ade) ** 2), + "baseline_score_mae": torch.mean(torch.abs(baseline_scores - true_ade)), + "corrected_score_mae": torch.mean(torch.abs(corrected_scores - true_ade)), + } + for prefix, scores in (("baseline", baseline_scores), ("corrected", corrected_scores)): + for key, value in selection_metrics(scores, true_ade).items(): + metrics[f"{prefix}_{key}"] = value + return metrics + + +def selection_metrics(score_matrix: torch.Tensor, ade_per_mode: torch.Tensor) -> dict[str, torch.Tensor]: + batch_size = score_matrix.size(0) + batch_idx = torch.arange(batch_size, device=score_matrix.device) + picked_idx = score_matrix.argmin(dim=1) + picked_ade = ade_per_mode[batch_idx, picked_idx] + oracle_ade = ade_per_mode.min(dim=1).values + + oracle_ranking = ade_per_mode.argsort(dim=1) + oracle_rank_of = oracle_ranking.argsort(dim=1) + picked_rank = oracle_rank_of[batch_idx, picked_idx].float() + + n_modes = ade_per_mode.size(1) + scorer_ranks = score_matrix.argsort(dim=1).argsort(dim=1).float() + oracle_ranks = oracle_rank_of.float() + d = oracle_ranks - scorer_ranks + denom = max(n_modes * (n_modes * n_modes - 1), 1) + spearman = 1 - 6 * (d.square().sum(dim=1) / denom) + + metrics = { + "picked_ade": picked_ade.mean(), + "oracle_ade": oracle_ade.mean(), + "regret": (picked_ade - oracle_ade).mean(), + "mean_rank": picked_rank.mean(), + "spearman": spearman.mean(), + } + for topk in (1, 5, 10): + if topk <= n_modes: + metrics[f"top{topk}_acc"] = (picked_rank < topk).float().mean() + return metrics + + +class SAEScorerResidual(nn.Module): + def __init__( + self, + sae_dim: int, + traj_dim: int, + hidden_dim: int = 512, + dropout: float = 0.1, + ) -> None: + super().__init__() + self.sae_tower = nn.Sequential( + nn.LayerNorm(sae_dim), + nn.Linear(sae_dim, hidden_dim), + nn.GELU(), + nn.Dropout(dropout), + ) + self.traj_tower = nn.Sequential( + nn.LayerNorm(traj_dim), + nn.Linear(traj_dim, hidden_dim), + nn.GELU(), + nn.Dropout(dropout), + ) + self.score_tower = nn.Sequential( + nn.Linear(1, hidden_dim // 2), + nn.GELU(), + nn.Dropout(dropout), + ) + fusion_dim = hidden_dim * 2 + hidden_dim // 2 + self.head = nn.Sequential( + nn.LayerNorm(fusion_dim), + nn.Linear(fusion_dim, hidden_dim), + nn.GELU(), + nn.Dropout(dropout), + nn.Linear(hidden_dim, hidden_dim // 2), + nn.GELU(), + nn.Dropout(dropout), + nn.Linear(hidden_dim // 2, 1), + ) + + def forward( + self, + sae_hidden: torch.Tensor, + trajectory: torch.Tensor, + score_pred: torch.Tensor, + ) -> torch.Tensor: + if score_pred.ndim == 1: + score_pred = score_pred.unsqueeze(-1) + features = torch.cat( + [ + self.sae_tower(sae_hidden), + self.traj_tower(trajectory), + self.score_tower(score_pred), + ], + dim=-1, + ) + return self.head(features).squeeze(-1) + + +def infer_residual_dims( + model: LitModel, sae: SparseAutoencoder, latent_source: str = "sae" +) -> tuple[int, int]: + feature_dim = sae.encoder.weight.shape[0] if latent_source == "sae" else sae.input_dim + return feature_dim, model.model.horizon * 2 + + +class LitSAEScorerResidual(pl.LightningModule): + def __init__( + self, + *, + model_checkpoint_path: str, + sae_checkpoint_path: str, + block_idx: int = 3, + lr: float = 1e-3, + weight_decay: float = 1e-3, + hidden_dim: int = 256, + dropout: float = 0.2, + loss_type: str = "huber", + topk_train: int = 10, + topk_weight: float = 3.0, + rest_weight: float = 1.0, + rank_loss_weight: float = 0.25, + rank_temperature: float = 1.0, + latent_source: str = "sae", + ) -> None: + super().__init__() + if latent_source not in LATENT_SOURCES: + raise ValueError(f"latent_source must be one of {LATENT_SOURCES}, got {latent_source!r}") + self.latent_source = latent_source + + backbone_model, sae, capture, _ = load_model_and_sae( + model_checkpoint_path=model_checkpoint_path, + sae_checkpoint_path=sae_checkpoint_path, + block_idx=block_idx, + ) + freeze_module(backbone_model) + freeze_module(sae) + set_eval_mode(backbone_model, sae) + + self.backbone_model = backbone_model + self.sae = sae + self.capture = capture + self.sae_dim, self.traj_dim = infer_residual_dims(backbone_model, sae, latent_source) + self.residual_model = SAEScorerResidual( + sae_dim=self.sae_dim, + traj_dim=self.traj_dim, + hidden_dim=hidden_dim, + dropout=dropout, + ) + self.save_hyperparameters({ + key: value + for key, value in locals().items() + if key not in {"self", "backbone_model", "sae", "capture"} + } | {"sae_dim": self.sae_dim, "traj_dim": self.traj_dim}) + + def _set_backbone_eval(self) -> None: + set_eval_mode(self.backbone_model, self.sae) + + def remove_hook(self) -> None: + if getattr(self, "capture", None) is not None: + self.capture.remove() + self.capture = None + + def transfer_batch_to_device(self, batch, device, dataloader_idx): + if not isinstance(batch, dict): + return super().transfer_batch_to_device(batch, device, dataloader_idx) + + if "IMAGES_JPEG" in batch: + images_jpeg = batch["IMAGES_JPEG"] + batch_wo_jpeg = dict(batch) + batch_wo_jpeg.pop("IMAGES_JPEG", None) + moved = super().transfer_batch_to_device(batch_wo_jpeg, device, dataloader_idx) + moved["IMAGES"] = self.backbone_model.decode_batch_jpeg(images_jpeg, device=device) + return moved + + return super().transfer_batch_to_device(batch, device, dataloader_idx) + + def on_fit_start(self) -> None: + self._set_backbone_eval() + + def on_train_epoch_start(self) -> None: + self._set_backbone_eval() + + def on_validation_epoch_start(self) -> None: + self._set_backbone_eval() + + def on_fit_end(self) -> None: + self.remove_hook() + + def teardown(self, stage: str | None) -> None: + self.remove_hook() + return super().teardown(stage) + + def configure_optimizers(self): + return torch.optim.AdamW( + self.residual_model.parameters(), + lr=self.hparams.lr, + weight_decay=self.hparams.weight_decay, + ) + + def _log_step_metric( + self, + name: str, + value: torch.Tensor, + *, + batch_size: int, + prog_bar: bool = False, + ) -> None: + self.log( + name, + value, + on_step=False, + on_epoch=True, + prog_bar=prog_bar, + batch_size=batch_size, + sync_dist=False, + ) + + def _shared_step(self, batch: dict[str, Any], stage: str) -> torch.Tensor: + self._set_backbone_eval() + + with torch.no_grad(): + batch_targets = extract_batch_targets( + self.backbone_model, + self.sae, + self.capture, + batch, + self.device, + self.latent_source, + ) + + flat_batch = flatten_proposal_batch(batch_targets) + proposal_weights = proposal_weights_from_ade( + batch_targets["true_ade"], + topk=self.hparams.topk_train, + topk_weight=self.hparams.topk_weight, + rest_weight=self.hparams.rest_weight, + ) + flat_weights = proposal_weights.reshape(-1) + + residual_pred_flat = self.residual_model( + flat_batch["sae_hidden"], + flat_batch["trajectory"], + flat_batch["score_pred"], + ) + residual_pred = residual_pred_flat.view_as(batch_targets["score_pred"]) + reg_loss = correction_loss( + residual_pred_flat, + flat_batch["residual_target"], + loss_type=self.hparams.loss_type, + weights=flat_weights, + ) + corrected_scores = batch_targets["score_pred"] + residual_pred + rank_loss = pairwise_ranking_loss( + corrected_scores, + batch_targets["true_ade"], + proposal_weights=proposal_weights, + temperature=self.hparams.rank_temperature, + ) + loss = reg_loss + self.hparams.rank_loss_weight * rank_loss + + sample_count = batch_targets["score_pred"].size(0) + score_count = batch_targets["score_pred"].numel() + comparison_metrics = score_comparison_metrics( + batch_targets["score_pred"], + corrected_scores, + batch_targets["true_ade"], + ) + + self._log_step_metric(f"{stage}_loss", loss, batch_size=score_count, prog_bar=(stage == "val")) + self._log_step_metric(f"{stage}_reg_loss", reg_loss, batch_size=score_count) + self._log_step_metric(f"{stage}_rank_loss", rank_loss, batch_size=score_count) + for key, value in comparison_metrics.items(): + self._log_step_metric( + f"{stage}_{key}", + value, + batch_size=score_count if key in SCORE_METRIC_NAMES else sample_count, + prog_bar=(stage == "val" and key == "corrected_picked_ade"), + ) + + return loss + + def training_step(self, batch: dict[str, Any], batch_idx: int) -> torch.Tensor: + return self._shared_step(batch, "train") + + def validation_step(self, batch: dict[str, Any], batch_idx: int) -> torch.Tensor: + return self._shared_step(batch, "val") + + +def _load_residual_state_from_lightning_checkpoint( + checkpoint: dict[str, Any], +) -> tuple[dict[str, torch.Tensor], dict[str, Any]]: + hparams = dict(checkpoint.get("hyper_parameters", {})) + state_dict = checkpoint.get("state_dict", {}) + residual_prefix = "residual_model." + residual_state = { + key[len(residual_prefix):]: value + for key, value in state_dict.items() + if key.startswith(residual_prefix) + } + if not residual_state: + raise KeyError("Lightning checkpoint does not contain residual_model.* weights") + return residual_state, hparams + + +def build_residual_model_from_metadata( + metadata: dict[str, Any], + *, + device: torch.device, +) -> SAEScorerResidual: + model = SAEScorerResidual( + sae_dim=metadata["sae_dim"], + traj_dim=metadata["traj_dim"], + hidden_dim=metadata["hidden_dim"], + dropout=metadata["dropout"], + ) + model = model.to(device) + model.eval() + return model + + +def load_residual_model( + checkpoint_path: str, + device: torch.device, +) -> tuple[SAEScorerResidual, dict[str, Any]]: + checkpoint = torch.load(checkpoint_path, map_location="cpu", weights_only=False) + + if "model_state_dict" in checkpoint: + model = build_residual_model_from_metadata(checkpoint, device=device) + model.load_state_dict(checkpoint["model_state_dict"], strict=True) + return model, checkpoint + + residual_state, metadata = _load_residual_state_from_lightning_checkpoint(checkpoint) + model = build_residual_model_from_metadata(metadata, device=device) + model.load_state_dict(residual_state, strict=True) + metadata["args"] = metadata.copy() + return model, metadata + + +def save_residual_model( + output_path: str, + model: SAEScorerResidual, + *, + sae_dim: int, + traj_dim: int, + hidden_dim: int, + dropout: float, + best_metric: float, + args: dict[str, Any], +) -> None: + path = Path(output_path) + path.parent.mkdir(parents=True, exist_ok=True) + torch.save( + { + "model_state_dict": model.state_dict(), + "sae_dim": sae_dim, + "traj_dim": traj_dim, + "hidden_dim": hidden_dim, + "dropout": dropout, + "best_metric": best_metric, + "args": args, + }, + path, + ) + + +def correction_loss( + residual_pred: torch.Tensor, + residual_target: torch.Tensor, + loss_type: str, + weights: torch.Tensor | None = None, +) -> torch.Tensor: + if loss_type == "huber": + loss = F.smooth_l1_loss(residual_pred, residual_target, reduction="none") + elif loss_type == "mse": + loss = F.mse_loss(residual_pred, residual_target, reduction="none") + else: + raise ValueError(f"Unsupported loss type: {loss_type}") + + if weights is not None: + weights = weights.to(loss.device, dtype=loss.dtype) + loss = loss * weights + return loss.sum() / weights.sum().clamp_min(1e-8) + + return loss.mean() + + +def proposal_weights_from_ade( + ade_per_mode: torch.Tensor, + *, + topk: int, + topk_weight: float, + rest_weight: float, +) -> torch.Tensor: + if ade_per_mode.ndim != 2: + raise ValueError(f"Expected ade_per_mode to be 2D, got shape {tuple(ade_per_mode.shape)}") + + n_modes = ade_per_mode.size(1) + topk = max(1, min(topk, n_modes)) + oracle_ranks = ade_per_mode.argsort(dim=1).argsort(dim=1) + weights = torch.full_like(ade_per_mode, float(rest_weight)) + weights = torch.where( + oracle_ranks < topk, + torch.full_like(weights, float(topk_weight)), + weights, + ) + return weights / weights.mean().clamp_min(1e-8) + + +def pairwise_ranking_loss( + score_matrix: torch.Tensor, + ade_per_mode: torch.Tensor, + *, + proposal_weights: torch.Tensor | None = None, + temperature: float = 1.0, +) -> torch.Tensor: + if score_matrix.ndim != 2 or ade_per_mode.ndim != 2: + raise ValueError("pairwise_ranking_loss expects 2D score and ADE matrices") + if score_matrix.shape != ade_per_mode.shape: + raise ValueError( + f"Score/ADE shape mismatch: {tuple(score_matrix.shape)} vs {tuple(ade_per_mode.shape)}" + ) + + score_diff = score_matrix[:, :, None] - score_matrix[:, None, :] + ade_diff = ade_per_mode[:, :, None] - ade_per_mode[:, None, :] + order_sign = torch.sign(-ade_diff) + + pair_mask = torch.triu( + torch.ones_like(order_sign, dtype=torch.bool), + diagonal=1, + ) & (order_sign != 0) + pair_loss = F.softplus((order_sign * score_diff) / max(temperature, 1e-6)) + + if proposal_weights is not None: + pair_weights = 0.5 * ( + proposal_weights[:, :, None] + proposal_weights[:, None, :] + ) + pair_loss = pair_loss * pair_weights + + pair_loss = pair_loss.masked_select(pair_mask) + if pair_loss.numel() == 0: + return score_matrix.new_tensor(0.0) + return pair_loss.mean() + + +def build_arg_parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser() + parser.add_argument("--model_checkpoint_path", type=str, required=True) + parser.add_argument("--sae_checkpoint_path", type=str, required=True) + parser.add_argument("--data_dir", type=str, required=True) + parser.add_argument("--output_path", type=str, required=True) + parser.add_argument("--train_items", type=int, default=250_000) + parser.add_argument("--val_items", type=int, default=25_000) + parser.add_argument("--batch_size", type=int, default=16) + parser.add_argument("--num_workers", type=int, default=14) + parser.add_argument("--block_idx", type=int, default=3) + parser.add_argument( + "--latent_source", + type=str, + default="sae", + choices=LATENT_SOURCES, + help="'sae' uses sparse latents; 'raw' is the control arm using the dense " + "block activation the SAE reconstructs.", + ) + parser.add_argument("--epochs", type=int, default=8) + parser.add_argument("--lr", type=float, default=1e-3) + parser.add_argument("--weight_decay", type=float, default=1e-3) + parser.add_argument("--hidden_dim", type=int, default=256) + parser.add_argument("--dropout", type=float, default=0.2) + parser.add_argument("--loss_type", type=str, default="huber", choices=("mse", "huber")) + parser.add_argument("--topk_train", type=int, default=10) + parser.add_argument("--topk_weight", type=float, default=3.0) + parser.add_argument("--rest_weight", type=float, default=1.0) + parser.add_argument("--rank_loss_weight", type=float, default=0.25) + parser.add_argument("--rank_temperature", type=float, default=1.0) + parser.add_argument("--log_dir", type=str, default=None) + parser.add_argument("--checkpoint_dir", type=str, default=None) + return parser + + +def make_run_name() -> str: + return f"sae_scorer_residual_{datetime.now().strftime('%Y%m%d_%H%M%S')}" + + +def export_args(args: argparse.Namespace) -> dict[str, Any]: + return { + **vars(args), + "train_split": DEFAULT_TRAIN_SPLIT, + "val_split": DEFAULT_VAL_SPLIT, + "patience": DEFAULT_PATIENCE, + "min_delta": DEFAULT_MIN_DELTA, + "run_name": None, + "log_every_n_steps": DEFAULT_LOG_EVERY_N_STEPS, + "wandb_project": "robotvision", + "seed": DEFAULT_SEED, + } + + +def export_best_residual_checkpoint( + lightning_checkpoint_path: str, + output_path: str, + args: dict[str, Any], + best_metric: float, +) -> None: + checkpoint = torch.load(lightning_checkpoint_path, map_location="cpu", weights_only=False) + residual_state, metadata = _load_residual_state_from_lightning_checkpoint(checkpoint) + residual_model = build_residual_model_from_metadata(metadata, device=torch.device("cpu")) + residual_model.load_state_dict(residual_state, strict=True) + save_residual_model( + output_path, + residual_model, + sae_dim=metadata["sae_dim"], + traj_dim=metadata["traj_dim"], + hidden_dim=metadata["hidden_dim"], + dropout=metadata["dropout"], + best_metric=best_metric, + args=args, + ) + + +def build_loggers(log_dir: Path, run_name: str) -> list[CSVLogger | WandbLogger]: + loggers: list[CSVLogger | WandbLogger] = [CSVLogger(log_dir.as_posix(), name=run_name)] + loggers.append( + WandbLogger( + name=run_name, + save_dir=log_dir.as_posix(), + project="robotvision", + log_model=True, + ) + ) + return loggers + + +def build_callbacks(checkpoint_dir: Path) -> list[ModelCheckpoint | EarlyStopping]: + return [ + ModelCheckpoint( + monitor=VAL_MONITOR, + mode="min", + save_top_k=1, + dirpath=checkpoint_dir.as_posix(), + filename="sae-scorer-residual-{epoch:02d}-{val_corrected_picked_ade:.4f}", + ), + EarlyStopping( + monitor=VAL_MONITOR, + mode="min", + patience=DEFAULT_PATIENCE, + min_delta=DEFAULT_MIN_DELTA, + ), + ] + + +def trainer_precision() -> tuple[str, str]: + accelerator = "cuda" if torch.cuda.is_available() else "cpu" + precision = "32-true" + if accelerator == "cuda": + precision = "bf16-mixed" if torch.cuda.is_bf16_supported() else "16-mixed" + torch.set_float32_matmul_precision("medium") + return accelerator, precision + + +def main() -> None: + parser = build_arg_parser() + args = parser.parse_args() + + pl.seed_everything(DEFAULT_SEED, workers=True) + + train_loader, val_loader = build_train_val_loaders( + data_dir=args.data_dir, + train_items=args.train_items, + val_items=args.val_items, + batch_size=args.batch_size, + num_workers=args.num_workers, + seed=DEFAULT_SEED, + ) + lit_model = LitSAEScorerResidual( + model_checkpoint_path=args.model_checkpoint_path, + sae_checkpoint_path=args.sae_checkpoint_path, + block_idx=args.block_idx, + lr=args.lr, + weight_decay=args.weight_decay, + hidden_dim=args.hidden_dim, + dropout=args.dropout, + loss_type=args.loss_type, + topk_train=args.topk_train, + topk_weight=args.topk_weight, + rest_weight=args.rest_weight, + rank_loss_weight=args.rank_loss_weight, + rank_temperature=args.rank_temperature, + latent_source=args.latent_source, + ) + + base_path = Path(args.data_dir).parent + run_name = make_run_name() + log_dir = Path(args.log_dir) if args.log_dir is not None else base_path / "logs" + checkpoint_dir = Path(args.checkpoint_dir) if args.checkpoint_dir is not None else base_path / "checkpoints" + log_dir.mkdir(parents=True, exist_ok=True) + checkpoint_dir.mkdir(parents=True, exist_ok=True) + + callbacks = build_callbacks(checkpoint_dir) + checkpoint_callback = callbacks[0] + accelerator, precision = trainer_precision() + + trainer = pl.Trainer( + max_epochs=args.epochs, + accelerator=accelerator, + devices=1, + precision=precision, + logger=build_loggers(log_dir, run_name), + callbacks=callbacks, + log_every_n_steps=DEFAULT_LOG_EVERY_N_STEPS, + ) + trainer.fit(lit_model, train_loader, val_loader) + + if checkpoint_callback.best_model_path: + best_metric = float(checkpoint_callback.best_model_score.item()) + export_best_residual_checkpoint( + checkpoint_callback.best_model_path, + args.output_path, + export_args(args), + best_metric=best_metric, + ) + print(f"Best Lightning checkpoint: {checkpoint_callback.best_model_path}") + print(f"Exported residual checkpoint: {args.output_path}") + print(f"Best val corrected pick ADE: {best_metric:.6f}") + else: + raise RuntimeError("Trainer finished without producing a best checkpoint") + + +if __name__ == "__main__": + main() diff --git a/src/camera-based-e2e/sae_utils.py b/src/camera-based-e2e/sae_utils.py index 9908a1b..5c93d1c 100644 --- a/src/camera-based-e2e/sae_utils.py +++ b/src/camera-based-e2e/sae_utils.py @@ -28,15 +28,37 @@ class ActivationCapture: def __init__(self) -> None: self.activations: torch.Tensor | None = None + self._handle = None def clear(self) -> None: self.activations = None + def register(self, module) -> "ActivationCapture": + """Attach to `module` and retain the handle so the hook can be removed.""" + self.remove() + self._handle = module.register_forward_hook(self) + return self + + def remove(self) -> None: + if self._handle is not None: + self._handle.remove() + self._handle = None + def __call__(self, module, inputs, output) -> None: del module, inputs self.activations = output.detach() +def freeze_module(module) -> None: + for param in module.parameters(): + param.requires_grad = False + + +def set_eval_mode(*modules) -> None: + for module in modules: + module.eval() + + def planner_token_key(block_idx: int) -> str: return f"planner_query_tok_block_{block_idx}" @@ -89,6 +111,35 @@ def infer_sae_paths(run_root: Path, split: str, sae_block: int) -> tuple[Path, P return ckpt_path, token_path, None +def normalize_compiled_state_dict(state_dict: dict, module) -> dict: + """Reconcile ``torch.compile``'s ``._orig_mod.`` key prefixes with `module`. + + A checkpoint saved from a compiled SAE carries ``encoder._orig_mod.weight`` + where an eager module expects ``encoder.weight`` (and vice versa). Loading + across that boundary silently drops every tensor under ``strict=False``, + leaving a randomly initialised SAE that still "loads" successfully. + """ + state_compiled = any("._orig_mod." in key for key in state_dict) + module_compiled = any("._orig_mod." in key for key in module.state_dict()) + if state_compiled == module_compiled: + return state_dict + + normalized = {} + for key, value in state_dict.items(): + for prefix in ("encoder.", "decoder."): + if module_compiled: + if key.startswith(prefix): + key = key.replace(prefix, f"{prefix}_orig_mod.", 1) + break + else: + compiled_prefix = f"{prefix}_orig_mod." + if key.startswith(compiled_prefix): + key = key.replace(compiled_prefix, prefix, 1) + break + normalized[key] = value + return normalized + + def build_sae_from_checkpoint(ckpt: dict, legacy_norm: dict | None = None) -> SparseAutoencoder: sae_type = ckpt.get("sae_type", LEGACY_SAE_TYPE) if sae_type == TOPK_AUX_SAE_TYPE: @@ -109,7 +160,7 @@ def build_sae_from_checkpoint(ckpt: dict, legacy_norm: dict | None = None) -> Sp sae_type=LEGACY_SAE_TYPE, use_encoder_bias=True, ) - model.load_state_dict(ckpt["state_dict"], strict=False) + model.load_state_dict(normalize_compiled_state_dict(ckpt["state_dict"], model), strict=False) if legacy_norm is not None: model.set_legacy_normalization( mean=legacy_norm["mean"], @@ -190,9 +241,7 @@ def load_model_and_sae( sae = build_sae_from_checkpoint(sae_checkpoint) sae.eval() - capture = ActivationCapture() - target_layer = get_sae_target_layer(model, block_idx) - target_layer.register_forward_hook(capture) + capture = ActivationCapture().register(get_sae_target_layer(model, block_idx)) if device is not None: model = model.to(device) diff --git a/src/camera-based-e2e/test_sae_scorer_residual.py b/src/camera-based-e2e/test_sae_scorer_residual.py new file mode 100644 index 0000000..3dd4802 --- /dev/null +++ b/src/camera-based-e2e/test_sae_scorer_residual.py @@ -0,0 +1,209 @@ +""" +Evaluate whether an SAE-informed residual model improves scorer proposal selection. +""" + +from __future__ import annotations + +import argparse +import json +from collections import defaultdict +from pathlib import Path + +import pytorch_lightning as pl +import torch +from torch.utils.data import DataLoader + +from loader import WaymoE2E +from models.base_model import collate_with_images +from sae_scorer_correction import ( + DEFAULT_SEED, + DEFAULT_VAL_SPLIT, + SCORE_METRIC_NAMES, + LATENT_SOURCES, + checkpoint_value, + correction_loss, + default_index_file, + default_num_workers, + extract_batch_targets, + flatten_proposal_batch, + load_model_and_sae, + load_residual_model, + resolve_required_arg, + score_comparison_metrics, +) + + +def evaluate( + residual_model, + loader, + model, + sae, + capture, + device: torch.device, + *, + loss_type: str, + latent_source: str = "sae", +) -> dict[str, float]: + residual_model.eval() + + totals = defaultdict(float) + n_samples = 0 + n_scores = 0 + + with torch.no_grad(): + for batch in loader: + batch_targets = extract_batch_targets( + model, sae, capture, batch, device, latent_source + ) + flat_batch = flatten_proposal_batch(batch_targets) + residual_pred_flat = residual_model( + flat_batch["sae_hidden"], + flat_batch["trajectory"], + flat_batch["score_pred"], + ) + residual_loss = correction_loss( + residual_pred_flat, + flat_batch["residual_target"], + loss_type=loss_type, + ) + corrected_scores = batch_targets["score_pred"] + residual_pred_flat.view_as(batch_targets["score_pred"]) + + sample_count = batch_targets["score_pred"].size(0) + score_count = batch_targets["score_pred"].numel() + n_samples += sample_count + n_scores += score_count + + totals["residual_loss"] += residual_loss.item() * score_count + for key, value in score_comparison_metrics( + batch_targets["score_pred"], + corrected_scores, + batch_targets["true_ade"], + ).items(): + weight = score_count if key in SCORE_METRIC_NAMES else sample_count + totals[key] += value.item() * weight + + if n_samples == 0 or n_scores == 0: + raise RuntimeError("No evaluation batches were processed") + + out = {} + for key, total in totals.items(): + denom = n_scores if "score_" in key or key == "residual_loss" else n_samples + out[key] = total / denom + + out["delta_score_mse"] = out["corrected_score_mse"] - out["baseline_score_mse"] + out["delta_score_mae"] = out["corrected_score_mae"] - out["baseline_score_mae"] + out["delta_picked_ade"] = out["corrected_picked_ade"] - out["baseline_picked_ade"] + out["delta_regret"] = out["corrected_regret"] - out["baseline_regret"] + out["delta_mean_rank"] = out["corrected_mean_rank"] - out["baseline_mean_rank"] + return out + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--model_checkpoint_path", type=str, default=None) + parser.add_argument("--sae_checkpoint_path", type=str, default=None) + parser.add_argument("--residual_checkpoint_path", type=str, required=True) + parser.add_argument("--data_dir", type=str, required=True) + parser.add_argument("--split", type=str, default=DEFAULT_VAL_SPLIT, choices=["train", "val"]) + parser.add_argument("--n_items", type=int, default=25_000) + parser.add_argument("--batch_size", type=int, default=16) + parser.add_argument("--num_workers", type=int, default=default_num_workers()) + parser.add_argument("--block_idx", type=int, default=None) + parser.add_argument("--loss_type", type=str, default=None, choices=["mse", "huber"]) + parser.add_argument("--latent_source", type=str, default=None, choices=LATENT_SOURCES) + parser.add_argument("--output_json", type=str, default=None) + args = parser.parse_args() + + pl.seed_everything(DEFAULT_SEED, workers=True) + + device = torch.device("cpu" if not torch.cuda.is_available() else "cuda") + + residual_model, residual_checkpoint = load_residual_model( + args.residual_checkpoint_path, + device=device, + ) + model_checkpoint_path = resolve_required_arg( + args.model_checkpoint_path, + residual_checkpoint, + "model_checkpoint_path", + ) + sae_checkpoint_path = resolve_required_arg( + args.sae_checkpoint_path, + residual_checkpoint, + "sae_checkpoint_path", + ) + block_idx = int(resolve_required_arg(args.block_idx, residual_checkpoint, "block_idx")) + loss_type = args.loss_type or checkpoint_value(residual_checkpoint, "loss_type") or "huber" + latent_source = ( + args.latent_source or checkpoint_value(residual_checkpoint, "latent_source") or "sae" + ) + + loader = DataLoader( + WaymoE2E( + indexFile=default_index_file(args.split), + data_dir=args.data_dir, + n_items=args.n_items, + seed=DEFAULT_SEED, + ), + batch_size=args.batch_size, + num_workers=args.num_workers, + collate_fn=collate_with_images, + persistent_workers=False, + pin_memory=False, + shuffle=False, + ) + + model, sae, capture, _ = load_model_and_sae( + model_checkpoint_path=model_checkpoint_path, + sae_checkpoint_path=sae_checkpoint_path, + block_idx=block_idx, + device=device, + ) + + try: + metrics = evaluate( + residual_model, + loader, + model, + sae, + capture, + device, + loss_type=loss_type, + latent_source=latent_source, + ) + finally: + capture.remove() + + print(f"Residual model checkpoint: {args.residual_checkpoint_path}") + print(f"Backbone model checkpoint: {model_checkpoint_path}") + print(f"SAE checkpoint: {sae_checkpoint_path}") + print(f"Using block index: {block_idx}") + print(f"Using residual loss type: {loss_type}") + print(f"Using latent source: {latent_source}") + if "best_metric" in residual_checkpoint: + print(f"Saved best validation corrected pick ADE: {residual_checkpoint['best_metric']:.6f}") + print(f"Score MSE: {metrics['baseline_score_mse']:.6f} -> {metrics['corrected_score_mse']:.6f} (delta {metrics['delta_score_mse']:+.6f})") + print(f"Score MAE: {metrics['baseline_score_mae']:.6f} -> {metrics['corrected_score_mae']:.6f} (delta {metrics['delta_score_mae']:+.6f})") + print(f"Picked ADE: {metrics['baseline_picked_ade']:.6f} -> {metrics['corrected_picked_ade']:.6f} (delta {metrics['delta_picked_ade']:+.6f})") + print(f"Oracle ADE: {metrics['baseline_oracle_ade']:.6f}") + print(f"Regret: {metrics['baseline_regret']:.6f} -> {metrics['corrected_regret']:.6f} (delta {metrics['delta_regret']:+.6f})") + print(f"Mean rank: {metrics['baseline_mean_rank']:.6f} -> {metrics['corrected_mean_rank']:.6f} (delta {metrics['delta_mean_rank']:+.6f})") + for topk in (1, 5, 10): + key = f"top{topk}_acc" + baseline_key = f"baseline_{key}" + corrected_key = f"corrected_{key}" + if baseline_key in metrics: + delta = metrics[corrected_key] - metrics[baseline_key] + print(f"{key}: {metrics[baseline_key]:.6f} -> {metrics[corrected_key]:.6f} (delta {delta:+.6f})") + print(f"Spearman: {metrics['baseline_spearman']:.6f} -> {metrics['corrected_spearman']:.6f}") + + if args.output_json is not None: + output_path = Path(args.output_json) + output_path.parent.mkdir(parents=True, exist_ok=True) + with output_path.open("w") as f: + json.dump(metrics, f, indent=2, sort_keys=True) + print(f"Saved metrics to {args.output_json}") + + +if __name__ == "__main__": + main()