diff --git a/.codex b/.codex new file mode 100644 index 0000000..e69de29 diff --git a/.gitignore b/.gitignore index 53eac37..0d7aded 100644 --- a/.gitignore +++ b/.gitignore @@ -217,3 +217,7 @@ cython_debug/ marimo/_static/ marimo/_lsp/ __marimo__/ + +*.pt +*.ckpt +.codex/ \ No newline at end of file diff --git a/sae_checkpoints/sae_checkpoints.tar.gz b/sae_checkpoints/sae_checkpoints.tar.gz new file mode 100644 index 0000000..1ae8720 Binary files /dev/null and b/sae_checkpoints/sae_checkpoints.tar.gz differ diff --git a/sbatch/extract_tok.sbatch b/sbatch/extract_tok.sbatch new file mode 100644 index 0000000..2746137 --- /dev/null +++ b/sbatch/extract_tok.sbatch @@ -0,0 +1,30 @@ +#!/bin/bash +#SBATCH --job-name=extract_both +#SBATCH --output=/scratch/gilbreth/chang899/codes/int/src/camera-based-e2e/log/extract_both-%j.out +#SBATCH --error=/scratch/gilbreth/chang899/codes/int/src/camera-based-e2e/log/extract_both-%j.err +# +#SBATCH --account=labi +#SBATCH --partition=a10 +#SBATCH --gres=gpu:1 +#SBATCH --time=84:00:00 +#SBATCH --mem=64G +#SBATCH --cpus-per-task=8 + +module load anaconda +conda activate /scratch/gilbreth/chang899/conda_envs/lead_ltf + +cd /scratch/gilbreth/chang899/codes/int/src/camera-based-e2e + +mkdir -p log + +srun python extract_planner_tok.py \ + --checkpoint camera-e2e-epoch=04-val_loss=2.90.ckpt \ + --data_dir /scratch/gilbreth/chang899/waymo_data/waymo_open_dataset_end_to_end_camera_v_1_0_0 \ + --index_file index_train.pkl \ + --output_path planner_tokens_train.pt + +srun python extract_planner_tok.py \ + --checkpoint camera-e2e-epoch=04-val_loss=2.90.ckpt \ + --data_dir /scratch/gilbreth/chang899/waymo_data/waymo_open_dataset_end_to_end_camera_v_1_0_0 \ + --index_file index_val.pkl \ + --output_path planner_tokens_val.pt diff --git a/sbatch/train_SAE.sbatch b/sbatch/train_SAE.sbatch new file mode 100644 index 0000000..426b30f --- /dev/null +++ b/sbatch/train_SAE.sbatch @@ -0,0 +1,24 @@ +#!/bin/bash +#SBATCH --job-name=train_sae +#SBATCH --output=/scratch/gilbreth/chang899/codes/int/src/camera-based-e2e/log/train_sae-%j.out +#SBATCH --error=/scratch/gilbreth/chang899/codes/int/src/camera-based-e2e/log/train_sae-%j.err +# +#SBATCH --account=csso +#SBATCH --partition=a10 +#SBATCH --gres=gpu:1 +#SBATCH --time=84:00:00 +#SBATCH --mem=64G +#SBATCH --cpus-per-task=8 + +module load anaconda +conda activate /scratch/gilbreth/chang899/conda_envs/lead_ltf + +cd /scratch/gilbreth/chang899/codes/int/src/camera-based-e2e + +mkdir -p log + +srun python train_sae.py \ + --output_dir /scratch/gilbreth/chang899/codes/int/src/camera-based-e2e/output \ + --train_dataset planner_tokens_train.pt \ + --val_dataset planner_tokens_val.pt + diff --git a/sbatch/train_deep_mono.sbatch b/sbatch/train_deep_mono.sbatch new file mode 100644 index 0000000..e55de7e --- /dev/null +++ b/sbatch/train_deep_mono.sbatch @@ -0,0 +1,21 @@ +#!/bin/bash +#SBATCH --job-name=train_deep_mono +#SBATCH --output=/scratch/gilbreth/chang899/codes/int/src/camera-based-e2e/log/deep_mono-%j.out +#SBATCH --error=/scratch/gilbreth/chang899/codes/int/src/camera-based-e2e/log/deep_mono-%j.err +# +#SBATCH --account=labi +#SBATCH --partition=a10 +#SBATCH --gres=gpu:1 +#SBATCH --time=84:00:00 +#SBATCH --mem=64G +#SBATCH --cpus-per-task=8 + +module load anaconda +conda activate /scratch/gilbreth/chang899/conda_envs/lead_ltf + +cd /scratch/gilbreth/chang899/codes/int/src/camera-based-e2e + +mkdir -p log + +srun python train.py \ + --data_dir /scratch/gilbreth/chang899/waymo_data \ No newline at end of file diff --git a/src/camera-based-e2e/analyze_sae_ade_intervention.py b/src/camera-based-e2e/analyze_sae_ade_intervention.py new file mode 100644 index 0000000..109d9a0 --- /dev/null +++ b/src/camera-based-e2e/analyze_sae_ade_intervention.py @@ -0,0 +1,336 @@ +import argparse +import csv +from pathlib import Path + +import torch + +from extract_planner_tok import load_model +from models.sae import SparseAutoencoder +from sae_utils import ( + DEFAULT_SAE_BLOCK, + build_sae_from_checkpoint, + collate_dataset_indices, + dataset_from_token_blob, + default_device, + encode_tensor_batchwise, + load_sae_bundle, + prepare_replay_context, + resolve_token_tensor, +) + + +def parse_int_list(text: str) -> list[int]: + return [int(part.strip()) for part in text.split(",") if part.strip()] + + +def parse_feature_spec(text: str, latent_dim: int) -> list[int]: + raw = text.strip().lower() + if raw in {"all", "*"}: + return list(range(latent_dim)) + + if ":" in raw and "," not in raw: + parts = raw.split(":") + if len(parts) not in {2, 3}: + raise ValueError(f"Unsupported feature slice: {text}") + start = int(parts[0]) if parts[0] else 0 + stop = int(parts[1]) if parts[1] else latent_dim + step = int(parts[2]) if len(parts) == 3 and parts[2] else 1 + return list(range(start, stop, step)) + + return parse_int_list(text) + + +def parse_float_list(text: str) -> list[float]: + return [float(part.strip()) for part in text.split(",") if part.strip()] + + +def clip_feature_list( + features: list[int], + latent_dim: int, + feature_start: int | None, + feature_end: int | None, +) -> list[int]: + clipped = [] + for feature_idx in features: + if feature_idx < 0 or feature_idx >= latent_dim: + continue + if feature_start is not None and feature_idx < feature_start: + continue + if feature_end is not None and feature_idx >= feature_end: + continue + clipped.append(feature_idx) + return clipped + + +def rank_along_levels(values: torch.Tensor) -> torch.Tensor: + order = torch.argsort(values, dim=0) + ranks = torch.empty_like(order, dtype=torch.float32) + base = torch.arange(values.size(0), device=values.device, dtype=torch.float32).unsqueeze(1) + ranks.scatter_(0, order, base.expand_as(order)) + return ranks + + +def spearman_vs_level(curves: torch.Tensor) -> torch.Tensor: + if curves.size(0) < 2: + return torch.zeros(curves.size(1), dtype=torch.float32, device=curves.device) + x = torch.arange(curves.size(0), device=curves.device, dtype=torch.float32) + x = x - x.mean() + x_denom = torch.sqrt((x * x).sum()).clamp_min(1e-6) + + y = rank_along_levels(curves) + y = y - y.mean(dim=0, keepdim=True) + y_denom = torch.sqrt((y * y).sum(dim=0)).clamp_min(1e-6) + return (x[:, None] * y).sum(dim=0) / (x_denom * y_denom) + + +def selected_and_oracle_ade( + out: dict, + future_batch: torch.Tensor, + num_proposals: int, + horizon: int, +) -> tuple[torch.Tensor, torch.Tensor]: + batch = future_batch.size(0) + traj = out["trajectory"].view(batch, num_proposals, horizon, 2) + scores = out["scores"] + dist = torch.norm(traj - future_batch[:, None], dim=-1) + ade_per_mode = dist.mean(dim=-1) + row_idx = torch.arange(batch, device=future_batch.device) + selected_idx = scores.argmin(dim=1) + selected_ade = ade_per_mode[row_idx, selected_idx] + oracle_ade = ade_per_mode.min(dim=1).values + return selected_ade, oracle_ade + + +def summarize_feature( + feature_idx: int, + token_tensor_cpu: torch.Tensor, + z_all: torch.Tensor, + scales: torch.Tensor, + sae: SparseAutoencoder, + planner_model, + lit_model, + sae_block: int, + past: torch.Tensor, + future: torch.Tensor, + relevant_scenes: int, + alphas: list[float], + scene_selection: str, + random_seed: int, + dataset, + device: torch.device, +) -> dict: + feature_act = z_all[:, feature_idx] + active_mask = feature_act > 0 + active_count = int(active_mask.sum().item()) + if scene_selection == "relevant": + keep = min(active_count, relevant_scenes) + else: + keep = min(z_all.size(0), relevant_scenes) + if keep == 0: + return { + "feature_idx": feature_idx, + "active_count": 0, + "relevant_scene_count": 0, + "scene_selection": scene_selection, + "random_seed": random_seed, + "mean_delta_selected_ade": 0.0, + "frac_improved_selected_ade": 0.0, + "frac_monotone_down_selected_ade": 0.0, + "mean_rho_selected_ade": 0.0, + "mean_delta_selected_ade_hard": 0.0, + "frac_improved_selected_ade_hard": 0.0, + "mean_delta_oracle_ade": 0.0, + "frac_improved_oracle_ade": 0.0, + } + + if scene_selection == "relevant": + rel_idx = torch.topk(feature_act, k=keep).indices + elif scene_selection == "random": + generator = torch.Generator(device="cpu") + generator.manual_seed(random_seed) + rel_idx = torch.randperm(z_all.size(0), generator=generator)[:keep] + else: + raise ValueError(f"Unsupported scene_selection={scene_selection}") + + z_rel = z_all[rel_idx].clone().to(device) + future_rel = future[rel_idx].to(device) + base_x = token_tensor_cpu[rel_idx].to(device) + base_activation = z_rel[:, feature_idx].clone() + scale = float(scales[feature_idx].item()) + + if sae_block == DEFAULT_SAE_BLOCK: + replay_context = None + past_rel = past[rel_idx].to(device) + else: + batch = collate_dataset_indices(dataset, rel_idx) + replay_context = prepare_replay_context(planner_model, lit_model, batch, device=device) + past_rel = replay_context["past"] + + selected_curves = [] + oracle_curves = [] + with torch.no_grad(): + for alpha in alphas: + z_mod = z_rel.clone() + z_mod[:, feature_idx] = (base_activation + alpha * scale).clamp_min(0.0) + recon_query = sae.decode_to_input(z_mod, reference_x=base_x) + if sae_block == DEFAULT_SAE_BLOCK: + out = planner_model.forward_from_planner_query_tok(recon_query, past_rel) + else: + out = planner_model.forward_from_block_query_tok( + recon_query, + past_rel, + replay_context["tokens"], + start_block=sae_block, + ) + selected_ade, oracle_ade = selected_and_oracle_ade( + out=out, + future_batch=future_rel, + num_proposals=planner_model.n_proposals, + horizon=planner_model.horizon, + ) + selected_curves.append(selected_ade) + oracle_curves.append(oracle_ade) + + selected_curves = torch.stack(selected_curves, dim=0) + oracle_curves = torch.stack(oracle_curves, dim=0) + + base_selected = selected_curves[0] + final_selected = selected_curves[-1] + delta_selected = final_selected - base_selected + + base_oracle = oracle_curves[0] + final_oracle = oracle_curves[-1] + delta_oracle = final_oracle - base_oracle + + hard_thresh = torch.quantile(base_selected, 0.75) + hard_mask = base_selected >= hard_thresh + + diffs = selected_curves[1:] - selected_curves[:-1] + monotone_down = (diffs <= 1e-8).all(dim=0) + rho_selected = spearman_vs_level(selected_curves) + + return { + "feature_idx": feature_idx, + "active_count": active_count, + "relevant_scene_count": keep, + "scene_selection": scene_selection, + "random_seed": random_seed, + "intervention_scale": scale, + "mean_delta_selected_ade": float(delta_selected.mean().item()), + "frac_improved_selected_ade": float((delta_selected < 0).float().mean().item()), + "frac_monotone_down_selected_ade": float(monotone_down.float().mean().item()), + "mean_rho_selected_ade": float(rho_selected.mean().item()), + "mean_delta_selected_ade_hard": float(delta_selected[hard_mask].mean().item()), + "frac_improved_selected_ade_hard": float((delta_selected[hard_mask] < 0).float().mean().item()), + "mean_delta_oracle_ade": float(delta_oracle.mean().item()), + "frac_improved_oracle_ade": float((delta_oracle < 0).float().mean().item()), + } + + +def write_csv(path: Path, rows: list[dict]) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + if not rows: + return + with path.open("w", newline="") as f: + writer = csv.DictWriter(f, fieldnames=list(rows[0].keys())) + writer.writeheader() + writer.writerows(rows) + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument("--run_root", type=str, required=True) + parser.add_argument("--planner_checkpoint", type=str, required=True) + parser.add_argument("--split", type=str, default="val", choices=["train", "val"]) + parser.add_argument("--features", type=str, required=True) + parser.add_argument("--sae_block", type=int, default=3) + parser.add_argument("--data_dir", type=str, default=None) + parser.add_argument("--index_file", type=str, default=None) + parser.add_argument("--feature_start", type=int, default=None) + parser.add_argument("--feature_end", type=int, default=None) + parser.add_argument("--alphas", type=str, default="0,0.5,1.0,2.0") + parser.add_argument("--relevant_scenes_per_feature", type=int, default=384) + parser.add_argument("--scene_selection", type=str, default="relevant", choices=["relevant", "random"]) + parser.add_argument("--random_seeds", type=str, default="0") + parser.add_argument("--batch_size", type=int, default=4096) + parser.add_argument("--device", type=str, default=default_device()) + parser.add_argument("--output_csv", type=str, required=True) + args = parser.parse_args() + + device = torch.device(args.device) + run_root = Path(args.run_root) + + bundle = load_sae_bundle(run_root, args.split, args.sae_block, map_location="cpu") + sae_ckpt = bundle["ckpt"] + token_blob = bundle["token_blob"] + token_tensor, token_key = resolve_token_tensor(token_blob, args.sae_block) + + sae = build_sae_from_checkpoint(sae_ckpt, bundle["legacy_norm"]) + sae.to(device) + sae.eval() + + planner_model, lit_model = load_model(args.planner_checkpoint, device=device) + planner_model.eval() + + past = token_blob["past"].float() + future = token_blob["future"].float() + + z_all = encode_tensor_batchwise( + sae, + token_tensor, + batch_size=args.batch_size, + device=device, + ) + latent_dim = sae_ckpt["latent_dim"] + features = clip_feature_list( + parse_feature_spec(args.features, latent_dim=latent_dim), + latent_dim=latent_dim, + feature_start=args.feature_start, + feature_end=args.feature_end, + ) + + active_mask = z_all > 0 + active_count = active_mask.sum(dim=0) + active_sum = z_all.sum(dim=0) + active_sum_sq = (z_all * z_all).sum(dim=0) + active_mean = active_sum / active_count.clamp_min(1) + active_var = active_sum_sq / active_count.clamp_min(1) - active_mean.square() + active_std = torch.sqrt(active_var.clamp_min(0.0)) + scales = torch.maximum(active_std, 0.25 * active_mean).clamp_min(0.05) + + dataset = None + if args.sae_block != DEFAULT_SAE_BLOCK: + dataset = dataset_from_token_blob( + token_blob, + data_dir=args.data_dir, + index_file=args.index_file, + ) + + rows = [] + for random_seed in parse_int_list(args.random_seeds): + for feature_idx in features: + row = summarize_feature( + feature_idx=feature_idx, + token_tensor_cpu=token_tensor, + z_all=z_all, + scales=scales, + sae=sae, + planner_model=planner_model, + lit_model=lit_model, + sae_block=args.sae_block, + past=past, + future=future, + relevant_scenes=args.relevant_scenes_per_feature, + alphas=parse_float_list(args.alphas), + scene_selection=args.scene_selection, + random_seed=random_seed, + dataset=dataset, + device=device, + ) + rows.append(row) + print(row) + + rows.sort(key=lambda row: row["mean_delta_selected_ade"]) + write_csv(Path(args.output_csv), rows) + print(f"Used token key {token_key}") + print(f"Saved CSV to {args.output_csv}") diff --git a/src/camera-based-e2e/analyze_sae_control.py b/src/camera-based-e2e/analyze_sae_control.py new file mode 100644 index 0000000..5d8d24e --- /dev/null +++ b/src/camera-based-e2e/analyze_sae_control.py @@ -0,0 +1,524 @@ +import argparse +import csv +from pathlib import Path + +import torch + +from extract_planner_tok import load_model +from models.sae import SparseAutoencoder +from sae_utils import ( + DEFAULT_SAE_BLOCK, + build_sae_from_checkpoint, + collate_dataset_indices, + dataset_from_token_blob, + default_analysis_dir, + default_device, + encode_tensor_batchwise, + load_sae_bundle, + planner_inputs_from_collated_batch, + prepare_replay_context, + resolve_token_tensor, +) + + +STAT_NAMES = ( + "final_lateral_disp", + "avg_curvature", + "brake_mag", + "accel_mag", + "score_margin", + "proposal_spread", +) + + +def parse_float_list(text: str) -> list[float]: + return [float(part.strip()) for part in text.split(",") if part.strip()] + + +def infer_checkpoint_path(run_root: Path, token_blob: dict) -> str: + meta_path = token_blob.get("meta", {}).get("checkpoint") + if meta_path and Path(meta_path).exists(): + return meta_path + + repo_default = Path(__file__).resolve().parent / "camera-e2e-epoch=04-val_loss=2.90.ckpt" + if repo_default.exists(): + return str(repo_default) + + raise FileNotFoundError( + "Could not infer planner checkpoint path. Pass --planner_checkpoint explicitly." + ) + + +def reshape_trajectory(trajectory: torch.Tensor, num_proposals: int, horizon: int) -> torch.Tensor: + if trajectory.ndim == 4: + return trajectory + return trajectory.view(trajectory.size(0), num_proposals, horizon, 2) + + +def reshape_controls(controls: torch.Tensor, num_proposals: int, horizon: int) -> torch.Tensor: + return controls.view(controls.size(0), num_proposals, horizon, 2) + + +def average_curvature(traj: torch.Tensor) -> torch.Tensor: + seg = traj[:, 1:] - traj[:, :-1] + seg_norm = torch.norm(seg, dim=-1).clamp_min(1e-6) + heading = torch.atan2(seg[..., 1], seg[..., 0]) + d_heading = torch.atan2( + torch.sin(heading[:, 1:] - heading[:, :-1]), + torch.cos(heading[:, 1:] - heading[:, :-1]), + ).abs() + ds = 0.5 * (seg_norm[:, 1:] + seg_norm[:, :-1]) + curvature = d_heading / ds.clamp_min(1e-6) + return curvature.mean(dim=1) + + +def compute_output_stats( + trajectory: torch.Tensor, + scores: torch.Tensor, + controls: torch.Tensor, + num_proposals: int, + horizon: int, +) -> dict[str, torch.Tensor]: + traj = reshape_trajectory(trajectory, num_proposals=num_proposals, horizon=horizon) + ctrl = reshape_controls(controls, num_proposals=num_proposals, horizon=horizon) + + best_idx = scores.argmin(dim=1) + row_idx = torch.arange(traj.size(0), device=traj.device) + + selected_traj = traj[row_idx, best_idx] + selected_ctrl = ctrl[row_idx, best_idx] + accel = selected_ctrl[..., 0] + + sorted_scores = scores.sort(dim=1).values + if scores.size(1) > 1: + score_margin = sorted_scores[:, 1] - sorted_scores[:, 0] + else: + score_margin = torch.zeros(scores.size(0), device=scores.device, dtype=scores.dtype) + + traj_mean = traj.mean(dim=1, keepdim=True) + proposal_spread = torch.norm(traj - traj_mean, dim=-1).mean(dim=(1, 2)) + + return { + "final_lateral_disp": selected_traj[:, -1, 1], + "avg_curvature": average_curvature(selected_traj), + "brake_mag": (-accel).clamp_min(0).mean(dim=1), + "accel_mag": accel.clamp_min(0).mean(dim=1), + "score_margin": score_margin, + "proposal_spread": proposal_spread, + } + + +def compute_global_stat_stds( + token_blob: dict, + num_proposals: int, + horizon: int, +) -> dict[str, float]: + stats = compute_output_stats( + trajectory=token_blob["trajectory"].float(), + scores=token_blob["scores"].float(), + controls=token_blob["controls"].float(), + num_proposals=num_proposals, + horizon=horizon, + ) + return { + name: float(values.std(unbiased=False).clamp_min(1e-6).item()) + for name, values in stats.items() + } + + +def rank_along_levels(values: torch.Tensor) -> torch.Tensor: + order = torch.argsort(values, dim=0) + ranks = torch.empty_like(order, dtype=torch.float32) + base = torch.arange(values.size(0), device=values.device, dtype=torch.float32).unsqueeze(1) + ranks.scatter_(0, order, base.expand_as(order)) + return ranks + + +def spearman_vs_level(curves: torch.Tensor) -> torch.Tensor: + if curves.size(0) < 2: + return torch.zeros(curves.size(1), dtype=torch.float32) + x = torch.arange(curves.size(0), device=curves.device, dtype=torch.float32) + x = x - x.mean() + x_denom = torch.sqrt((x * x).sum()).clamp_min(1e-6) + + y = rank_along_levels(curves) + y = y - y.mean(dim=0, keepdim=True) + y_denom = torch.sqrt((y * y).sum(dim=0)).clamp_min(1e-6) + return (x[:, None] * y).sum(dim=0) / (x_denom * y_denom) + + +def intervention_levels_for_feature( + base_activation: torch.Tensor, + scale: float, + alphas: list[float], +) -> list[torch.Tensor]: + return [(base_activation + alpha * scale).clamp_min(0.0) for alpha in alphas] + + +def run_intervention_outputs( + *, + sae_block: int, + sae: SparseAutoencoder, + planner_model, + lit_model, + token_tensor_cpu: torch.Tensor, + past_cpu: torch.Tensor, + relevant_indices: torch.Tensor, + level_values: list[torch.Tensor], + base_z_cpu: torch.Tensor, + dataset, + device: torch.device, +) -> list[dict[str, torch.Tensor]]: + base_x = token_tensor_cpu[relevant_indices].to(device) + base_z = base_z_cpu[relevant_indices].to(device) + + replay_context = None + if sae_block == DEFAULT_SAE_BLOCK: + past = past_cpu[relevant_indices].to(device) + else: + batch = collate_dataset_indices(dataset, relevant_indices) + replay_context = prepare_replay_context(planner_model, lit_model, batch, device=device) + past = replay_context["past"] + + outputs = [] + sae.eval() + planner_model.eval() + with torch.no_grad(): + for level in level_values: + z_mod = base_z.clone() + feature_idx = None + # Caller overwrites a single feature before passing level values, + # so reuse the matching level tensor positionally below. + if level.ndim != 1: + raise ValueError("Expected per-scene latent level vector.") + # Identify the edited feature by comparing shapes later in caller. + # z_mod is updated in caller before use. + outputs.append((z_mod, past, base_x, replay_context)) + return outputs + + +def analyze_feature( + feature_idx: int, + token_tensor_cpu: torch.Tensor, + base_z_cpu: torch.Tensor, + feature_active_count: int, + relevant_indices: torch.Tensor, + scale: float, + alphas: list[float], + sae: SparseAutoencoder, + planner_model, + lit_model, + sae_block: int, + past_cpu: torch.Tensor, + global_stat_stds: dict[str, float], + dataset, + device: torch.device, +) -> tuple[list[dict], dict]: + if relevant_indices.numel() == 0: + rows = [] + for stat_name in STAT_NAMES: + rows.append( + { + "feature_idx": feature_idx, + "stat_name": stat_name, + "active_count": feature_active_count, + "relevant_scene_count": 0, + "intervention_scale": scale, + "mean_delta_max": 0.0, + "std_effect_max": 0.0, + "mean_rho": 0.0, + "frac_consistent": 0.0, + "selectivity_share": 0.0, + "target_vs_next_ratio": 0.0, + "control_score": 0.0, + } + ) + summary = { + "feature_idx": feature_idx, + "best_stat": "NONE", + "best_control_score": 0.0, + "best_std_effect_max": 0.0, + "best_mean_rho": 0.0, + "best_frac_consistent": 0.0, + "best_selectivity_share": 0.0, + "active_count": feature_active_count, + "relevant_scene_count": 0, + "intervention_scale": scale, + "runner_up_stat": "NONE", + } + return rows, summary + + base_x = token_tensor_cpu[relevant_indices].to(device) + base_z = base_z_cpu[relevant_indices].to(device) + base_activation = base_z[:, feature_idx].clone() + level_values = intervention_levels_for_feature(base_activation, scale=scale, alphas=alphas) + + replay_context = None + if sae_block == DEFAULT_SAE_BLOCK: + past = past_cpu[relevant_indices].to(device) + else: + batch = collate_dataset_indices(dataset, relevant_indices) + replay_context = prepare_replay_context(planner_model, lit_model, batch, device=device) + past = replay_context["past"] + + stat_curves = {name: [] for name in STAT_NAMES} + sae.eval() + planner_model.eval() + with torch.no_grad(): + for level in level_values: + z_mod = base_z.clone() + z_mod[:, feature_idx] = level + recon_query = sae.decode_to_input(z_mod, reference_x=base_x) + if sae_block == DEFAULT_SAE_BLOCK: + out = planner_model.forward_from_planner_query_tok(recon_query, past) + else: + out = planner_model.forward_from_block_query_tok( + recon_query, + past, + replay_context["tokens"], + start_block=sae_block, + ) + stats = compute_output_stats( + trajectory=out["trajectory"], + scores=out["scores"], + controls=out["controls"], + num_proposals=planner_model.n_proposals, + horizon=planner_model.horizon, + ) + for stat_name in STAT_NAMES: + stat_curves[stat_name].append(stats[stat_name].cpu()) + + rows = [] + std_effects = {} + metric_cache = {} + for stat_name in STAT_NAMES: + curves = torch.stack(stat_curves[stat_name], dim=0) + delta = curves - curves[0:1] + delta_max = delta[-1] + mean_delta_max = float(delta_max.mean().item()) + std_effect_max = mean_delta_max / global_stat_stds[stat_name] + rho_scene = spearman_vs_level(curves) + mean_rho = float(rho_scene.mean().item()) + sign = 1.0 if mean_delta_max >= 0 else -1.0 + consistent = ((sign * rho_scene) > 0.5) & ((sign * delta_max) > 0) + frac_consistent = float(consistent.float().mean().item()) + std_effects[stat_name] = abs(std_effect_max) + metric_cache[stat_name] = { + "mean_delta_max": mean_delta_max, + "std_effect_max": std_effect_max, + "mean_rho": mean_rho, + "frac_consistent": frac_consistent, + } + + effect_sum = sum(std_effects.values()) + 1e-8 + sorted_effects = sorted(std_effects.items(), key=lambda item: item[1], reverse=True) + + for stat_name in STAT_NAMES: + next_best = max( + (value for other_name, value in std_effects.items() if other_name != stat_name), + default=0.0, + ) + selectivity_share = std_effects[stat_name] / effect_sum + target_vs_next_ratio = std_effects[stat_name] / (next_best + 1e-8) + control_score = ( + std_effects[stat_name] + * abs(metric_cache[stat_name]["mean_rho"]) + * metric_cache[stat_name]["frac_consistent"] + * selectivity_share + ) + rows.append( + { + "feature_idx": feature_idx, + "stat_name": stat_name, + "active_count": feature_active_count, + "relevant_scene_count": int(relevant_indices.numel()), + "intervention_scale": scale, + "mean_delta_max": metric_cache[stat_name]["mean_delta_max"], + "std_effect_max": metric_cache[stat_name]["std_effect_max"], + "mean_rho": metric_cache[stat_name]["mean_rho"], + "frac_consistent": metric_cache[stat_name]["frac_consistent"], + "selectivity_share": selectivity_share, + "target_vs_next_ratio": target_vs_next_ratio, + "control_score": control_score, + } + ) + + best_row = max(rows, key=lambda row: row["control_score"]) + summary = { + "feature_idx": feature_idx, + "best_stat": best_row["stat_name"], + "best_control_score": best_row["control_score"], + "best_std_effect_max": best_row["std_effect_max"], + "best_mean_rho": best_row["mean_rho"], + "best_frac_consistent": best_row["frac_consistent"], + "best_selectivity_share": best_row["selectivity_share"], + "active_count": feature_active_count, + "relevant_scene_count": int(relevant_indices.numel()), + "intervention_scale": scale, + "runner_up_stat": sorted_effects[1][0] if len(sorted_effects) > 1 else "NONE", + } + return rows, summary + + +def write_csv(path: Path, rows: list[dict]) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + if not rows: + return + fieldnames = [] + seen = set() + for row in rows: + for key in row.keys(): + if key not in seen: + seen.add(key) + fieldnames.append(key) + with path.open("w", newline="") as f: + writer = csv.DictWriter(f, fieldnames=fieldnames) + writer.writeheader() + writer.writerows(rows) + + +def print_top_features(summary_rows: list[dict], stat_rows: list[dict], top_k: int) -> None: + print("Top controllable features overall:") + top_features = sorted(summary_rows, key=lambda row: row["best_control_score"], reverse=True)[:top_k] + for rank, row in enumerate(top_features, start=1): + print( + f" {rank}. feature={row['feature_idx']} best_stat={row['best_stat']} " + f"score={row['best_control_score']:.4f} std_effect={row['best_std_effect_max']:+.4f} " + f"rho={row['best_mean_rho']:+.4f} consistent={row['best_frac_consistent']:.3f} " + f"selectivity={row['best_selectivity_share']:.3f}" + ) + print("") + + for stat_name in STAT_NAMES: + top_rows = sorted( + (row for row in stat_rows if row["stat_name"] == stat_name), + key=lambda row: row["control_score"], + reverse=True, + )[:top_k] + print(f"Top features for {stat_name}:") + for rank, row in enumerate(top_rows, start=1): + print( + f" {rank}. feature={row['feature_idx']} score={row['control_score']:.4f} " + f"std_effect={row['std_effect_max']:+.4f} rho={row['mean_rho']:+.4f} " + f"consistent={row['frac_consistent']:.3f} selectivity={row['selectivity_share']:.3f}" + ) + print("") + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument("--run_root", type=str, required=True) + parser.add_argument("--planner_checkpoint", type=str, default=None) + parser.add_argument("--split", type=str, default="val", choices=["train", "val"]) + parser.add_argument("--output_dir", type=str, default=None) + parser.add_argument("--sae_block", type=int, default=3) + parser.add_argument("--data_dir", type=str, default=None) + parser.add_argument("--index_file", type=str, default=None) + parser.add_argument("--device", type=str, default=default_device()) + parser.add_argument("--encode_batch_size", type=int, default=4096) + parser.add_argument("--relevant_scenes_per_feature", type=int, default=64) + parser.add_argument("--alphas", type=str, default="0,0.5,1.0,2.0") + parser.add_argument("--min_scale", type=float, default=0.05) + parser.add_argument("--max_features", type=int, default=None) + parser.add_argument("--top_k", type=int, default=15) + args = parser.parse_args() + + run_root = Path(args.run_root) + output_dir = default_analysis_dir(run_root, args.sae_block, args.output_dir) + output_dir.mkdir(parents=True, exist_ok=True) + device = torch.device(args.device) + alphas = parse_float_list(args.alphas) + + bundle = load_sae_bundle(run_root, args.split, args.sae_block, map_location="cpu") + sae_ckpt = bundle["ckpt"] + token_blob = bundle["token_blob"] + token_tensor, token_key = resolve_token_tensor(token_blob, args.sae_block) + + planner_checkpoint = args.planner_checkpoint or infer_checkpoint_path(run_root, token_blob) + planner_model, lit_model = load_model(planner_checkpoint, device=device) + + global_stat_stds = compute_global_stat_stds( + token_blob=token_blob, + num_proposals=planner_model.n_proposals, + horizon=planner_model.horizon, + ) + print("Global stat stds:") + for stat_name in STAT_NAMES: + print(f" {stat_name}: {global_stat_stds[stat_name]:.4f}") + print("") + + sae = build_sae_from_checkpoint(sae_ckpt, bundle["legacy_norm"]) + sae.to(device) + sae.eval() + + past_cpu = token_blob["past"].float() + base_z_cpu = encode_tensor_batchwise( + sae, + token_tensor, + batch_size=args.encode_batch_size, + device=device, + ) + + active_mask = base_z_cpu > 0 + active_count = active_mask.sum(dim=0) + active_sum = base_z_cpu.sum(dim=0) + active_sum_sq = (base_z_cpu * base_z_cpu).sum(dim=0) + active_mean = active_sum / active_count.clamp_min(1) + active_var = active_sum_sq / active_count.clamp_min(1) - active_mean.square() + active_std = torch.sqrt(active_var.clamp_min(0)) + intervention_scale = torch.maximum(active_std, 0.25 * active_mean).clamp_min(args.min_scale) + + top_k = min(args.relevant_scenes_per_feature, base_z_cpu.size(0)) + _, top_indices = torch.topk(base_z_cpu, k=top_k, dim=0) + + dataset = None + if args.sae_block != DEFAULT_SAE_BLOCK: + dataset = dataset_from_token_blob( + token_blob, + data_dir=args.data_dir, + index_file=args.index_file, + ) + + feature_limit = base_z_cpu.size(1) if args.max_features is None else min(args.max_features, base_z_cpu.size(1)) + stat_rows = [] + summary_rows = [] + + for feature_idx in range(feature_limit): + feature_active = int(active_count[feature_idx].item()) + keep = min(feature_active, top_k) + relevant_indices = top_indices[:keep, feature_idx] if keep > 0 else torch.empty(0, dtype=torch.long) + scale = float(intervention_scale[feature_idx].item()) + + feature_stat_rows, feature_summary = analyze_feature( + feature_idx=feature_idx, + token_tensor_cpu=token_tensor, + base_z_cpu=base_z_cpu, + feature_active_count=feature_active, + relevant_indices=relevant_indices, + scale=scale, + alphas=alphas, + sae=sae, + planner_model=planner_model, + lit_model=lit_model, + sae_block=args.sae_block, + past_cpu=past_cpu, + global_stat_stds=global_stat_stds, + dataset=dataset, + device=device, + ) + stat_rows.extend(feature_stat_rows) + summary_rows.append(feature_summary) + + if (feature_idx + 1) % 100 == 0 or feature_idx + 1 == feature_limit: + print(f"Processed {feature_idx + 1}/{feature_limit} features") + + summary_csv = output_dir / f"sae_control_summary_block_{args.sae_block}_{args.split}.csv" + stat_csv = output_dir / f"sae_control_stat_rows_block_{args.sae_block}_{args.split}.csv" + write_csv(summary_csv, summary_rows) + write_csv(stat_csv, stat_rows) + + print("") + print(f"Used token key {token_key}") + print_top_features(summary_rows, stat_rows, top_k=args.top_k) + print(f"Saved feature summary to {summary_csv}") + print(f"Saved per-stat rows to {stat_csv}") diff --git a/src/camera-based-e2e/analyze_sae_error.py b/src/camera-based-e2e/analyze_sae_error.py new file mode 100644 index 0000000..ee09c21 --- /dev/null +++ b/src/camera-based-e2e/analyze_sae_error.py @@ -0,0 +1,287 @@ +import argparse +import csv +import math +from pathlib import Path + +import torch +from torch.utils.data import DataLoader, TensorDataset + +from models.sae import SparseAutoencoder +from sae_utils import ( + build_sae_from_checkpoint, + default_analysis_dir, + default_device, + load_sae_bundle, + resolve_token_tensor, +) + + +INTENT_NAMES = { + 0: "UNKNOWN", + 1: "GO_STRAIGHT", + 2: "GO_LEFT", + 3: "GO_RIGHT", +} +def pearson_from_sums( + n: int, + sum_x: torch.Tensor, + sum_x2: torch.Tensor, + sum_y: float, + sum_y2: float, + sum_xy: torch.Tensor, +) -> tuple[torch.Tensor, float, float]: + n_float = float(n) + mean_x = sum_x / n_float + mean_y = sum_y / n_float + var_x = torch.clamp(sum_x2 / n_float - mean_x.square(), min=0.0) + var_y = max(sum_y2 / n_float - mean_y * mean_y, 0.0) + std_x = torch.sqrt(var_x) + std_y = math.sqrt(var_y) + cov_xy = sum_xy / n_float - mean_x * mean_y + + r = torch.zeros_like(mean_x) + if std_y > 0: + mask = std_x > 0 + r[mask] = cov_xy[mask] / (std_x[mask] * std_y) + return r, mean_y, std_y + + +def compute_ade_metrics(token_blob: dict) -> dict: + future = token_blob["future"].float() # (N, T, 2) + trajectory = token_blob["trajectory"].float() # (N, K*T*2) + scores = token_blob["scores"].float() # (N, K) + intent = token_blob["intent"].long() + + n = future.shape[0] + t = future.shape[1] + k = scores.shape[1] + + pred = trajectory.view(n, k, t, 2) + dist = torch.norm(pred - future[:, None, :, :], dim=-1) + ade_per_mode = dist.mean(dim=-1) + + selected_idx = scores.argmin(dim=1) + row_idx = torch.arange(n) + selected_ade = ade_per_mode[row_idx, selected_idx] + oracle_ade = ade_per_mode.min(dim=1).values + regret = selected_ade - oracle_ade + + return { + "selected_ade": selected_ade, + "oracle_ade": oracle_ade, + "regret": regret, + "intent": intent, + } + + +def compute_feature_stats( + model: SparseAutoencoder, + token_tensor: torch.Tensor, + metric_tensors: dict[str, torch.Tensor], + batch_size: int, + device: torch.device, +) -> dict: + metric_names = list(metric_tensors.keys()) + latent_dim = model.encoder.out_features + n_samples = token_tensor.shape[0] + + thresholds = {} + for name, values in metric_tensors.items(): + thresholds[name] = { + "low": torch.quantile(values, 0.10).item(), + "high": torch.quantile(values, 0.90).item(), + } + + dataset = TensorDataset( + token_tensor, + *[metric_tensors[name] for name in metric_names], + ) + loader = DataLoader(dataset, batch_size=batch_size, shuffle=False, num_workers=0) + + sum_x = torch.zeros(latent_dim, dtype=torch.float64) + sum_x2 = torch.zeros(latent_dim, dtype=torch.float64) + + metric_state = {} + for name in metric_names: + metric_state[name] = { + "sum_y": 0.0, + "sum_y2": 0.0, + "sum_xy": torch.zeros(latent_dim, dtype=torch.float64), + "high_count": 0, + "low_count": 0, + "high_sum_x": torch.zeros(latent_dim, dtype=torch.float64), + "low_sum_x": torch.zeros(latent_dim, dtype=torch.float64), + } + + model.eval() + with torch.no_grad(): + for batch in loader: + batch_x = batch[0].to(device, non_blocking=True) + z = model.encode(batch_x).cpu().to(torch.float64) + + sum_x += z.sum(dim=0) + sum_x2 += (z * z).sum(dim=0) + + for i, name in enumerate(metric_names, start=1): + y = batch[i].cpu().to(torch.float64) + state = metric_state[name] + state["sum_y"] += float(y.sum().item()) + state["sum_y2"] += float((y * y).sum().item()) + state["sum_xy"] += (z * y.unsqueeze(1)).sum(dim=0) + + high_mask = y >= thresholds[name]["high"] + low_mask = y <= thresholds[name]["low"] + if high_mask.any(): + state["high_count"] += int(high_mask.sum().item()) + state["high_sum_x"] += z[high_mask].sum(dim=0) + if low_mask.any(): + state["low_count"] += int(low_mask.sum().item()) + state["low_sum_x"] += z[low_mask].sum(dim=0) + + out = {"n_samples": n_samples, "thresholds": thresholds, "metrics": {}} + for name in metric_names: + state = metric_state[name] + r, mean_y, std_y = pearson_from_sums( + n=n_samples, + sum_x=sum_x, + sum_x2=sum_x2, + sum_y=state["sum_y"], + sum_y2=state["sum_y2"], + sum_xy=state["sum_xy"], + ) + high_mean = state["high_sum_x"] / max(state["high_count"], 1) + low_mean = state["low_sum_x"] / max(state["low_count"], 1) + out["metrics"][name] = { + "r": r, + "mean_y": mean_y, + "std_y": std_y, + "high_threshold": thresholds[name]["high"], + "low_threshold": thresholds[name]["low"], + "high_count": state["high_count"], + "low_count": state["low_count"], + "high_mean": high_mean, + "low_mean": low_mean, + "delta_high_low": high_mean - low_mean, + } + return out + + +def write_csv(stats: dict, output_csv: Path) -> None: + metric_names = list(stats["metrics"].keys()) + latent_dim = len(next(iter(stats["metrics"].values()))["r"]) + rows = [] + for feature_idx in range(latent_dim): + row = {"feature_idx": feature_idx} + best_metric = None + best_abs_r = -1.0 + for metric in metric_names: + metric_stats = stats["metrics"][metric] + r_val = float(metric_stats["r"][feature_idx].item()) + delta_val = float(metric_stats["delta_high_low"][feature_idx].item()) + row[f"r_{metric}"] = r_val + row[f"delta_high_low_{metric}"] = delta_val + if abs(r_val) > best_abs_r: + best_abs_r = abs(r_val) + best_metric = metric + row["best_abs_r"] = best_abs_r + row["best_metric"] = best_metric + rows.append(row) + + rows.sort(key=lambda item: item["r_selected_ade"], reverse=True) + fieldnames = list(rows[0].keys()) + output_csv.parent.mkdir(parents=True, exist_ok=True) + with output_csv.open("w", newline="") as f: + writer = csv.DictWriter(f, fieldnames=fieldnames) + writer.writeheader() + writer.writerows(rows) + + +def print_metric_summary(metric_name: str, metric_stats: dict, top_k: int) -> None: + r = metric_stats["r"] + delta = metric_stats["delta_high_low"] + top_pos = torch.argsort(r, descending=True)[:top_k].tolist() + top_neg = torch.argsort(r, descending=False)[:top_k].tolist() + + print( + f"{metric_name}: mean={metric_stats['mean_y']:.4f} std={metric_stats['std_y']:.4f} " + f"low10<={metric_stats['low_threshold']:.4f} high10>={metric_stats['high_threshold']:.4f}" + ) + print(f"Top {top_k} positive correlations:") + for rank, feature_idx in enumerate(top_pos, start=1): + print( + f" {rank}. feature={feature_idx} " + f"r={float(r[feature_idx].item()):+.4f} " + f"delta_high_low={float(delta[feature_idx].item()):+.4f}" + ) + print(f"Top {top_k} negative correlations:") + for rank, feature_idx in enumerate(top_neg, start=1): + print( + f" {rank}. feature={feature_idx} " + f"r={float(r[feature_idx].item()):+.4f} " + f"delta_high_low={float(delta[feature_idx].item()):+.4f}" + ) + print("") + + +def print_ade_by_intent(metrics: dict) -> None: + selected_ade = metrics["selected_ade"] + oracle_ade = metrics["oracle_ade"] + regret = metrics["regret"] + intent = metrics["intent"] + + print("ADE by intent:") + for intent_id in sorted(torch.unique(intent).tolist()): + mask = intent == intent_id + name = INTENT_NAMES.get(intent_id, str(intent_id)) + print( + f" {name}: n={int(mask.sum().item())} " + f"selected_ADE={float(selected_ade[mask].mean().item()):.4f} " + f"oracle_ADE={float(oracle_ade[mask].mean().item()):.4f} " + f"regret={float(regret[mask].mean().item()):.4f}" + ) + print("") + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument("--run_root", type=str, required=True) + parser.add_argument("--split", type=str, default="val", choices=["train", "val"]) + parser.add_argument("--output_dir", type=str, default=None) + parser.add_argument("--sae_block", type=int, default=3) + parser.add_argument("--batch_size", type=int, default=4096) + parser.add_argument("--device", type=str, default=default_device()) + parser.add_argument("--top_k", type=int, default=15) + args = parser.parse_args() + + run_root = Path(args.run_root) + output_dir = default_analysis_dir(run_root, args.sae_block, args.output_dir) + output_dir.mkdir(parents=True, exist_ok=True) + + device = torch.device(args.device) + bundle = load_sae_bundle(run_root, args.split, args.sae_block, map_location="cpu") + ckpt = bundle["ckpt"] + token_blob = bundle["token_blob"] + + model = build_sae_from_checkpoint(ckpt, bundle["legacy_norm"]) + model.to(device) + + token_tensor, token_key = resolve_token_tensor(token_blob, args.sae_block) + error_metrics = compute_ade_metrics(token_blob) + intent = error_metrics.pop("intent") + + stats = compute_feature_stats( + model=model, + token_tensor=token_tensor, + metric_tensors=error_metrics, + batch_size=args.batch_size, + device=device, + ) + + print_ade_by_intent({**error_metrics, "intent": intent}) + for metric_name, metric_stats in stats["metrics"].items(): + print_metric_summary(metric_name, metric_stats, top_k=args.top_k) + + output_csv = output_dir / f"sae_error_correlation_block_{args.sae_block}_{args.split}.csv" + write_csv(stats, output_csv) + print(f"Used token key {token_key}") + print(f"Saved CSV to {output_csv}") diff --git a/src/camera-based-e2e/analyze_sae_gated_ade.py b/src/camera-based-e2e/analyze_sae_gated_ade.py new file mode 100644 index 0000000..86df5f9 --- /dev/null +++ b/src/camera-based-e2e/analyze_sae_gated_ade.py @@ -0,0 +1,749 @@ +import argparse +import csv +from pathlib import Path + +import torch + +from extract_planner_tok import load_model +from sae_utils import ( + DEFAULT_SAE_BLOCK, + build_sae_from_checkpoint, + collate_dataset_indices, + dataset_from_token_blob, + default_device, + load_sae_bundle, + prepare_replay_context, + resolve_token_tensor, +) + + +PROXY_RULE_SPECS = ( + ("selected_score_down", "selected_score", -1.0), + ("score_margin_up", "score_margin", 1.0), + ("score_entropy_down", "score_entropy", -1.0), + ("proposal_spread_down", "proposal_spread", -1.0), + ("proposal_spread_up", "proposal_spread", 1.0), +) + + +def parse_int_list(text: str) -> list[int]: + return [int(part.strip()) for part in text.split(",") if part.strip()] + + +def parse_float_list(text: str) -> list[float]: + return [float(part.strip()) for part in text.split(",") if part.strip()] + + +def compute_selected_and_oracle_ade( + trajectory_flat: torch.Tensor, + scores: torch.Tensor, + future: torch.Tensor, + num_proposals: int, + horizon: int, +) -> tuple[torch.Tensor, torch.Tensor]: + batch = future.size(0) + traj = trajectory_flat.view(batch, num_proposals, horizon, 2) + dist = torch.norm(traj - future[:, None], dim=-1) + ade_per_mode = dist.mean(dim=-1) + row_idx = torch.arange(batch, device=future.device) + selected_idx = scores.argmin(dim=1) + selected_ade = ade_per_mode[row_idx, selected_idx] + oracle_ade = ade_per_mode.min(dim=1).values + return selected_ade, oracle_ade + + +def compute_output_stats( + trajectory_flat: torch.Tensor, + scores: torch.Tensor, + num_proposals: int, + horizon: int, +) -> dict[str, torch.Tensor]: + batch = scores.size(0) + traj = trajectory_flat.view(batch, num_proposals, horizon, 2) + row_idx = torch.arange(batch, device=scores.device) + selected_idx = scores.argmin(dim=1) + selected_score = scores[row_idx, selected_idx] + + if scores.size(1) > 1: + sorted_scores = scores.sort(dim=1).values + score_margin = sorted_scores[:, 1] - sorted_scores[:, 0] + else: + score_margin = torch.zeros(batch, device=scores.device, dtype=scores.dtype) + + score_probs = torch.softmax(-scores, dim=1) + score_entropy = -(score_probs * score_probs.clamp_min(1e-8).log()).sum(dim=1) + + traj_mean = traj.mean(dim=1, keepdim=True) + proposal_spread = torch.norm(traj - traj_mean, dim=-1).mean(dim=(1, 2)) + + return { + "selected_score": selected_score, + "score_margin": score_margin, + "score_entropy": score_entropy, + "proposal_spread": proposal_spread, + } + + +def compute_threshold_specs( + feature_act: torch.Tensor, + quantiles: list[float], +) -> list[dict]: + specs = [{"threshold_name": "always_on", "threshold_value": float("-inf")}] + specs.append({"threshold_name": "active_only", "threshold_value": 0.0}) + + positive = feature_act[feature_act > 0] + if positive.numel() == 0: + return specs + + seen = {spec["threshold_name"] for spec in specs} + for q in quantiles: + threshold_value = float(torch.quantile(positive, q).item()) + threshold_name = f"q{int(round(q * 100)):02d}_active" + if threshold_name in seen: + continue + specs.append( + { + "threshold_name": threshold_name, + "threshold_value": threshold_value, + } + ) + seen.add(threshold_name) + return specs + + +def compute_proxy_threshold_specs( + positive_values: torch.Tensor, + quantiles: list[float], +) -> list[dict]: + specs = [{"proxy_threshold_name": "gt_zero", "proxy_threshold_value": 0.0}] + if positive_values.numel() == 0: + return specs + + seen = {"gt_zero"} + for q in quantiles: + threshold_value = float(torch.quantile(positive_values, q).item()) + threshold_name = f"pos_q{int(round(q * 100)):02d}" + if threshold_name in seen: + continue + specs.append( + { + "proxy_threshold_name": threshold_name, + "proxy_threshold_value": threshold_value, + } + ) + seen.add(threshold_name) + return specs + + +def parse_feature_spec(text: str, latent_dim: int) -> list[int]: + raw = text.strip().lower() + if raw in {"all", "*"}: + return list(range(latent_dim)) + + if ":" in raw and "," not in raw: + parts = raw.split(":") + if len(parts) not in {2, 3}: + raise ValueError(f"Unsupported feature slice: {text}") + start = int(parts[0]) if parts[0] else 0 + stop = int(parts[1]) if parts[1] else latent_dim + step = int(parts[2]) if len(parts) == 3 and parts[2] else 1 + return list(range(start, stop, step)) + + return parse_int_list(text) + + +def clip_feature_list( + features: list[int], + latent_dim: int, + feature_start: int | None, + feature_end: int | None, +) -> list[int]: + clipped = [] + for feature_idx in features: + if feature_idx < 0 or feature_idx >= latent_dim: + continue + if feature_start is not None and feature_idx < feature_start: + continue + if feature_end is not None and feature_idx >= feature_end: + continue + clipped.append(feature_idx) + return clipped + + +def safe_mean(value_sum: float, count: int) -> float: + if count <= 0: + return 0.0 + return value_sum / count + + +def write_csv(path: Path, rows: list[dict]) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + if not rows: + return + with path.open("w", newline="") as f: + writer = csv.DictWriter(f, fieldnames=list(rows[0].keys())) + writer.writeheader() + writer.writerows(rows) + + +def derive_default_path(base_path: Path, suffix: str) -> Path: + return base_path.with_name(f"{base_path.stem}{suffix}{base_path.suffix}") + + +def build_proxy_row( + *, + feature_idx: int, + active_count: int, + alpha: float, + scale: float, + threshold_name: str, + threshold_value: float, + gate_rate: float, + intervened_count: int, + total_count: int, + delta_selected: torch.Tensor, + delta_oracle: torch.Tensor, + accept_mask: torch.Tensor, + proxy_rule: str, + proxy_threshold_name: str, + proxy_threshold_value: float, +) -> dict: + accepted_count = int(accept_mask.sum().item()) + if accepted_count > 0: + delta_selected_accept = delta_selected[accept_mask] + delta_oracle_accept = delta_oracle[accept_mask] + improved_selected_accept = int((delta_selected_accept < 0).sum().item()) + improved_oracle_accept = int((delta_oracle_accept < 0).sum().item()) + selected_sum = float(delta_selected_accept.sum().item()) + oracle_sum = float(delta_oracle_accept.sum().item()) + else: + improved_selected_accept = 0 + improved_oracle_accept = 0 + selected_sum = 0.0 + oracle_sum = 0.0 + + return { + "feature_idx": feature_idx, + "active_count": active_count, + "alpha": alpha, + "intervention_scale": scale, + "threshold_name": threshold_name, + "threshold_value": threshold_value, + "gate_rate": gate_rate, + "intervened_scene_count": intervened_count, + "proxy_rule": proxy_rule, + "proxy_threshold_name": proxy_threshold_name, + "proxy_threshold_value": proxy_threshold_value, + "accepted_scene_count": accepted_count, + "accepted_rate_all": accepted_count / total_count, + "accepted_rate_intervened": accepted_count / intervened_count if intervened_count > 0 else 0.0, + "mean_delta_selected_ade_accepted": safe_mean(selected_sum, accepted_count), + "frac_improved_selected_ade_accepted": ( + improved_selected_accept / accepted_count if accepted_count > 0 else 0.0 + ), + "mean_delta_oracle_ade_accepted": safe_mean(oracle_sum, accepted_count), + "frac_improved_oracle_ade_accepted": ( + improved_oracle_accept / accepted_count if accepted_count > 0 else 0.0 + ), + "mean_delta_selected_ade_if_applied": selected_sum / total_count, + "frac_improved_selected_ade_if_applied": improved_selected_accept / total_count, + } + + +def choose_best_proxy_row(proxy_rows: list[dict], min_accept_count: int) -> dict: + eligible = [row for row in proxy_rows if row["accepted_scene_count"] >= min_accept_count] + candidates = eligible if eligible else proxy_rows + return min( + candidates, + key=lambda row: ( + row["mean_delta_selected_ade_accepted"] if row["accepted_scene_count"] > 0 else float("inf"), + row["mean_delta_selected_ade_if_applied"], + -row["accepted_scene_count"], + ), + ) + + +def summarize_best_rows(rows: list[dict]) -> list[dict]: + by_feature: dict[int, list[dict]] = {} + for row in rows: + by_feature.setdefault(row["feature_idx"], []).append(row) + + summary_rows = [] + for feature_idx, feature_rows in by_feature.items(): + best_row = min( + feature_rows, + key=lambda row: ( + row["best_proxy_mean_delta_selected_ade_accepted"] + if row["best_proxy_accept_count"] > 0 + else float("inf"), + row["best_proxy_mean_delta_selected_ade_if_applied"], + row["mean_delta_selected_ade_intervened"], + ), + ) + summary_rows.append(best_row) + + summary_rows.sort( + key=lambda row: ( + row["best_proxy_mean_delta_selected_ade_accepted"] + if row["best_proxy_accept_count"] > 0 + else float("inf"), + row["best_proxy_mean_delta_selected_ade_if_applied"], + ) + ) + return summary_rows + + +def evaluate_setting( + *, + feature_idx: int, + active_count: int, + alpha: float, + scale: float, + threshold_name: str, + threshold_value: float, + intervene_idx: torch.Tensor, + token_tensor_cpu: torch.Tensor, + z_all_cpu: torch.Tensor, + past_cpu: torch.Tensor, + future_cpu: torch.Tensor, + baseline_selected: torch.Tensor, + baseline_oracle: torch.Tensor, + baseline_stats_cpu: dict[str, torch.Tensor], + hard_mask_all: torch.Tensor, + sae, + planner_model, + lit_model, + sae_block: int, + dataset, + batch_size: int, + device: torch.device, + proxy_metric_quantiles: list[float], + min_proxy_accept_count: int, +) -> tuple[dict, list[dict]]: + total_count = baseline_selected.numel() + intervene_idx = intervene_idx.to(torch.long) + intervened_count = int(intervene_idx.numel()) + gate_rate = intervened_count / total_count + + if intervened_count == 0: + proxy_rows = [ + build_proxy_row( + feature_idx=feature_idx, + active_count=active_count, + alpha=alpha, + scale=scale, + threshold_name=threshold_name, + threshold_value=threshold_value, + gate_rate=gate_rate, + intervened_count=0, + total_count=total_count, + delta_selected=torch.zeros(0), + delta_oracle=torch.zeros(0), + accept_mask=torch.zeros(0, dtype=torch.bool), + proxy_rule="keep_all_intervened", + proxy_threshold_name="all", + proxy_threshold_value=float("-inf"), + ) + ] + best_proxy = proxy_rows[0] + row = { + "feature_idx": feature_idx, + "active_count": active_count, + "alpha": alpha, + "intervention_scale": scale, + "threshold_name": threshold_name, + "threshold_value": threshold_value, + "gate_rate": gate_rate, + "mean_delta_selected_ade": 0.0, + "frac_improved_selected_ade": 0.0, + "mean_delta_oracle_ade": 0.0, + "frac_improved_oracle_ade": 0.0, + "intervened_scene_count": 0, + "mean_delta_selected_ade_intervened": 0.0, + "frac_improved_selected_ade_intervened": 0.0, + "intervened_hard_scene_count": 0, + "mean_delta_selected_ade_hard_intervened": 0.0, + "frac_improved_selected_ade_hard_intervened": 0.0, + "best_proxy_rule": best_proxy["proxy_rule"], + "best_proxy_threshold_name": best_proxy["proxy_threshold_name"], + "best_proxy_threshold_value": best_proxy["proxy_threshold_value"], + "best_proxy_accept_count": best_proxy["accepted_scene_count"], + "best_proxy_accept_rate_all": best_proxy["accepted_rate_all"], + "best_proxy_accept_rate_intervened": best_proxy["accepted_rate_intervened"], + "best_proxy_mean_delta_selected_ade_accepted": best_proxy["mean_delta_selected_ade_accepted"], + "best_proxy_frac_improved_selected_ade_accepted": best_proxy["frac_improved_selected_ade_accepted"], + "best_proxy_mean_delta_selected_ade_if_applied": best_proxy["mean_delta_selected_ade_if_applied"], + "best_proxy_frac_improved_selected_ade_if_applied": best_proxy["frac_improved_selected_ade_if_applied"], + } + return row, proxy_rows + + sum_delta_selected = 0.0 + sum_delta_oracle = 0.0 + improved_selected = 0 + improved_oracle = 0 + improved_selected_hard = 0 + intervened_hard_count = 0 + sum_delta_selected_hard = 0.0 + + delta_selected_chunks = [] + delta_oracle_chunks = [] + directional_value_chunks = {name: [] for name, _, _ in PROXY_RULE_SPECS} + + sae.eval() + planner_model.eval() + with torch.no_grad(): + for start in range(0, intervened_count, batch_size): + batch_indices = intervene_idx[start : start + batch_size] + batch_future = future_cpu[batch_indices].to(device) + batch_baseline_selected = baseline_selected[batch_indices].to(device) + batch_baseline_oracle = baseline_oracle[batch_indices].to(device) + batch_hard_mask = hard_mask_all[batch_indices] + + batch_x = token_tensor_cpu[batch_indices].to(device) + z_batch = z_all_cpu[batch_indices].to(device) + act_batch = z_batch[:, feature_idx] + z_batch[:, feature_idx] = (act_batch + alpha * scale).clamp_min(0.0) + + recon_query = sae.decode_to_input(z_batch, reference_x=batch_x) + if sae_block == DEFAULT_SAE_BLOCK: + batch_past = past_cpu[batch_indices].to(device) + out = planner_model.forward_from_planner_query_tok(recon_query, batch_past) + else: + batch = collate_dataset_indices(dataset, batch_indices) + replay_context = prepare_replay_context(planner_model, lit_model, batch, device=device) + out = planner_model.forward_from_block_query_tok( + recon_query, + replay_context["past"], + replay_context["tokens"], + start_block=sae_block, + ) + + selected_ade, oracle_ade = compute_selected_and_oracle_ade( + trajectory_flat=out["trajectory"], + scores=out["scores"], + future=batch_future, + num_proposals=planner_model.n_proposals, + horizon=planner_model.horizon, + ) + out_stats = compute_output_stats( + trajectory_flat=out["trajectory"], + scores=out["scores"], + num_proposals=planner_model.n_proposals, + horizon=planner_model.horizon, + ) + + delta_selected = (selected_ade - batch_baseline_selected).detach().cpu() + delta_oracle = (oracle_ade - batch_baseline_oracle).detach().cpu() + + sum_delta_selected += float(delta_selected.sum().item()) + sum_delta_oracle += float(delta_oracle.sum().item()) + improved_selected += int((delta_selected < 0).sum().item()) + improved_oracle += int((delta_oracle < 0).sum().item()) + + hard_mask_cpu = batch_hard_mask.cpu() + intervened_hard_count += int(hard_mask_cpu.sum().item()) + if hard_mask_cpu.any(): + hard_delta = delta_selected[hard_mask_cpu] + sum_delta_selected_hard += float(hard_delta.sum().item()) + improved_selected_hard += int((hard_delta < 0).sum().item()) + + delta_selected_chunks.append(delta_selected) + delta_oracle_chunks.append(delta_oracle) + + for proxy_name, stat_name, sign in PROXY_RULE_SPECS: + baseline_stat = baseline_stats_cpu[stat_name][batch_indices].to(device) + directional_values = sign * (out_stats[stat_name] - baseline_stat) + directional_value_chunks[proxy_name].append(directional_values.detach().cpu()) + + delta_selected_all = torch.cat(delta_selected_chunks, dim=0) + delta_oracle_all = torch.cat(delta_oracle_chunks, dim=0) + + proxy_rows = [ + build_proxy_row( + feature_idx=feature_idx, + active_count=active_count, + alpha=alpha, + scale=scale, + threshold_name=threshold_name, + threshold_value=threshold_value, + gate_rate=gate_rate, + intervened_count=intervened_count, + total_count=total_count, + delta_selected=delta_selected_all, + delta_oracle=delta_oracle_all, + accept_mask=torch.ones(intervened_count, dtype=torch.bool), + proxy_rule="keep_all_intervened", + proxy_threshold_name="all", + proxy_threshold_value=float("-inf"), + ) + ] + + for proxy_name, _, _ in PROXY_RULE_SPECS: + directional_values = torch.cat(directional_value_chunks[proxy_name], dim=0) + positive_values = directional_values[directional_values > 0] + for spec in compute_proxy_threshold_specs(positive_values, proxy_metric_quantiles): + accept_mask = directional_values > spec["proxy_threshold_value"] + proxy_rows.append( + build_proxy_row( + feature_idx=feature_idx, + active_count=active_count, + alpha=alpha, + scale=scale, + threshold_name=threshold_name, + threshold_value=threshold_value, + gate_rate=gate_rate, + intervened_count=intervened_count, + total_count=total_count, + delta_selected=delta_selected_all, + delta_oracle=delta_oracle_all, + accept_mask=accept_mask, + proxy_rule=proxy_name, + proxy_threshold_name=spec["proxy_threshold_name"], + proxy_threshold_value=spec["proxy_threshold_value"], + ) + ) + + best_proxy = choose_best_proxy_row(proxy_rows, min_accept_count=min_proxy_accept_count) + + row = { + "feature_idx": feature_idx, + "active_count": active_count, + "alpha": alpha, + "intervention_scale": scale, + "threshold_name": threshold_name, + "threshold_value": threshold_value, + "gate_rate": gate_rate, + "mean_delta_selected_ade": sum_delta_selected / total_count, + "frac_improved_selected_ade": improved_selected / total_count, + "mean_delta_oracle_ade": sum_delta_oracle / total_count, + "frac_improved_oracle_ade": improved_oracle / total_count, + "intervened_scene_count": intervened_count, + "mean_delta_selected_ade_intervened": safe_mean(sum_delta_selected, intervened_count), + "frac_improved_selected_ade_intervened": ( + improved_selected / intervened_count if intervened_count > 0 else 0.0 + ), + "intervened_hard_scene_count": intervened_hard_count, + "mean_delta_selected_ade_hard_intervened": safe_mean(sum_delta_selected_hard, intervened_hard_count), + "frac_improved_selected_ade_hard_intervened": ( + improved_selected_hard / intervened_hard_count if intervened_hard_count > 0 else 0.0 + ), + "best_proxy_rule": best_proxy["proxy_rule"], + "best_proxy_threshold_name": best_proxy["proxy_threshold_name"], + "best_proxy_threshold_value": best_proxy["proxy_threshold_value"], + "best_proxy_accept_count": best_proxy["accepted_scene_count"], + "best_proxy_accept_rate_all": best_proxy["accepted_rate_all"], + "best_proxy_accept_rate_intervened": best_proxy["accepted_rate_intervened"], + "best_proxy_mean_delta_selected_ade_accepted": best_proxy["mean_delta_selected_ade_accepted"], + "best_proxy_frac_improved_selected_ade_accepted": best_proxy["frac_improved_selected_ade_accepted"], + "best_proxy_mean_delta_selected_ade_if_applied": best_proxy["mean_delta_selected_ade_if_applied"], + "best_proxy_frac_improved_selected_ade_if_applied": best_proxy["frac_improved_selected_ade_if_applied"], + } + return row, proxy_rows + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument("--run_root", type=str, required=True) + parser.add_argument("--planner_checkpoint", type=str, required=True) + parser.add_argument("--split", type=str, default="val", choices=["train", "val"]) + parser.add_argument("--features", type=str, default="all") + parser.add_argument("--sae_block", type=int, default=3) + parser.add_argument("--data_dir", type=str, default=None) + parser.add_argument("--index_file", type=str, default=None) + parser.add_argument("--feature_start", type=int, default=None) + parser.add_argument("--feature_end", type=int, default=None) + parser.add_argument("--alphas", type=str, default="1.0,2.0") + parser.add_argument("--threshold_quantiles", type=str, default="0.5,0.75,0.9,0.95") + parser.add_argument("--proxy_metric_quantiles", type=str, default="0.75,0.9") + parser.add_argument("--min_proxy_accept_count", type=int, default=32) + parser.add_argument("--batch_size", type=int, default=2048) + parser.add_argument("--device", type=str, default=default_device()) + parser.add_argument("--output_csv", type=str, required=True) + parser.add_argument("--output_proxy_csv", type=str, default=None) + parser.add_argument("--output_best_csv", type=str, default=None) + args = parser.parse_args() + + device = torch.device(args.device) + run_root = Path(args.run_root) + output_csv = Path(args.output_csv) + output_proxy_csv = Path(args.output_proxy_csv) if args.output_proxy_csv else None + output_best_csv = ( + Path(args.output_best_csv) + if args.output_best_csv + else derive_default_path(output_csv, "_best_by_feature") + ) + + alphas = parse_float_list(args.alphas) + threshold_quantiles = parse_float_list(args.threshold_quantiles) + proxy_metric_quantiles = parse_float_list(args.proxy_metric_quantiles) + + bundle = load_sae_bundle(run_root, args.split, args.sae_block, map_location="cpu") + sae_ckpt = bundle["ckpt"] + token_blob = bundle["token_blob"] + token_tensor, token_key = resolve_token_tensor(token_blob, args.sae_block) + + sae = build_sae_from_checkpoint(sae_ckpt, bundle["legacy_norm"]) + sae.to(device) + sae.eval() + + planner_model, lit_model = load_model(args.planner_checkpoint, device=device) + planner_model.eval() + + latent_dim = sae_ckpt["latent_dim"] + features = clip_feature_list( + parse_feature_spec(args.features, latent_dim=latent_dim), + latent_dim=latent_dim, + feature_start=args.feature_start, + feature_end=args.feature_end, + ) + + past_cpu = token_blob["past"].float() + future_cpu = token_blob["future"].float() + scores_cpu = token_blob["scores"].float() + trajectory_cpu = token_blob["trajectory"].float() + + baseline_selected, baseline_oracle = compute_selected_and_oracle_ade( + trajectory_flat=trajectory_cpu, + scores=scores_cpu, + future=future_cpu, + num_proposals=planner_model.n_proposals, + horizon=planner_model.horizon, + ) + baseline_stats_cpu = compute_output_stats( + trajectory_flat=trajectory_cpu, + scores=scores_cpu, + num_proposals=planner_model.n_proposals, + horizon=planner_model.horizon, + ) + hard_threshold = torch.quantile(baseline_selected, 0.75) + hard_mask_all = baseline_selected >= hard_threshold + + z_all_cpu = [] + with torch.no_grad(): + for start in range(0, len(token_tensor), args.batch_size): + batch_x = token_tensor[start : start + args.batch_size].to(device) + z_all_cpu.append(sae.encode(batch_x).cpu()) + z_all_cpu = torch.cat(z_all_cpu, dim=0) + + active_mask = z_all_cpu > 0 + active_count = active_mask.sum(dim=0) + active_sum = z_all_cpu.sum(dim=0) + active_sum_sq = (z_all_cpu * z_all_cpu).sum(dim=0) + active_mean = active_sum / active_count.clamp_min(1) + active_var = active_sum_sq / active_count.clamp_min(1) - active_mean.square() + active_std = torch.sqrt(active_var.clamp_min(0.0)) + scales = torch.maximum(active_std, 0.25 * active_mean).clamp_min(0.05) + + dataset = None + if args.sae_block != DEFAULT_SAE_BLOCK: + dataset = dataset_from_token_blob( + token_blob, + data_dir=args.data_dir, + index_file=args.index_file, + ) + + all_indices = torch.arange(len(token_tensor)) + rows = [] + proxy_rows = [] if output_proxy_csv is not None else None + + print( + f"Running gated ADE analysis for {len(features)} features " + f"(feature_start={args.feature_start}, feature_end={args.feature_end}, token_key={token_key})", + flush=True, + ) + + for order_idx, feature_idx in enumerate(features, start=1): + feature_act = z_all_cpu[:, feature_idx] + threshold_specs = compute_threshold_specs(feature_act, threshold_quantiles) + active_feature_count = int(active_count[feature_idx].item()) + scale = float(scales[feature_idx].item()) + print( + f"[{order_idx}/{len(features)}] feature={feature_idx} active={active_feature_count} scale={scale:.4f}", + flush=True, + ) + + for alpha in alphas: + for spec in threshold_specs: + if spec["threshold_name"] == "always_on": + intervene_idx = all_indices + else: + intervene_idx = torch.nonzero( + feature_act > spec["threshold_value"], + as_tuple=False, + ).squeeze(1) + + row, setting_proxy_rows = evaluate_setting( + feature_idx=feature_idx, + active_count=active_feature_count, + alpha=alpha, + scale=scale, + threshold_name=spec["threshold_name"], + threshold_value=spec["threshold_value"], + intervene_idx=intervene_idx, + token_tensor_cpu=token_tensor, + z_all_cpu=z_all_cpu, + past_cpu=past_cpu, + future_cpu=future_cpu, + baseline_selected=baseline_selected, + baseline_oracle=baseline_oracle, + baseline_stats_cpu=baseline_stats_cpu, + hard_mask_all=hard_mask_all, + sae=sae, + planner_model=planner_model, + lit_model=lit_model, + sae_block=args.sae_block, + dataset=dataset, + batch_size=args.batch_size, + device=device, + proxy_metric_quantiles=proxy_metric_quantiles, + min_proxy_accept_count=args.min_proxy_accept_count, + ) + rows.append(row) + if proxy_rows is not None: + proxy_rows.extend(setting_proxy_rows) + + print( + { + "feature_idx": feature_idx, + "alpha": alpha, + "threshold_name": spec["threshold_name"], + "gate_rate": row["gate_rate"], + "mean_delta_selected_ade_intervened": row["mean_delta_selected_ade_intervened"], + "best_proxy_rule": row["best_proxy_rule"], + "best_proxy_mean_delta_selected_ade_accepted": row["best_proxy_mean_delta_selected_ade_accepted"], + "best_proxy_accept_count": row["best_proxy_accept_count"], + }, + flush=True, + ) + + rows.sort( + key=lambda row: ( + row["best_proxy_mean_delta_selected_ade_accepted"] + if row["best_proxy_accept_count"] > 0 + else float("inf"), + row["best_proxy_mean_delta_selected_ade_if_applied"], + row["mean_delta_selected_ade_intervened"], + ) + ) + write_csv(output_csv, rows) + + best_rows = summarize_best_rows(rows) + write_csv(output_best_csv, best_rows) + + if output_proxy_csv is not None and proxy_rows is not None: + proxy_rows.sort( + key=lambda row: ( + row["mean_delta_selected_ade_accepted"] + if row["accepted_scene_count"] > 0 + else float("inf"), + row["mean_delta_selected_ade_if_applied"], + -row["accepted_scene_count"], + ) + ) + write_csv(output_proxy_csv, proxy_rows) + + print(f"Used token key {token_key}", flush=True) + print(f"Saved detailed rows to {output_csv}", flush=True) + print(f"Saved best-per-feature rows to {output_best_csv}", flush=True) + if output_proxy_csv is not None: + print(f"Saved proxy-rule rows to {output_proxy_csv}", flush=True) diff --git a/src/camera-based-e2e/analyze_sae_gated_ade_old.py b/src/camera-based-e2e/analyze_sae_gated_ade_old.py new file mode 100644 index 0000000..ac1bb7d --- /dev/null +++ b/src/camera-based-e2e/analyze_sae_gated_ade_old.py @@ -0,0 +1,275 @@ +import argparse +import csv +from pathlib import Path + +import torch + +from extract_planner_tok import load_model +from models.sae import SparseAutoencoder + + +def parse_int_list(text: str) -> list[int]: + return [int(part.strip()) for part in text.split(",") if part.strip()] + + +def parse_float_list(text: str) -> list[float]: + return [float(part.strip()) for part in text.split(",") if part.strip()] + + +def normalize(x: torch.Tensor, mean: torch.Tensor, std: torch.Tensor) -> torch.Tensor: + return (x - mean) / std + + +def denormalize(x: torch.Tensor, mean: torch.Tensor, std: torch.Tensor) -> torch.Tensor: + return x * std + mean + + +def compute_selected_and_oracle_ade( + trajectory_flat: torch.Tensor, + scores: torch.Tensor, + future: torch.Tensor, + num_proposals: int, + horizon: int, +) -> tuple[torch.Tensor, torch.Tensor]: + batch = future.size(0) + traj = trajectory_flat.view(batch, num_proposals, horizon, 2) + dist = torch.norm(traj - future[:, None], dim=-1) + ade_per_mode = dist.mean(dim=-1) + row_idx = torch.arange(batch, device=future.device) + selected_idx = scores.argmin(dim=1) + selected_ade = ade_per_mode[row_idx, selected_idx] + oracle_ade = ade_per_mode.min(dim=1).values + return selected_ade, oracle_ade + + +def compute_threshold_specs( + feature_act: torch.Tensor, + quantiles: list[float], +) -> list[dict]: + specs = [{"threshold_name": "always_on", "threshold_value": float("-inf")}] + specs.append({"threshold_name": "active_only", "threshold_value": 0.0}) + + positive = feature_act[feature_act > 0] + if positive.numel() == 0: + return specs + + seen = {spec["threshold_name"] for spec in specs} + for q in quantiles: + threshold_value = float(torch.quantile(positive, q).item()) + threshold_name = f"q{int(round(q * 100)):02d}_active" + if threshold_name in seen: + continue + specs.append( + { + "threshold_name": threshold_name, + "threshold_value": threshold_value, + } + ) + seen.add(threshold_name) + return specs + + +def write_csv(path: Path, rows: list[dict]) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + if not rows: + return + with path.open("w", newline="") as f: + writer = csv.DictWriter(f, fieldnames=list(rows[0].keys())) + writer.writeheader() + writer.writerows(rows) + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument("--run_root", type=str, required=True) + parser.add_argument("--planner_checkpoint", type=str, required=True) + parser.add_argument("--split", type=str, default="val", choices=["train", "val"]) + parser.add_argument("--features", type=str, required=True) + parser.add_argument("--alphas", type=str, default="1.0,2.0") + parser.add_argument("--threshold_quantiles", type=str, default="0.5,0.75,0.9,0.95,0.99") + parser.add_argument("--batch_size", type=int, default=2048) + parser.add_argument("--device", type=str, default="cuda" if torch.cuda.is_available() else "cpu") + parser.add_argument("--output_csv", type=str, required=True) + args = parser.parse_args() + + device = torch.device(args.device) + run_root = Path(args.run_root) + features = parse_int_list(args.features) + alphas = parse_float_list(args.alphas) + threshold_quantiles = parse_float_list(args.threshold_quantiles) + + sae_ckpt = torch.load(run_root / "model" / "sae_checkpoint.pt", map_location="cpu") + norm = torch.load(run_root / "model" / "sae_normalization.pt", map_location="cpu") + token_blob = torch.load(run_root / "tokens" / f"planner_tokens_{args.split}.pt", map_location="cpu") + + sae = SparseAutoencoder( + input_dim=sae_ckpt["input_dim"], + latent_dim=sae_ckpt["latent_dim"], + ) + sae.load_state_dict(sae_ckpt["state_dict"]) + sae.to(device) + sae.eval() + + planner_model, _ = load_model(args.planner_checkpoint, device=device) + planner_model.eval() + + mean = norm["mean"].to(device) + std = norm["std"].to(device) + + token_tensor = token_blob["planner_query_tok"].float() + past_cpu = token_blob["past"].float() + future_cpu = token_blob["future"].float() + scores_cpu = token_blob["scores"].float() + trajectory_cpu = token_blob["trajectory"].float() + + # Compute baseline ADE using the saved planner outputs. + baseline_selected, baseline_oracle = compute_selected_and_oracle_ade( + trajectory_flat=trajectory_cpu, + scores=scores_cpu, + future=future_cpu, + num_proposals=planner_model.n_proposals, + horizon=planner_model.horizon, + ) + hard_threshold = torch.quantile(baseline_selected, 0.75) + hard_mask_all = baseline_selected >= hard_threshold + + # Encode once for all validation scenes and keep on CPU for thresholding. + latent_chunks = [] + with torch.no_grad(): + for start in range(0, len(token_tensor), args.batch_size): + batch_x = token_tensor[start : start + args.batch_size].to(device) + latent_chunks.append(sae.encode(normalize(batch_x, mean, std)).cpu()) + z_all_cpu = torch.cat(latent_chunks, dim=0) + + active_mask = z_all_cpu > 0 + active_count = active_mask.sum(dim=0) + active_sum = z_all_cpu.sum(dim=0) + active_sum_sq = (z_all_cpu * z_all_cpu).sum(dim=0) + active_mean = active_sum / active_count.clamp_min(1) + active_var = active_sum_sq / active_count.clamp_min(1) - active_mean.square() + active_std = torch.sqrt(active_var.clamp_min(0.0)) + scales = torch.maximum(active_std, 0.25 * active_mean).clamp_min(0.05) + + rows = [] + for feature_idx in features: + feature_act = z_all_cpu[:, feature_idx] + threshold_specs = compute_threshold_specs(feature_act, threshold_quantiles) + active_feature_count = int(active_count[feature_idx].item()) + scale = float(scales[feature_idx].item()) + + for alpha in alphas: + accum = {} + for spec in threshold_specs: + accum[(alpha, spec["threshold_name"])] = { + "sum_delta_selected": 0.0, + "sum_delta_oracle": 0.0, + "count_improved_selected": 0, + "count_improved_oracle": 0, + "count_monotone_selected_proxy": 0, + "count_intervened": 0, + "sum_delta_selected_intervened": 0.0, + "count_improved_selected_intervened": 0, + "count_intervened_hard": 0, + "sum_delta_selected_hard": 0.0, + "count_improved_selected_hard": 0, + } + + with torch.no_grad(): + for start in range(0, len(token_tensor), args.batch_size): + end = min(start + args.batch_size, len(token_tensor)) + batch_tokens = token_tensor[start:end].to(device) + batch_past = past_cpu[start:end].to(device) + batch_future = future_cpu[start:end].to(device) + batch_baseline_selected = baseline_selected[start:end].to(device) + batch_baseline_oracle = baseline_oracle[start:end].to(device) + batch_hard_mask = hard_mask_all[start:end].to(device) + z_batch = z_all_cpu[start:end].to(device) + act_batch = z_batch[:, feature_idx] + + for spec in threshold_specs: + if spec["threshold_name"] == "always_on": + intervene_mask = torch.ones_like(act_batch, dtype=torch.bool) + else: + intervene_mask = act_batch > spec["threshold_value"] + + z_mod = z_batch.clone() + if intervene_mask.any(): + z_mod[intervene_mask, feature_idx] = ( + act_batch[intervene_mask] + alpha * scale + ).clamp_min(0.0) + + recon_norm = sae.decode(z_mod) + recon_query = denormalize(recon_norm, mean, std) + out = planner_model.forward_from_planner_query_tok(recon_query, batch_past) + selected_ade, oracle_ade = compute_selected_and_oracle_ade( + trajectory_flat=out["trajectory"], + scores=out["scores"], + future=batch_future, + num_proposals=planner_model.n_proposals, + horizon=planner_model.horizon, + ) + + delta_selected = selected_ade - batch_baseline_selected + delta_oracle = oracle_ade - batch_baseline_oracle + key = (alpha, spec["threshold_name"]) + state = accum[key] + state["sum_delta_selected"] += float(delta_selected.sum().item()) + state["sum_delta_oracle"] += float(delta_oracle.sum().item()) + state["count_improved_selected"] += int((delta_selected < 0).sum().item()) + state["count_improved_oracle"] += int((delta_oracle < 0).sum().item()) + state["count_intervened"] += int(intervene_mask.sum().item()) + state["sum_delta_selected_intervened"] += float(delta_selected[intervene_mask].sum().item()) + state["count_improved_selected_intervened"] += int((delta_selected[intervene_mask] < 0).sum().item()) + hard_and_intervened = batch_hard_mask & intervene_mask + state["count_intervened_hard"] += int(hard_and_intervened.sum().item()) + state["sum_delta_selected_hard"] += float(delta_selected[hard_and_intervened].sum().item()) + state["count_improved_selected_hard"] += int((delta_selected[hard_and_intervened] < 0).sum().item()) + + total_count = len(token_tensor) + for spec in threshold_specs: + key = (alpha, spec["threshold_name"]) + state = accum[key] + intervened_count = state["count_intervened"] + intervened_hard_count = state["count_intervened_hard"] + rows.append( + { + "feature_idx": feature_idx, + "active_count": active_feature_count, + "alpha": alpha, + "threshold_name": spec["threshold_name"], + "threshold_value": spec["threshold_value"], + "gate_rate": intervened_count / total_count, + "mean_delta_selected_ade": state["sum_delta_selected"] / total_count, + "frac_improved_selected_ade": state["count_improved_selected"] / total_count, + "mean_delta_oracle_ade": state["sum_delta_oracle"] / total_count, + "frac_improved_oracle_ade": state["count_improved_oracle"] / total_count, + "intervened_scene_count": intervened_count, + "mean_delta_selected_ade_intervened": ( + state["sum_delta_selected_intervened"] / intervened_count + if intervened_count > 0 + else 0.0 + ), + "frac_improved_selected_ade_intervened": ( + state["count_improved_selected_intervened"] / intervened_count + if intervened_count > 0 + else 0.0 + ), + "intervened_hard_scene_count": intervened_hard_count, + "mean_delta_selected_ade_hard_intervened": ( + state["sum_delta_selected_hard"] / intervened_hard_count + if intervened_hard_count > 0 + else 0.0 + ), + "frac_improved_selected_ade_hard_intervened": ( + state["count_improved_selected_hard"] / intervened_hard_count + if intervened_hard_count > 0 + else 0.0 + ), + } + ) + + print(rows[-1]) + + rows.sort(key=lambda row: row["mean_delta_selected_ade"]) + write_csv(Path(args.output_csv), rows) + print(f"Saved CSV to {args.output_csv}") diff --git a/src/camera-based-e2e/analyze_sae_intent.py b/src/camera-based-e2e/analyze_sae_intent.py new file mode 100644 index 0000000..dae16f5 --- /dev/null +++ b/src/camera-based-e2e/analyze_sae_intent.py @@ -0,0 +1,268 @@ +import argparse +import csv +import math +from pathlib import Path + +import torch +from torch.utils.data import DataLoader, TensorDataset + +from models.sae import SparseAutoencoder +from sae_utils import ( + build_sae_from_checkpoint, + default_analysis_dir, + default_device, + load_sae_bundle, + resolve_token_tensor, +) + + +INTENT_NAMES = { + 0: "UNKNOWN", + 1: "GO_STRAIGHT", + 2: "GO_LEFT", + 3: "GO_RIGHT", +} + + +def safe_div(num: torch.Tensor, den: torch.Tensor) -> torch.Tensor: + out = torch.zeros_like(num) + mask = den != 0 + out[mask] = num[mask] / den[mask] + return out +def compute_stats( + model: SparseAutoencoder, + token_tensor: torch.Tensor, + intent_tensor: torch.Tensor, + batch_size: int, + device: torch.device, +) -> dict: + dataset = TensorDataset(token_tensor, intent_tensor) + loader = DataLoader(dataset, batch_size=batch_size, shuffle=False, num_workers=0) + + latent_dim = model.encoder.out_features + all_intents = sorted(torch.unique(intent_tensor).tolist()) + + total_count = 0 + total_sum = torch.zeros(latent_dim, dtype=torch.float64) + total_sum_sq = torch.zeros(latent_dim, dtype=torch.float64) + total_active = torch.zeros(latent_dim, dtype=torch.float64) + + class_counts = {intent: 0 for intent in all_intents} + class_sums = {intent: torch.zeros(latent_dim, dtype=torch.float64) for intent in all_intents} + class_active = {intent: torch.zeros(latent_dim, dtype=torch.float64) for intent in all_intents} + + model.eval() + with torch.no_grad(): + for batch_x, batch_intent in loader: + batch_x = batch_x.to(device, non_blocking=True) + batch_intent = batch_intent.to(device, non_blocking=True) + + z = model.encode(batch_x) + z_cpu = z.cpu().to(torch.float64) + intent_cpu = batch_intent.cpu() + + total_count += z_cpu.shape[0] + total_sum += z_cpu.sum(dim=0) + total_sum_sq += (z_cpu * z_cpu).sum(dim=0) + total_active += (z_cpu > 0).sum(dim=0) + + for intent in all_intents: + mask = intent_cpu == intent + n = int(mask.sum().item()) + if n == 0: + continue + selected = z_cpu[mask] + class_counts[intent] += n + class_sums[intent] += selected.sum(dim=0) + class_active[intent] += (selected > 0).sum(dim=0) + + total_n = torch.tensor(float(total_count), dtype=torch.float64) + mean_all = total_sum / total_n + var_all = safe_div(total_sum_sq, total_n) - mean_all.square() + var_all = torch.clamp(var_all, min=0.0) + std_all = torch.sqrt(var_all) + active_rate_all = total_active / total_n + + ss_total = total_sum_sq - total_n * mean_all.square() + ss_between = torch.zeros_like(ss_total) + mean_by_intent = {} + active_rate_by_intent = {} + point_biserial_r = {} + + for intent in all_intents: + class_n = float(class_counts[intent]) + class_sum = class_sums[intent] + class_active_sum = class_active[intent] + class_n_tensor = torch.tensor(class_n, dtype=torch.float64) + other_n = float(total_count - class_counts[intent]) + + mean_intent = safe_div(class_sum, class_n_tensor) + mean_by_intent[intent] = mean_intent + active_rate_by_intent[intent] = safe_div(class_active_sum, class_n_tensor) + + ss_between += class_n_tensor * (mean_intent - mean_all).square() + + if class_n == 0 or other_n == 0: + point_biserial_r[intent] = torch.zeros(latent_dim, dtype=torch.float64) + continue + + other_mean = safe_div(total_sum - class_sum, torch.tensor(other_n, dtype=torch.float64)) + p = class_n / total_count + q = 1.0 - p + scale = math.sqrt(p * q) + r = torch.zeros(latent_dim, dtype=torch.float64) + denom_mask = std_all > 0 + r[denom_mask] = ((mean_intent[denom_mask] - other_mean[denom_mask]) / std_all[denom_mask]) * scale + point_biserial_r[intent] = r + + eta_sq = torch.zeros_like(ss_total) + valid_total = ss_total > 0 + eta_sq[valid_total] = ss_between[valid_total] / ss_total[valid_total] + + return { + "all_intents": all_intents, + "total_count": total_count, + "mean_all": mean_all, + "std_all": std_all, + "active_rate_all": active_rate_all, + "mean_by_intent": mean_by_intent, + "active_rate_by_intent": active_rate_by_intent, + "point_biserial_r": point_biserial_r, + "eta_sq": eta_sq, + "class_counts": class_counts, + } + + +def write_csv(stats: dict, output_csv: Path) -> None: + all_intents = stats["all_intents"] + eta_sq = stats["eta_sq"] + + rows = [] + for feature_idx in range(len(eta_sq)): + row = { + "feature_idx": feature_idx, + "eta_sq": float(eta_sq[feature_idx].item()), + "mean_all": float(stats["mean_all"][feature_idx].item()), + "std_all": float(stats["std_all"][feature_idx].item()), + "active_rate_all": float(stats["active_rate_all"][feature_idx].item()), + } + + best_intent = None + best_abs_r = -1.0 + for intent in all_intents: + intent_name = INTENT_NAMES.get(intent, str(intent)) + r_val = float(stats["point_biserial_r"][intent][feature_idx].item()) + mean_val = float(stats["mean_by_intent"][intent][feature_idx].item()) + active_rate_val = float(stats["active_rate_by_intent"][intent][feature_idx].item()) + row[f"r_{intent_name}"] = r_val + row[f"mean_{intent_name}"] = mean_val + row[f"active_rate_{intent_name}"] = active_rate_val + if abs(r_val) > best_abs_r: + best_abs_r = abs(r_val) + best_intent = intent_name + + row["best_abs_r"] = best_abs_r + row["best_intent"] = best_intent + rows.append(row) + + rows.sort(key=lambda item: item["eta_sq"], reverse=True) + fieldnames = list(rows[0].keys()) + + output_csv.parent.mkdir(parents=True, exist_ok=True) + with output_csv.open("w", newline="") as f: + writer = csv.DictWriter(f, fieldnames=fieldnames) + writer.writeheader() + writer.writerows(rows) + + +def print_summary(stats: dict, top_k: int) -> None: + eta_sq = stats["eta_sq"] + top_eta = torch.argsort(eta_sq, descending=True)[:top_k].tolist() + + print("Intent counts:") + for intent in stats["all_intents"]: + name = INTENT_NAMES.get(intent, str(intent)) + print(f" {name}: {stats['class_counts'][intent]}") + + print("") + print(f"Top {top_k} SAE features by eta^2:") + for rank, feature_idx in enumerate(top_eta, start=1): + eta_val = float(eta_sq[feature_idx].item()) + + best_intent = None + best_r = None + best_abs_r = -1.0 + per_intent_bits = [] + for intent in stats["all_intents"]: + name = INTENT_NAMES.get(intent, str(intent)) + r_val = float(stats["point_biserial_r"][intent][feature_idx].item()) + mean_val = float(stats["mean_by_intent"][intent][feature_idx].item()) + active_rate_val = float(stats["active_rate_by_intent"][intent][feature_idx].item()) + per_intent_bits.append( + f"{name}: r={r_val:+.4f}, mean={mean_val:.4f}, active={active_rate_val:.3f}" + ) + if abs(r_val) > best_abs_r: + best_abs_r = abs(r_val) + best_intent = name + best_r = r_val + + print( + f"{rank}. feature={feature_idx} eta^2={eta_val:.5f} " + f"best={best_intent} r={best_r:+.4f} active_all={float(stats['active_rate_all'][feature_idx].item()):.3f}" + ) + print(" " + " | ".join(per_intent_bits)) + + print("") + for intent in stats["all_intents"]: + name = INTENT_NAMES.get(intent, str(intent)) + abs_r = torch.abs(stats["point_biserial_r"][intent]) + top_feats = torch.argsort(abs_r, descending=True)[:top_k].tolist() + print(f"Top {top_k} features for {name} by |r|:") + for rank, feature_idx in enumerate(top_feats, start=1): + r_val = float(stats["point_biserial_r"][intent][feature_idx].item()) + eta_val = float(stats["eta_sq"][feature_idx].item()) + mean_val = float(stats["mean_by_intent"][intent][feature_idx].item()) + active_rate_val = float(stats["active_rate_by_intent"][intent][feature_idx].item()) + print( + f" {rank}. feature={feature_idx} r={r_val:+.4f} " + f"eta^2={eta_val:.5f} mean={mean_val:.4f} active={active_rate_val:.3f}" + ) + print("") +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument("--run_root", type=str, required=True) + parser.add_argument("--split", type=str, default="val", choices=["train", "val"]) + parser.add_argument("--output_dir", type=str, default=None) + parser.add_argument("--sae_block", type=int, default=3) + parser.add_argument("--batch_size", type=int, default=4096) + parser.add_argument("--device", type=str, default=default_device()) + parser.add_argument("--top_k", type=int, default=15) + args = parser.parse_args() + + run_root = Path(args.run_root) + output_dir = default_analysis_dir(run_root, args.sae_block, args.output_dir) + output_dir.mkdir(parents=True, exist_ok=True) + device = torch.device(args.device) + bundle = load_sae_bundle(run_root, args.split, args.sae_block, map_location="cpu") + ckpt = bundle["ckpt"] + token_blob = bundle["token_blob"] + + model = build_sae_from_checkpoint(ckpt, bundle["legacy_norm"]) + model.to(device) + + token_tensor, token_key = resolve_token_tensor(token_blob, args.sae_block) + intent_tensor = token_blob["intent"].long() + + stats = compute_stats( + model=model, + token_tensor=token_tensor, + intent_tensor=intent_tensor, + batch_size=args.batch_size, + device=device, + ) + + output_csv = output_dir / f"sae_intent_correlation_block_{args.sae_block}_{args.split}.csv" + write_csv(stats, output_csv) + print_summary(stats, top_k=args.top_k) + print(f"Used token key {token_key}") + print(f"Saved CSV to {output_csv}") diff --git a/src/camera-based-e2e/analyze_sae_visual_gen.py b/src/camera-based-e2e/analyze_sae_visual_gen.py new file mode 100644 index 0000000..ec391c2 --- /dev/null +++ b/src/camera-based-e2e/analyze_sae_visual_gen.py @@ -0,0 +1,520 @@ +''' +NOTE: This must be run in a GPU env with sglang installed. + +Command for Illinois NCSA Delta GPU: +``` +export MODEL_ROOT=/work/nvme/bgxf/mgagvani/hfcache + export HF_HOME=$MODEL_ROOT/huggingface + export HF_HUB_CACHE=$HF_HOME/hub + export TRANSFORMERS_CACHE=$HF_HUB_CACHE + export PATH=/work/nvme/bgxf/mgagvani/conda/envs/robotvision/bin:$PATH + CUDA_VISIBLE_DEVICES=0 /work/nvme/bgxf/mgagvani/conda/envs/robotvision/bin/python -m sglang.launch_server \ + --model-path "$MODEL_ROOT/checkpoints/Qwen3.5-35B-A3B-GPTQ-Int4" \ + --served-model-name qwen3.5-35b-a3b \ + --host 127.0.0.1 \ + --port 8000 \ + --tp-size 1 \ + --mem-fraction-static 0.80 \ + --context-length 8192 \ + --dtype float16 \ + --disable-cuda-graph \ + --reasoning-parser qwen3 +``` + +NCSA Delta AI: +``` +export SGLANG_MAMBA_CONV_DTYPE=float16 +export LD_LIBRARY_PATH="$VIRTUAL_ENV/lib/python3.12/site-packages/nvidia/cu13/lib:$VIRTUAL_ENV/lib/python3.12/site-packages/torch/lib${LD_LIBRARY_PATH:+:$LD_LIBRARY_PATH}" +export MODEL_ROOT=/work/nvme/bgxf/mgagvani/hfcache +export HF_HOME=$MODEL_ROOT/huggingface +export HF_HUB_CACHE=$HF_HOME/hub +export TRANSFORMERS_CACHE=$HF_HUB_CACHE +CUDA_VISIBLE_DEVICES=0 /u/mgagvani/robotvision/.venv/bin/python -m sglang.launch_server \ + --model-path "$MODEL_ROOT/checkpoints/Qwen3.5-35B-A3B-GPTQ-Int4" \ + --served-model-name qwen3.5-35b-a3b \ + --host 127.0.0.1 \ + --port 8000 \ + --tp-size 1 \ + --mem-fraction-static 0.80 \ + --context-length 8192 \ + --dtype float16 \ + --mamba-ssm-dtype float16 \ + --linear-attn-prefill-backend triton \ + --linear-attn-backend flashinfer \ + --disable-cuda-graph \ + --reasoning-parser qwen3 +``` + +``` +export VLLM_DEEP_GEMM_WARMUP=skip +export MODEL_ROOT=/work/nvme/bgxf/mgagvani/hfcache +export HF_HOME=$MODEL_ROOT/huggingface +export HF_HUB_CACHE=$HF_HOME/hub +export TRANSFORMERS_CACHE=$HF_HUB_CACHE +CUDA_VISIBLE_DEVICES=0 /u/mgagvani/robotvision/.venv/bin/vllm serve \ + "$MODEL_ROOT/checkpoints/Qwen3.5-4B" \ + --served-model-name qwen3.5-4b \ + --host 127.0.0.1 \ + --port 8000 \ + --tensor-parallel-size 1 \ + --gpu-memory-utilization 0.35 \ + --max-num-seqs 64 \ + --max-model-len 8192 \ + --dtype float16 \ + --limit-mm-per-prompt '{"image": 1}' + --reasoning-parser qwen3 +``` + +Note that --tp-size 1 is for 1 GPU, if N gpus set to N. +''' + + +from ultralytics import YOLO +from loader import WaymoE2E +import argparse +import json +import io +from PIL import Image + +import os +from pathlib import Path +import dotenv +from google import genai +from google.genai import types +from openai import OpenAI +import base64 + +from diffusers import QwenImageEditPipeline, Flux2KleinPipeline, QwenImageEditPlusPipeline +import torch +torch.backends.cuda.enable_cudnn_sdp(False) +torch.backends.cuda.enable_flash_sdp(True) +torch.backends.cuda.enable_mem_efficient_sdp(True) +torch.backends.cuda.enable_math_sdp(True) + + +PROMPT = ''' +Choose one minimal counterfactual edit for an image editing model from one of three things: +1) The color of traffic lights in the image +2) The presence of stop signs in the image +3) The presence of pedestrians in the image. + +E.g., if there are no traffic lights, do not write a prompt involving traffic lights. +If there is, ask to change the color of the traffic light following the rule "RED -> GREEN, GREEN -> RED, YELLOW -> GREEN". +If there is no stop sign, ask to add one, or vice versa, only if at an intersection. +If there are no pedestrians, add some, and vice versa, only if there is a sidewalk or crosswalk visible. +If no condition is met, set edit_type and edit_direction to "no_change" and prompt to "NO CHANGE". +The prompt should be visually descriptive, at most 3 sentences. +Do not include unnecessary detail about tone/style, as we want the new image to be as similar to the original as possible aside from the specified edit. +Do not mention bounding box coordinates in the prompt, but you can reference relative positions of objects. +Make sure to mention "Edit the image minimally. Preserve original composition, lighting, texture, and all unrelated details." + +Return strict JSON only, with no markdown or extra text: +{ + "edit_type": "traffic_light_color | traffic_light_presence | pedestrian_presence | stop_sign_presence_location | no_change", + "edit_direction": "red_to_green | green_to_red | yellow_to_green | add_traffic_light | remove_traffic_light | add_pedestrian | remove_pedestrian | add_stop_sign | remove_stop_sign | move_stop_sign | no_change", + "prompt": "image editing prompt or NO CHANGE" +} +''' + +VALID_EDIT_TYPES = { + "traffic_light_color", + "traffic_light_presence", + "pedestrian_presence", + "stop_sign_presence_location", + "no_change", +} + +VALID_EDIT_DIRECTIONS = { + "red_to_green", + "green_to_red", + "yellow_to_green", + "add_traffic_light", + "remove_traffic_light", + "add_pedestrian", + "remove_pedestrian", + "add_stop_sign", + "remove_stop_sign", + "move_stop_sign", + "no_change", +} + +def load_yolo(model_path: str = "yolo26x.pt"): + # Load a model + model = YOLO(model_path) + return model + +def inference_yolo(model, images): + # Predict with the model + results = model(images, verbose=False, device="cuda:0") # predict on an image + + # Access the results + for result in results: + xywh = result.boxes.xywh # center-x, center-y, width, height + xywhn = result.boxes.xywhn # normalized + xyxy = result.boxes.xyxy # top-left-x, top-left-y, bottom-right-x, bottom-right-y + xyxyn = result.boxes.xyxyn # normalized + names = [result.names[cls.item()] for cls in result.boxes.cls.int()] # class name of each box + confs = result.boxes.conf # confidence score of each box + + return results + +def parse_edit_plan(raw_output: str) -> dict: + raw = (raw_output or "").strip() + if raw == "NO CHANGE": + return { + "edit_type": "no_change", + "edit_direction": "no_change", + "prompt": "NO CHANGE", + "raw_generator_output": raw_output, + } + + cleaned = raw + if cleaned.startswith("```"): + lines = cleaned.splitlines() + if lines and lines[0].startswith("```"): + lines = lines[1:] + if lines and lines[-1].startswith("```"): + lines = lines[:-1] + cleaned = "\n".join(lines).strip() + + try: + data = json.loads(cleaned) + except json.JSONDecodeError: + start = cleaned.find("{") + end = cleaned.rfind("}") + if start >= 0 and end > start: + try: + data = json.loads(cleaned[start : end + 1]) + except json.JSONDecodeError: + data = None + else: + data = None + + if not isinstance(data, dict): + return { + "edit_type": "no_change", + "edit_direction": "no_change", + "prompt": "NO CHANGE", + "raw_generator_output": raw_output, + "parse_error": "generator_output_was_not_valid_json", + } + + edit_type = str(data.get("edit_type", "")).strip() + edit_direction = str(data.get("edit_direction", "")).strip() + prompt = str(data.get("prompt", "")).strip() + + if edit_type not in VALID_EDIT_TYPES or edit_direction not in VALID_EDIT_DIRECTIONS: + return { + "edit_type": "no_change", + "edit_direction": "no_change", + "prompt": "NO CHANGE", + "raw_generator_output": raw_output, + "parse_error": "invalid_edit_type_or_direction", + } + + if edit_type == "no_change" or edit_direction == "no_change" or prompt == "NO CHANGE": + edit_type = "no_change" + edit_direction = "no_change" + prompt = "NO CHANGE" + elif "Edit the image minimally" not in prompt: + prompt = ( + f"{prompt} Edit the image minimally. Preserve original composition, " + "lighting, texture, and all unrelated details." + ) + + return { + "edit_type": edit_type, + "edit_direction": edit_direction, + "prompt": prompt, + "raw_generator_output": raw_output, + } + +def jpeg_tensor_to_image(jpeg) -> Image.Image: + if hasattr(jpeg, "numpy"): + jpeg = jpeg.numpy().tobytes() + elif hasattr(jpeg, "tobytes"): + jpeg = jpeg.tobytes() + return Image.open(io.BytesIO(jpeg)).convert("RGB") + +def yolo_scene_description(result) -> str: + names = [result.names[cls.item()] for cls in result.boxes.cls.int()] + confs = result.boxes.conf + xyxy = result.boxes.xyxy + lines = [] + for j, (name, conf) in enumerate(zip(names, confs)): + lines.append(f"{name} conf={float(conf.item()):.3f} at {xyxy[j].tolist()}") + return "\n".join(lines) + +def generate_gemini(scene_description: str): + # Load Gemini API key from .env file + dotenv.load_dotenv() + api_key = os.getenv("GEMINI_API_KEY") + if not api_key: + raise RuntimeError("GEMINI_API_KEY is not set. Add it to .env or export it before running.") + genai_client = genai.Client(api_key=api_key) + + # Generate prompt for image editing model + response = genai_client.models.generate_content( + model="gemini-flash-lite-latest", + contents=[ + types.Content( + role="user", + parts=[types.Part.from_text(text=scene_description)], + ), + ], + config=types.GenerateContentConfig( + system_instruction=PROMPT, + thinking_config=types.ThinkingConfig(thinking_level="MINIMAL"), + ), + ) + + return response.text + +def image_to_url(image: Image.Image) -> str: + buffered = io.BytesIO() + if image.mode != "RGB": + image = image.convert("RGB") + image.save(buffered, format="JPEG") + img_str = base64.b64encode(buffered.getvalue()).decode("utf-8") + return f"data:image/jpeg;base64,{img_str}" + +def generate_local(scene_description: str, image: Image.Image | None = None): + client = OpenAI( + base_url=os.getenv("BASE_URL", "http://127.0.0.1:8000/v1"), + api_key=os.getenv("API_KEY", "EMPTY"), + ) + user_content = [{"type": "text", "text": scene_description}] + if image is not None: + user_content.append({ + "type": "image_url", + "image_url": { + "url": image_to_url(image), + }, + }) + messages = [ + { + "role": "system", + "content": PROMPT.strip(), + }, + { + "role": "user", + "content": user_content, + }, + ] + response = client.chat.completions.create( + model=os.getenv("MODEL", "qwen3.5-4b"), + messages=messages, + max_tokens=512, + temperature=0.6, + extra_body={"chat_template_kwargs": {"enable_thinking": False}}, + ) + return response.choices[0].message.content.strip() + +def load_qwen_image_edit(): + pipeline = QwenImageEditPipeline.from_pretrained("Qwen/Qwen-Image-Edit") + print("pipeline loaded") + pipeline.to(torch.bfloat16) + pipeline.to("cuda") + pipeline.set_progress_bar_config(disable=None) + + return pipeline + +def generate_image_edit(pipeline, image, prompt): + inputs = { + "image": image, + "prompt": prompt, + "generator": torch.manual_seed(0), + "true_cfg_scale": 4.0, + "negative_prompt": " ", + "num_inference_steps": 50, + } + + with torch.inference_mode(): + output = pipeline(**inputs) + output_image = output.images[0] + + return output_image + +def load_flux2klein(): + pipe = Flux2KleinPipeline.from_pretrained( + "black-forest-labs/FLUX.2-klein-9b-kv", + torch_dtype=torch.bfloat16, + ) + pipe.to("cuda") + pipe.to(torch.bfloat16) + pipe.set_progress_bar_config(disable=None) + return pipe + +def generate_flux2klein_edit(pipe, image, prompt): + out = pipe( + prompt=prompt, + image=image, + num_inference_steps=4, + generator=torch.Generator("cuda").manual_seed(0), + ).images[0] + return out + +def load_firered(): + pipe = QwenImageEditPlusPipeline.from_pretrained( + "FireRedTeam/FireRed-Image-Edit-1.1", + torch_dtype=torch.bfloat16, + ) + pipe.to("cuda") + pipe.to(torch.bfloat16) + pipe.set_progress_bar_config(disable=None) + return pipe + +def generate_firered_edit(pipe, image, prompt): + # image must be a list + if not isinstance(image, list): + image = [image] + out = pipe( + prompt=prompt, + negative_prompt=" ", # TODO: negative prompt + image=image, + num_inference_steps=40, + true_cfg_scale=3.0, + generator=torch.Generator("cuda").manual_seed(0), + ).images[0] + return out + +def load_editor(name: str): + if name == "firered": + return load_firered() + if name == "qwen": + return load_qwen_image_edit() + if name == "flux": + return load_flux2klein() + raise ValueError(f"Unknown editor: {name}") + +def generate_edit_with_editor(editor_name: str, pipeline, image: Image.Image, prompt: str) -> Image.Image: + if editor_name == "firered": + return generate_firered_edit(pipeline, image, prompt) + if editor_name == "qwen": + return generate_image_edit(pipeline, image, prompt) + if editor_name == "flux": + return generate_flux2klein_edit(pipeline, image, prompt) + raise ValueError(f"Unknown editor: {editor_name}") + +def generate_prompt(generator: str, scene_description: str, image: Image.Image) -> str: + if generator == "gemini": + return generate_gemini(scene_description) + if generator == "local": + return generate_local(scene_description, image) + raise ValueError(f"Unknown generator: {generator}") + +def run_generate_edits(args) -> None: + output_dir = Path(args.output_dir) + edited_dir = output_dir / "edited" + edited_dir.mkdir(parents=True, exist_ok=True) + manifest_path = output_dir / "manifest.jsonl" + + yolo_model = load_yolo(args.yolo_model) + dataset = WaymoE2E( + indexFile=args.index_file, + data_dir=args.data_dir, + n_items=args.n_items, + ) + + end_idx = min(args.start_idx + args.max_items, len(dataset)) + dataset_indices = list(range(args.start_idx, end_idx)) + images = [] + samples = [] + for dataset_idx in dataset_indices: + sample = dataset[dataset_idx] + samples.append(sample) + images.append(jpeg_tensor_to_image(sample["IMAGES_JPEG"][args.camera_idx])) + + # Generate batches + batches, bs = [], 16 + for i in range(0, len(images), bs): + batches.append((images[i : i + bs], samples[i : i + bs], dataset_indices[i : i + bs])) + + results = [] + for batch_images, batch_samples, batch_indices in batches: + batch_results = inference_yolo(yolo_model, batch_images) + results.extend(batch_results) + + pipeline = None + + with manifest_path.open("w") as f: + for dataset_idx, sample, image, result in zip(dataset_indices, samples, images, results): + scene_description = yolo_scene_description(result) + raw_prompt = generate_prompt(args.generator, scene_description, image) + edit_plan = parse_edit_plan(raw_prompt) + prompt = edit_plan["prompt"] + + row = { + "dataset_idx": dataset_idx, + "name": sample["NAME"], + "camera_idx": args.camera_idx, + "index_file": args.index_file, + "n_items": args.n_items, + "status": "no_change", + "edit_type": edit_plan["edit_type"], + "edit_direction": edit_plan["edit_direction"], + "prompt": prompt, + "scene_description": scene_description, + "generator": args.generator, + "editor": args.editor, + "edited_path": None, + "raw_generator_output": edit_plan.get("raw_generator_output"), + } + if "parse_error" in edit_plan: + row["parse_error"] = edit_plan["parse_error"] + + print(f"dataset_idx={dataset_idx} edit_type={row['edit_type']} direction={row['edit_direction']}") + print(prompt) + + if row["edit_type"] != "no_change": + if pipeline is None: + hf_home = os.getenv("HF_HOME") + if hf_home is None: + print("WARNING: HF_HOME is not set; diffusers will use its default cache.") + pipeline = load_editor(args.editor) + edited_image = generate_edit_with_editor(args.editor, pipeline, image, prompt) + edited_path = edited_dir / f"{dataset_idx}.jpg" + edited_image.save(edited_path) + row["status"] = "edited" + row["edited_path"] = str(edited_path.relative_to(output_dir)) + else: + print(f"No change for dataset_idx={dataset_idx}, skipping edit.") + + f.write(json.dumps(row) + "\n") + f.flush() + + print(f"Saved manifest to {manifest_path}") + +def build_parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser() + subparsers = parser.add_subparsers(dest="command", required=True) + + gen = subparsers.add_parser("generate-edits") + gen.add_argument("--data_dir", type=str, required=True, help="Path to Waymo directory") + gen.add_argument("--index_file", type=str, default="index_val.pkl") + gen.add_argument("--output_dir", type=str, required=True) + gen.add_argument("--n_items", type=int, default=5_000) + gen.add_argument("--start_idx", type=int, default=0) + gen.add_argument("--max_items", type=int, default=10) + gen.add_argument("--camera_idx", type=int, default=1) + gen.add_argument("--yolo_model", type=str, default="yolo26x.pt") + gen.add_argument( + "--generator", + type=str, + choices=["local", "gemini"], + default="local", + help="Prompt generator backend. Assumes local LLM is already running for --generator local.", + ) + gen.add_argument( + "--editor", + type=str, + choices=["firered", "qwen", "flux"], + default="firered", + ) + gen.set_defaults(func=run_generate_edits) + return parser + +if __name__ == "__main__": + parser = build_parser() + args = parser.parse_args() + args.func(args) diff --git a/src/camera-based-e2e/analyze_sae_visual_gen_pt2.py b/src/camera-based-e2e/analyze_sae_visual_gen_pt2.py new file mode 100644 index 0000000..fa0cc0c --- /dev/null +++ b/src/camera-based-e2e/analyze_sae_visual_gen_pt2.py @@ -0,0 +1,711 @@ +import argparse +import csv +import json +from pathlib import Path + +import numpy as np +import torch + +from extract_planner_tok import load_model +from loader import WaymoE2E +from models.base_model import collate_with_images +from sae_utils import ( + build_sae_from_checkpoint, + default_analysis_dir, + default_device, + encode_tensor_batchwise, + load_sae_bundle, + planner_token_key, + resolve_token_tensor, +) + + +PLANNER_DELTA_NAMES = ( + "trajectory_l2_delta", + "delta_selected_ade", + "delta_oracle_ade", + "delta_score_margin", + "delta_proposal_spread", +) + + +def read_jsonl(path: Path) -> list[dict]: + with path.open() as f: + return [json.loads(line) for line in f if line.strip()] + + +def write_csv(path: Path, rows: list[dict], fieldnames: list[str] | None = None) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + if not rows: + return + if fieldnames is None: + fieldnames = list(rows[0].keys()) + with path.open("w", newline="") as f: + writer = csv.DictWriter(f, fieldnames=fieldnames) + writer.writeheader() + writer.writerows(rows) + + +def infer_checkpoint_path(run_root: Path, token_blob: dict, explicit_path: str | None) -> str: + if explicit_path is not None: + return explicit_path + meta_path = token_blob.get("meta", {}).get("checkpoint") + if meta_path and Path(meta_path).exists(): + return meta_path + repo_default = Path(__file__).resolve().parent / "camera-e2e-epoch=04-val_loss=2.90.ckpt" + if repo_default.exists(): + return str(repo_default) + raise FileNotFoundError("Could not infer planner checkpoint path. Pass --planner_checkpoint.") + + +def resolve_edited_path(manifest_path: Path, row: dict) -> Path: + edited_value = row.get("edited_path") + if not edited_value: + raise ValueError(f"Missing edited_path for dataset_idx={row.get('dataset_idx')}") + edited_path = Path(edited_value) + if not edited_path.is_absolute(): + edited_path = manifest_path.parent / edited_path + return edited_path + + +def jpeg_bytes_to_tensor(path: Path) -> torch.Tensor: + data = path.read_bytes() + return torch.from_numpy(np.frombuffer(data, dtype=np.uint8).copy()) + + +def load_manifest_rows( + manifest_path: Path, + *, + strict_manifest: bool, +) -> list[dict]: + rows = [] + for row in read_jsonl(manifest_path): + if row.get("status") != "edited": + continue + edited_path = resolve_edited_path(manifest_path, row) + if not edited_path.exists(): + message = f"Missing edited image for dataset_idx={row.get('dataset_idx')}: {edited_path}" + if strict_manifest: + raise FileNotFoundError(message) + print(f"WARNING: {message}; skipping") + continue + row = dict(row) + row["resolved_edited_path"] = str(edited_path) + rows.append(row) + return rows + + +def choose_manifest_value(rows: list[dict], key: str, override): + if override is not None: + return override + values = {row.get(key) for row in rows if row.get(key) is not None} + if len(values) == 1: + return values.pop() + if not values: + return None + raise ValueError(f"Manifest has multiple {key} values; pass --{key} explicitly.") + + +def build_token_name_lookup(token_blob: dict) -> dict[str, int]: + names = token_blob.get("names") + if names is None: + raise KeyError("Token blob is missing names; cannot align original tokens to manifest rows.") + + lookup = {} + duplicates = set() + for idx, name in enumerate(names): + if name in lookup: + duplicates.add(name) + lookup[name] = idx + if duplicates: + examples = ", ".join(sorted(duplicates)[:5]) + raise ValueError(f"Token blob has duplicate sample names, cannot build unambiguous lookup: {examples}") + return lookup + + +def make_pair_samples( + dataset: WaymoE2E, + rows: list[dict], + *, + allow_name_mismatch: bool, +) -> tuple[list[dict], list[dict]]: + orig_samples = [] + edited_samples = [] + for row in rows: + dataset_idx = int(row["dataset_idx"]) + sample = dataset[dataset_idx] + if not allow_name_mismatch and row.get("name") != sample["NAME"]: + raise ValueError( + f"Manifest name mismatch at dataset_idx={dataset_idx}: " + f"manifest={row.get('name')} dataset={sample['NAME']}" + ) + + camera_idx = int(row.get("camera_idx", 1)) + edited_sample = { + "PAST": sample["PAST"], + "FUTURE": sample["FUTURE"], + "INTENT": sample["INTENT"], + "NAME": sample["NAME"], + "IMAGES_JPEG": list(sample["IMAGES_JPEG"]), + } + edited_sample["IMAGES_JPEG"][camera_idx] = jpeg_bytes_to_tensor(Path(row["resolved_edited_path"])) + orig_samples.append(sample) + edited_samples.append(edited_sample) + return orig_samples, edited_samples + + +def selected_and_oracle_ade( + trajectory: torch.Tensor, + scores: torch.Tensor, + future: torch.Tensor, + *, + num_proposals: int, + horizon: int, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + batch = future.size(0) + traj = trajectory.view(batch, num_proposals, horizon, 2) + dist = torch.norm(traj - future[:, None], dim=-1) + ade_per_mode = dist.mean(dim=-1) + selected_idx = scores.argmin(dim=1) + row_idx = torch.arange(batch, device=future.device) + selected_ade = ade_per_mode[row_idx, selected_idx] + oracle_ade = ade_per_mode.min(dim=1).values + return selected_idx, selected_ade, oracle_ade + + +def score_margin(scores: torch.Tensor) -> torch.Tensor: + sorted_scores = scores.sort(dim=1).values + if scores.size(1) < 2: + return torch.zeros(scores.size(0), device=scores.device, dtype=scores.dtype) + return sorted_scores[:, 1] - sorted_scores[:, 0] + + +def proposal_spread(trajectory: torch.Tensor, *, num_proposals: int, horizon: int) -> torch.Tensor: + batch = trajectory.size(0) + traj = trajectory.view(batch, num_proposals, horizon, 2) + mean_traj = traj.mean(dim=1, keepdim=True) + return torch.norm(traj - mean_traj, dim=-1).mean(dim=(1, 2)) + + +def controls_view(controls: torch.Tensor, *, num_proposals: int, horizon: int) -> torch.Tensor: + return controls.view(controls.size(0), num_proposals, horizon, 2) + + +def compute_pair_metrics( + out_orig: dict, + out_edit: dict, + future: torch.Tensor, + *, + num_proposals: int, + horizon: int, +) -> dict[str, torch.Tensor]: + selected_idx_orig, selected_ade_orig, oracle_ade_orig = selected_and_oracle_ade( + out_orig["trajectory"], + out_orig["scores"], + future, + num_proposals=num_proposals, + horizon=horizon, + ) + selected_idx_edit, selected_ade_edit, oracle_ade_edit = selected_and_oracle_ade( + out_edit["trajectory"], + out_edit["scores"], + future, + num_proposals=num_proposals, + horizon=horizon, + ) + + score_margin_orig = score_margin(out_orig["scores"]) + score_margin_edit = score_margin(out_edit["scores"]) + spread_orig = proposal_spread(out_orig["trajectory"], num_proposals=num_proposals, horizon=horizon) + spread_edit = proposal_spread(out_edit["trajectory"], num_proposals=num_proposals, horizon=horizon) + + ctrl_orig = controls_view(out_orig["controls"], num_proposals=num_proposals, horizon=horizon) + ctrl_edit = controls_view(out_edit["controls"], num_proposals=num_proposals, horizon=horizon) + ctrl_delta = (ctrl_edit - ctrl_orig).abs() + + return { + "trajectory_l2_delta": torch.norm(out_edit["trajectory"] - out_orig["trajectory"], dim=1), + "selected_idx_orig": selected_idx_orig, + "selected_idx_edit": selected_idx_edit, + "selected_ade_orig": selected_ade_orig, + "selected_ade_edit": selected_ade_edit, + "delta_selected_ade": selected_ade_edit - selected_ade_orig, + "oracle_ade_orig": oracle_ade_orig, + "oracle_ade_edit": oracle_ade_edit, + "delta_oracle_ade": oracle_ade_edit - oracle_ade_orig, + "score_margin_orig": score_margin_orig, + "score_margin_edit": score_margin_edit, + "delta_score_margin": score_margin_edit - score_margin_orig, + "proposal_spread_orig": spread_orig, + "proposal_spread_edit": spread_edit, + "delta_proposal_spread": spread_edit - spread_orig, + "mean_abs_accel_delta": ctrl_delta[..., 0].mean(dim=(1, 2)), + "mean_abs_omega_delta": ctrl_delta[..., 1].mean(dim=(1, 2)), + } + + +def model_inputs_from_batch(batch: dict, *, lit_model, device: torch.device) -> dict: + return { + "PAST": batch["PAST"].to(device, non_blocking=True), + "IMAGES": lit_model.decode_batch_jpeg(batch["IMAGES_JPEG"], device=device), + "INTENT": batch["INTENT"].to(device, non_blocking=True), + } + + +def compute_baseline_stats( + *, + sae, + token_tensor: torch.Tensor, + batch_size: int, + device: torch.device, + min_scale: float, +) -> dict[str, torch.Tensor]: + z_all = encode_tensor_batchwise(sae, token_tensor, batch_size=batch_size, device=device) + all_mean = z_all.mean(dim=0) + all_std = z_all.std(dim=0, unbiased=False) + active_mask = z_all > 0 + active_count = active_mask.sum(dim=0) + active_sum = z_all.sum(dim=0) + active_sum_sq = (z_all * z_all).sum(dim=0) + denom = active_count.clamp_min(1) + active_mean = active_sum / denom + active_var = active_sum_sq / denom - active_mean.square() + active_std = torch.sqrt(active_var.clamp_min(0.0)) + feature_scale = torch.maximum(active_std, 0.25 * active_mean).clamp_min(min_scale) + return { + "all_mean": all_mean, + "all_std": all_std.clamp_min(1e-6), + "active_count": active_count, + "active_rate": active_count.float() / max(1, z_all.size(0)), + "active_mean": active_mean, + "active_std": active_std, + "feature_scale": feature_scale, + } + + +def replay_pairs( + *, + rows: list[dict], + dataset: WaymoE2E, + planner_model, + lit_model, + sae, + sae_block: int, + token_blob: dict, + token_tensor: torch.Tensor, + token_name_to_idx: dict[str, int], + batch_size: int, + device: torch.device, + allow_name_mismatch: bool, +) -> dict: + token_key = planner_token_key(sae_block) + pair_rows = [] + tokens_orig_chunks = [] + tokens_edit_chunks = [] + z_orig_chunks = [] + z_edit_chunks = [] + planner_metric_chunks = {name: [] for name in ( + "trajectory_l2_delta", + "selected_idx_orig", + "selected_idx_edit", + "selected_ade_orig", + "selected_ade_edit", + "delta_selected_ade", + "oracle_ade_orig", + "oracle_ade_edit", + "delta_oracle_ade", + "score_margin_orig", + "score_margin_edit", + "delta_score_margin", + "proposal_spread_orig", + "proposal_spread_edit", + "delta_proposal_spread", + "mean_abs_accel_delta", + "mean_abs_omega_delta", + )} + + planner_model.eval() + sae.eval() + with torch.no_grad(): + for start in range(0, len(rows), batch_size): + batch_rows = rows[start : start + batch_size] + _, edited_samples = make_pair_samples( + dataset, + batch_rows, + allow_name_mismatch=allow_name_mismatch, + ) + edit_batch = collate_with_images(edited_samples) + token_indices = [] + for row in batch_rows: + name = row.get("name") + if name not in token_name_to_idx: + raise KeyError(f"Manifest sample name is not present in token blob: {name}") + token_indices.append(token_name_to_idx[name]) + token_indices = torch.tensor(token_indices, dtype=torch.long) + + token_orig_cpu = token_tensor.index_select(0, token_indices) + token_orig = token_orig_cpu.to(device, non_blocking=True) + out_orig = { + "trajectory": token_blob["trajectory"].index_select(0, token_indices).to(device, non_blocking=True), + "scores": token_blob["scores"].index_select(0, token_indices).to(device, non_blocking=True), + "controls": token_blob["controls"].index_select(0, token_indices).to(device, non_blocking=True), + } + future = token_blob["future"].index_select(0, token_indices).to(device, non_blocking=True) + out_edit = planner_model( + model_inputs_from_batch(edit_batch, lit_model=lit_model, device=device), + return_block_tokens=True, + ) + + token_edit = out_edit[token_key] + z_orig = sae.encode(token_orig) + z_edit = sae.encode(token_edit) + metrics = compute_pair_metrics( + out_orig, + out_edit, + future, + num_proposals=planner_model.n_proposals, + horizon=planner_model.horizon, + ) + + tokens_orig_chunks.append(token_orig.detach().cpu()) + tokens_edit_chunks.append(token_edit.detach().cpu()) + z_orig_chunks.append(z_orig.detach().cpu()) + z_edit_chunks.append(z_edit.detach().cpu()) + for name, value in metrics.items(): + planner_metric_chunks[name].append(value.detach().cpu()) + + token_delta_norm = torch.norm(token_edit - token_orig, dim=1).detach().cpu() + latent_delta_norm = torch.norm(z_edit - z_orig, dim=1).detach().cpu() + for i, row in enumerate(batch_rows): + pair_row = { + "dataset_idx": int(row["dataset_idx"]), + "name": row.get("name"), + "camera_idx": int(row.get("camera_idx", 1)), + "edit_type": row.get("edit_type"), + "edit_direction": row.get("edit_direction"), + "edited_path": row["resolved_edited_path"], + "token_blob_idx": int(token_indices[i].item()), + "original_source": "token_blob", + "token_l2_delta": float(token_delta_norm[i].item()), + "latent_l2_delta": float(latent_delta_norm[i].item()), + "latent_mean_abs_delta": float((z_edit[i] - z_orig[i]).abs().mean().item()), + } + for metric_name, metric_value in metrics.items(): + value = metric_value[i] + if value.dtype in {torch.int8, torch.int16, torch.int32, torch.int64, torch.long}: + pair_row[metric_name] = int(value.item()) + else: + pair_row[metric_name] = float(value.item()) + pair_rows.append(pair_row) + + print(f"Processed {min(start + batch_size, len(rows))}/{len(rows)} edited pairs", flush=True) + + return { + "pair_rows": pair_rows, + "tokens_orig": torch.cat(tokens_orig_chunks, dim=0), + "tokens_edit": torch.cat(tokens_edit_chunks, dim=0), + "z_orig": torch.cat(z_orig_chunks, dim=0), + "z_edit": torch.cat(z_edit_chunks, dim=0), + "planner_metrics": { + name: torch.cat(chunks, dim=0) for name, chunks in planner_metric_chunks.items() + }, + } + + +def pearson_vector(x: torch.Tensor, y: torch.Tensor) -> torch.Tensor: + if x.size(0) < 2: + return torch.zeros(x.size(1), dtype=torch.float32) + x = x.float() + y = y.float() + x_center = x - x.mean(dim=0, keepdim=True) + y_center = y - y.mean() + x_denom = torch.sqrt((x_center * x_center).sum(dim=0)).clamp_min(1e-6) + y_denom = torch.sqrt((y_center * y_center).sum()).clamp_min(1e-6) + return (x_center * y_center[:, None]).sum(dim=0) / (x_denom * y_denom) + + +def summarize_group( + *, + group_kind: str, + group_value: str, + mask: torch.Tensor, + z_orig: torch.Tensor, + z_edit: torch.Tensor, + baseline_stats: dict[str, torch.Tensor], + planner_metrics: dict[str, torch.Tensor], +) -> list[dict]: + z_o = z_orig[mask] + z_e = z_edit[mask] + delta = z_e - z_o + abs_delta = delta.abs() + n = delta.size(0) + latent_dim = delta.size(1) + if n == 0: + return [] + + mean_delta = delta.mean(dim=0) + mean_abs_delta = abs_delta.mean(dim=0) + std_delta = delta.std(dim=0, unbiased=False) + stderr_delta = std_delta / (n ** 0.5) + paired_t = mean_delta / stderr_delta.clamp_min(1e-6) + cohen_dz = mean_delta / std_delta.clamp_min(1e-6) + sign_consistency = torch.where( + mean_delta >= 0, + (delta > 0).float().mean(dim=0), + (delta < 0).float().mean(dim=0), + ) + + orig_active = z_o > 0 + edit_active = z_e > 0 + orig_active_rate = orig_active.float().mean(dim=0) + edited_active_rate = edit_active.float().mean(dim=0) + flip_on_rate = ((~orig_active) & edit_active).float().mean(dim=0) + flip_off_rate = (orig_active & (~edit_active)).float().mean(dim=0) + + feature_scale = baseline_stats["feature_scale"] + all_std = baseline_stats["all_std"] + mean_delta_scale_units = mean_delta / feature_scale + mean_abs_delta_scale_units = mean_abs_delta / feature_scale + mean_delta_all_std_units = mean_delta / all_std + mean_abs_delta_all_std_units = mean_abs_delta / all_std + + corr_delta_selected = pearson_vector(delta, planner_metrics["delta_selected_ade"][mask]) + corr_abs_traj = pearson_vector(abs_delta, planner_metrics["trajectory_l2_delta"][mask]) + corr_delta_score_margin = pearson_vector(delta, planner_metrics["delta_score_margin"][mask]) + + rows = [] + for feature_idx in range(latent_dim): + rows.append( + { + "group_kind": group_kind, + "group_value": group_value, + "feature_idx": feature_idx, + "n": n, + "mean_delta": float(mean_delta[feature_idx].item()), + "mean_abs_delta": float(mean_abs_delta[feature_idx].item()), + "std_delta": float(std_delta[feature_idx].item()), + "stderr_delta": float(stderr_delta[feature_idx].item()), + "paired_t": float(paired_t[feature_idx].item()), + "cohen_dz": float(cohen_dz[feature_idx].item()), + "mean_delta_scale_units": float(mean_delta_scale_units[feature_idx].item()), + "mean_abs_delta_scale_units": float(mean_abs_delta_scale_units[feature_idx].item()), + "mean_delta_all_std_units": float(mean_delta_all_std_units[feature_idx].item()), + "mean_abs_delta_all_std_units": float(mean_abs_delta_all_std_units[feature_idx].item()), + "baseline_active_rate": float(baseline_stats["active_rate"][feature_idx].item()), + "orig_active_rate": float(orig_active_rate[feature_idx].item()), + "edited_active_rate": float(edited_active_rate[feature_idx].item()), + "flip_on_rate": float(flip_on_rate[feature_idx].item()), + "flip_off_rate": float(flip_off_rate[feature_idx].item()), + "sign_consistency": float(sign_consistency[feature_idx].item()), + "corr_delta_z_delta_selected_ade": float(corr_delta_selected[feature_idx].item()), + "corr_abs_delta_z_trajectory_l2_delta": float(corr_abs_traj[feature_idx].item()), + "corr_delta_z_delta_score_margin": float(corr_delta_score_margin[feature_idx].item()), + } + ) + return rows + + +def build_summary_rows( + *, + rows: list[dict], + z_orig: torch.Tensor, + z_edit: torch.Tensor, + baseline_stats: dict[str, torch.Tensor], + planner_metrics: dict[str, torch.Tensor], +) -> list[dict]: + out = [] + for group_kind in ("edit_type", "edit_direction"): + values = sorted({str(row.get(group_kind)) for row in rows}) + for value in values: + mask = torch.tensor([str(row.get(group_kind)) == value for row in rows], dtype=torch.bool) + out.extend( + summarize_group( + group_kind=group_kind, + group_value=value, + mask=mask, + z_orig=z_orig, + z_edit=z_edit, + baseline_stats=baseline_stats, + planner_metrics=planner_metrics, + ) + ) + return out + + +def save_feature_deltas_pt( + path: Path, + *, + rows: list[dict], + z_orig: torch.Tensor, + z_edit: torch.Tensor, + baseline_stats: dict[str, torch.Tensor], +) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + delta = z_edit - z_orig + torch.save( + { + "rows": rows, + "z_orig": z_orig, + "z_edit": z_edit, + "delta_z": delta, + "abs_delta_z": delta.abs(), + "delta_z_scale_units": delta / baseline_stats["feature_scale"], + "delta_z_all_std_units": delta / baseline_stats["all_std"], + "orig_active": z_orig > 0, + "edited_active": z_edit > 0, + "flip_on": (z_orig <= 0) & (z_edit > 0), + "flip_off": (z_orig > 0) & (z_edit <= 0), + "feature_scale": baseline_stats["feature_scale"], + "all_std": baseline_stats["all_std"], + "baseline_active_rate": baseline_stats["active_rate"], + }, + path, + ) + + +def build_top_feature_rows(summary_rows: list[dict], top_k: int) -> list[dict]: + grouped = {} + for row in summary_rows: + grouped.setdefault((row["group_kind"], row["group_value"]), []).append(row) + top_rows = [] + for (group_kind, group_value), group_rows in sorted(grouped.items()): + ranked = sorted( + group_rows, + key=lambda row: abs(row["mean_delta_scale_units"]), + reverse=True, + )[:top_k] + for rank, row in enumerate(ranked, start=1): + top_row = dict(row) + top_row["rank"] = rank + top_rows.append(top_row) + return top_rows + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--manifest", type=str, required=True) + parser.add_argument("--data_dir", type=str, required=True) + parser.add_argument("--run_root", type=str, required=True) + parser.add_argument("--planner_checkpoint", type=str, default=None) + parser.add_argument("--split", type=str, default="val", choices=["train", "val"]) + parser.add_argument("--sae_block", type=int, default=3) + parser.add_argument("--output_dir", type=str, default=None) + parser.add_argument("--batch_size", type=int, default=8) + parser.add_argument("--encode_batch_size", type=int, default=4096) + parser.add_argument("--device", type=str, default=default_device()) + parser.add_argument("--top_k", type=int, default=50) + parser.add_argument("--index_file", type=str, default=None) + parser.add_argument("--n_items", type=int, default=None) + parser.add_argument("--min_scale", type=float, default=0.05) + parser.add_argument("--allow_name_mismatch", action="store_true") + parser.add_argument("--strict_manifest", action=argparse.BooleanOptionalAction, default=True) + args = parser.parse_args() + + manifest_path = Path(args.manifest) + run_root = Path(args.run_root) + output_dir = default_analysis_dir(run_root, args.sae_block, args.output_dir) + output_dir.mkdir(parents=True, exist_ok=True) + device = torch.device(args.device) + + rows = load_manifest_rows(manifest_path, strict_manifest=args.strict_manifest) + if not rows: + raise ValueError("No edited rows found in manifest.") + + index_file = choose_manifest_value(rows, "index_file", args.index_file) + n_items = choose_manifest_value(rows, "n_items", args.n_items) + dataset = WaymoE2E(indexFile=index_file, data_dir=args.data_dir, n_items=n_items) + + bundle = load_sae_bundle(run_root, args.split, args.sae_block, map_location="cpu") + sae = build_sae_from_checkpoint(bundle["ckpt"], bundle["legacy_norm"]).to(device) + sae.eval() + token_tensor, token_key = resolve_token_tensor(bundle["token_blob"], args.sae_block) + token_name_to_idx = build_token_name_lookup(bundle["token_blob"]) + checkpoint_path = infer_checkpoint_path(run_root, bundle["token_blob"], args.planner_checkpoint) + planner_model, lit_model = load_model(checkpoint_path, device=device) + + print(f"Loaded {len(rows)} edited manifest rows") + print(f"Using token key {token_key}") + print(f"Using planner checkpoint {checkpoint_path}") + print("Computing baseline SAE stats") + baseline_stats = compute_baseline_stats( + sae=sae, + token_tensor=token_tensor, + batch_size=args.encode_batch_size, + device=device, + min_scale=args.min_scale, + ) + + print("Replaying original/edited pairs") + replay = replay_pairs( + rows=rows, + dataset=dataset, + planner_model=planner_model, + lit_model=lit_model, + sae=sae, + sae_block=args.sae_block, + token_blob=bundle["token_blob"], + token_tensor=token_tensor, + token_name_to_idx=token_name_to_idx, + batch_size=args.batch_size, + device=device, + allow_name_mismatch=args.allow_name_mismatch, + ) + + pair_csv = output_dir / f"sae_visual_gen_pairs_block_{args.sae_block}.csv" + feature_delta_pt = output_dir / f"sae_visual_gen_feature_deltas_block_{args.sae_block}.pt" + summary_csv = output_dir / f"sae_visual_gen_feature_summary_block_{args.sae_block}.csv" + top_csv = output_dir / f"sae_visual_gen_top_features_block_{args.sae_block}.csv" + cache_path = output_dir / f"sae_visual_gen_pairs_block_{args.sae_block}.pt" + + write_csv(pair_csv, replay["pair_rows"]) + print(f"Saved pair rows to {pair_csv}") + + print("Saving per-sample feature deltas") + save_feature_deltas_pt( + feature_delta_pt, + rows=rows, + z_orig=replay["z_orig"], + z_edit=replay["z_edit"], + baseline_stats=baseline_stats, + ) + print(f"Saved feature deltas to {feature_delta_pt}") + + print("Computing feature summaries") + summary_rows = build_summary_rows( + rows=rows, + z_orig=replay["z_orig"], + z_edit=replay["z_edit"], + baseline_stats=baseline_stats, + planner_metrics=replay["planner_metrics"], + ) + write_csv(summary_csv, summary_rows) + top_rows = build_top_feature_rows(summary_rows, args.top_k) + top_fieldnames = ["rank"] + [name for name in top_rows[0].keys() if name != "rank"] if top_rows else None + write_csv(top_csv, top_rows, fieldnames=top_fieldnames) + + torch.save( + { + "manifest": str(manifest_path), + "run_root": str(run_root), + "sae_block": args.sae_block, + "token_key": token_key, + "token_path": str(bundle["token_path"]), + "original_source": "token_blob", + "rows": rows, + "pair_rows": replay["pair_rows"], + "tokens_orig": replay["tokens_orig"], + "tokens_edit": replay["tokens_edit"], + "z_orig": replay["z_orig"], + "z_edit": replay["z_edit"], + "planner_metrics": replay["planner_metrics"], + "baseline_stats": baseline_stats, + }, + cache_path, + ) + print(f"Saved feature summary to {summary_csv}") + print(f"Saved top features to {top_csv}") + print(f"Saved tensor cache to {cache_path}") + + +if __name__ == "__main__": + main() diff --git a/src/camera-based-e2e/extract_planner_tok.py b/src/camera-based-e2e/extract_planner_tok.py new file mode 100644 index 0000000..46fc86c --- /dev/null +++ b/src/camera-based-e2e/extract_planner_tok.py @@ -0,0 +1,134 @@ +import argparse +from pathlib import Path + +import torch + +from loader import WaymoE2E +from models.base_model import LitModel, collate_with_images +from models.feature_extractors import SAMFeatures +from models.monocular import DeepMonocularModel +from sae_utils import DRVLA_SAE_VERSION, DRVLA_SOURCE_NOTE, DRVLA_SOURCE_URLS, planner_token_key + + +def load_model(checkpoint_path: str, device: torch.device) -> tuple[DeepMonocularModel, LitModel]: + out_dim = 20 * 2 + model = DeepMonocularModel( + feature_extractor=SAMFeatures( + model_name="timm/vit_pe_spatial_small_patch16_512.fb", frozen=True + ), + out_dim=out_dim, + n_blocks=4, + n_proposals=50, + ) + lit_model = LitModel.load_from_checkpoint( + checkpoint_path, + model=model, + lr=1e-4, + map_location="cpu", + weights_only=False, + ) + model = lit_model.model.to(device) + model.eval() + for param in model.parameters(): + param.requires_grad = False + return model, lit_model + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument("--checkpoint", type=str, required=True) + parser.add_argument("--data_dir", type=str, required=True) + parser.add_argument("--index_file", type=str, required=True) + parser.add_argument("--batch_size", type=int, default=32) + parser.add_argument("--output_path", type=str, required=True) + parser.add_argument("--n_items", type=int, default=None) + parser.add_argument("--nw", type=int, default=0, help="number of workers for dataloader") + parser.add_argument("--device", type=str, default="cuda") + args = parser.parse_args() + + device = torch.device(args.device) + + dataset = WaymoE2E( + indexFile=args.index_file, + data_dir=args.data_dir, + n_items=args.n_items, + ) + loader = torch.utils.data.DataLoader( + dataset, + batch_size=args.batch_size, + num_workers=args.nw, + collate_fn=collate_with_images, + persistent_workers=False, + pin_memory=False, + ) + + model, lit_model = load_model(args.checkpoint, device) + + planner_tok = [] + past = [] + future = [] + intent = [] + names = [] + traj = [] + score = [] + control = [] + block_tokens = None + + with torch.inference_mode(): + for batch in loader: + past.append(batch["PAST"]) + future.append(batch["FUTURE"]) + intent.append(batch["INTENT"]) + names.extend(batch["NAME"]) + + model_inputs = { + "PAST": batch["PAST"].to(device, non_blocking=True), + "IMAGES": lit_model.decode_batch_jpeg(batch["IMAGES_JPEG"], device=device), + "INTENT": batch["INTENT"].to(device, non_blocking=True), + } + out = model(model_inputs, return_block_tokens=True) + planner_tok.append(out["planner_query_tok"].cpu()) + traj.append(out["trajectory"].cpu()) + score.append(out["scores"].cpu()) + control.append(out["controls"].cpu()) + if block_tokens is None: + block_tokens = { + planner_token_key(block_idx): [] + for block_idx in range(model.cfg.n_blocks) + } + for key in block_tokens: + block_tokens[key].append(out[key].cpu()) + + final = { + "planner_query_tok": torch.cat(planner_tok, dim=0), + "past": torch.cat(past, dim=0), + "future": torch.cat(future, dim=0), + "intent": torch.cat(intent, dim=0), + "names": names, + "meta": { + "checkpoint": str(Path(args.checkpoint).resolve()), + "index_file": str(Path(args.index_file).resolve()), + "data_dir": str(Path(args.data_dir).resolve()), + "num_samples": len(names), + "n_items": args.n_items, + "n_blocks": model.cfg.n_blocks, + "sae_version": DRVLA_SAE_VERSION, + "source_note": DRVLA_SOURCE_NOTE, + "source_urls": list(DRVLA_SOURCE_URLS), + "activation_kind": "post_transformer_block_query", + "planner_query_token_keys": [ + planner_token_key(block_idx) for block_idx in range(model.cfg.n_blocks) + ], + }, + "trajectory": torch.cat(traj, dim=0), + "scores": torch.cat(score, dim=0), + "controls": torch.cat(control, dim=0), + } + if block_tokens is not None: + for key, value in block_tokens.items(): + final[key] = torch.cat(value, dim=0) + final["planner_query_tok"] = final[planner_token_key(model.cfg.n_blocks - 1)] + + output_path = Path(args.output_path) + output_path.parent.mkdir(parents=True, exist_ok=True) + torch.save(final, output_path) diff --git a/src/camera-based-e2e/merge_sae_block_analysis.py b/src/camera-based-e2e/merge_sae_block_analysis.py new file mode 100644 index 0000000..e8a0ff7 --- /dev/null +++ b/src/camera-based-e2e/merge_sae_block_analysis.py @@ -0,0 +1,51 @@ +import argparse +import csv +from collections import defaultdict +from pathlib import Path + + +def read_csv_rows(path: Path) -> list[dict]: + with path.open("r", newline="") as f: + return list(csv.DictReader(f)) + + +def write_csv_rows(path: Path, rows: list[dict]) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + if not rows: + return + with path.open("w", newline="") as f: + writer = csv.DictWriter(f, fieldnames=list(rows[0].keys())) + writer.writeheader() + writer.writerows(rows) + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument("--run_root", type=str, required=True) + parser.add_argument("--analysis_dir", type=str, default=None) + args = parser.parse_args() + + run_root = Path(args.run_root) + analysis_root = Path(args.analysis_dir) if args.analysis_dir else run_root / "analysis" + merged_dir = analysis_root / "merged" + + grouped_paths: dict[Path, list[tuple[int, Path]]] = defaultdict(list) + for block_dir in sorted(analysis_root.glob("block_*")): + if not block_dir.is_dir(): + continue + try: + block_idx = int(block_dir.name.split("_")[-1]) + except ValueError: + continue + for csv_path in sorted(block_dir.rglob("*.csv")): + relative_path = csv_path.relative_to(block_dir) + grouped_paths[relative_path].append((block_idx, csv_path)) + + for relative_path, entries in grouped_paths.items(): + merged_rows = [] + for block_idx, csv_path in entries: + for row in read_csv_rows(csv_path): + merged_rows.append({"sae_block": block_idx, **row}) + output_path = merged_dir / relative_path + write_csv_rows(output_path, merged_rows) + print(f"Merged {len(entries)} block files into {output_path}") diff --git a/src/camera-based-e2e/models/base_model.py b/src/camera-based-e2e/models/base_model.py index 17782e6..1caa69d 100644 --- a/src/camera-based-e2e/models/base_model.py +++ b/src/camera-based-e2e/models/base_model.py @@ -26,6 +26,9 @@ def __init__(self, model: nn.Module, lr: float, lr_vision: float | None = 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: @@ -81,12 +84,17 @@ def decode_batch_jpeg( self, images_jpeg: list[list[torch.Tensor]], device: torch.device | None = None, - ) -> list[torch.Tensor]: + ) -> list[torch.Tensor | None]: + model_cfg = getattr(self.model, "cfg", None) + cam_idxs_used = tuple(getattr(model_cfg, "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] + # Flatten cameras flat_encoded, cam_sizes = [], [] - for cam in images_jpeg: - cam_sizes.append(len(cam)) + for cam_idx, cam in selected: + cam_sizes.append((cam_idx, len(cam))) for jpg in cam: t = jpg if isinstance(jpg, torch.Tensor) else torch.frombuffer(memoryview(jpg), dtype=torch.uint8) # decode_jpeg requires the raw jpeg bytes to be on cpu @@ -94,18 +102,30 @@ def decode_batch_jpeg( t = t.cpu() flat_encoded.append(t) - flat_decoded = torchvision.io.decode_jpeg( - flat_encoded, - mode=torchvision.io.ImageReadMode.UNCHANGED, - device=decode_device, - ) # list of (C, H, W) gpu tensors + try: + flat_decoded = torchvision.io.decode_jpeg( + flat_encoded, + mode=torchvision.io.ImageReadMode.UNCHANGED, + device=decode_device, + ) # list of (C, H, W) gpu tensors + except Exception as e: + if "nvJPEG" not in str(e): + raise e + self.has_nvjpeg = False + flat_decoded = torchvision.io.decode_jpeg( + flat_encoded, + mode=torchvision.io.ImageReadMode.UNCHANGED, + device="cpu", + ) + if torch.device(decode_device).type == "cuda": + flat_decoded = [img.to(decode_device, non_blocking=True) for img in flat_decoded] - out = [] + out = [None] * len(images_jpeg) idx = 0 - for n in cam_sizes: - cam_list = flat_decoded[idx: idx+n] + for cam_idx, n in cam_sizes: + out[cam_idx] = torch.stack(flat_decoded[idx:idx+n], dim=0) idx += n - out.append(torch.stack(cam_list, dim=0)) # (B, C, H, W) + return out def on_fit_start(self) -> None: diff --git a/src/camera-based-e2e/models/monocular.py b/src/camera-based-e2e/models/monocular.py index c310027..44dc856 100644 --- a/src/camera-based-e2e/models/monocular.py +++ b/src/camera-based-e2e/models/monocular.py @@ -1,6 +1,7 @@ import torch import torch.nn as nn import torch.nn.functional as F +from dataclasses import dataclass from math import sqrt from .blocks import TransformerBlock @@ -77,6 +78,22 @@ def forward(self, x: dict) -> torch.Tensor: attention = self.attn_norm(attention) return self.decoder(attention.squeeze(1)) # (B, 40) +@dataclass +class DeepMonocularConfig: + # arch + n_blocks: int = 4 + n_proposals: int = 50 + cam_idxs_used: list = (1,) # front only + # kinematics + dt: float = 0.25 + max_accel: float = 8.0 + max_omega: float = 1.0 + # adversarial training + adv_enabled: bool = True + adv_lambda: float = 0.1 + adv_epsilon: float = 0.10 + adv_steps: int = 3 + class DeepMonocularModel(nn.Module): def __init__( self, @@ -89,14 +106,21 @@ def __init__( max_omega: float = 1.0, ): super().__init__() + self.cfg = DeepMonocularConfig( + n_blocks=n_blocks, + n_proposals=n_proposals, + dt=dt, + max_accel=max_accel, + max_omega=max_omega, + ) self.features = feature_extractor self.feature_dim = sum(self.features.dims) if out_dim % 2 != 0: raise ValueError(f"out_dim must be even for (x,y) rollout, got {out_dim}") self.horizon = out_dim // 2 - self.dt = dt - self.max_accel = max_accel - self.max_omega = max_omega + self.dt = self.cfg.dt + self.max_accel = self.cfg.max_accel + self.max_omega = self.cfg.max_omega # Initial Query Projection (Intent + Past -> C) query_input_dim = 3 + 16 * 6 @@ -116,7 +140,7 @@ def __init__( # Deep network rather than single attention in MonocularModel self.blocks = nn.ModuleList([ TransformerBlock(self.feature_dim, num_heads=8, mlp_dim=self.feature_dim*4) - for _ in range(n_blocks) + for _ in range(self.cfg.n_blocks) ]) # For Supervised Depth Loss -> (B, 128, 128) @@ -130,7 +154,7 @@ def __init__( nn.Conv2d(32, 1, 1) ) - self.n_proposals = n_proposals + self.n_proposals = self.cfg.n_proposals self.traj_decoder = nn.Sequential( nn.Linear(self.feature_dim, self.feature_dim), nn.GELU(), @@ -175,57 +199,135 @@ def bicycle_model(self, control_pred: torch.Tensor, past: torch.Tensor) -> torch traj_xy = torch.stack(xy_steps, dim=2) # (B, K, T, 2) return traj_xy, traj_xy.reshape(traj_xy.size(0), -1), accel, omega # (B, K*T*2) - def forward(self, x): - # Copied from MonocularModel - # past: (B, 16, 6), intent: int - past, images, intent = x['PAST'], x['IMAGES'], x['INTENT'] - - # Ref: https://github.com/waymo-research/waymo-open-dataset/blob/5f8a1cd42491210e7de629b6f8fc09b65e0cbe99/src/waymo_open_dataset/dataset.proto#L50%20%20order%20=%20[2,%201,%203] + def prepare_visual_tokens( + self, + images, + ) -> tuple[torch.Tensor, torch.Tensor]: front_cam = images[1] - # Doesn't need no_grad b/c DINO/SAMFeatures will freeze if needed - feats_vit = self.features(front_cam) # list or tensor - + feats_vit = self.features(front_cam) if len(feats_vit) == 1 and isinstance(feats_vit, list): feats_vit = feats_vit[0] - feats = self.visual_adapter(feats_vit) # (B, C, H, W) - - # Depth Supervision - output_depth = F.softplus(self.depth_gen(feats).squeeze(1)) # (B, 128, 128) + feats = self.visual_adapter(feats_vit) + output_depth = F.softplus(self.depth_gen(feats).squeeze(1)) - # tokens: handle list of features or single tensor - # TODO: is this made redundant by if statement above? if isinstance(feats, (list, tuple)): - tokens = torch.cat([f.flatten(2) for f in feats], dim=1) # (B, C_total, N) + tokens = torch.cat([f.flatten(2) for f in feats], dim=1) else: - tokens = feats.flatten(2) # (B, C, N) - tokens = torch.permute(tokens, (0, 2, 1)) + self.positional_encoding # (B, N, C_total) - - # copy procedure to build query_0 from MonocularModel + tokens = feats.flatten(2) + tokens = torch.permute(tokens, (0, 2, 1)) + self.positional_encoding + return tokens, output_depth + + def prepare_initial_query( + self, + past: torch.Tensor, + intent: torch.Tensor, + ) -> torch.Tensor: intent_onehot = F.one_hot((intent - 1).long(), num_classes=3).float() past_flat = past.view(past.size(0), -1) - query: torch.Tensor = self.query_init(torch.cat([intent_onehot, past_flat], dim=1)).unsqueeze(1) + return self.query_init(torch.cat([intent_onehot, past_flat], dim=1)).unsqueeze(1) - for block in self.blocks: - query = block(query, tokens) + def forward_transformer_blocks( + self, + query: torch.Tensor, + tokens: torch.Tensor, + *, + start_block: int = 0, + return_block_outputs: bool = False, + ) -> tuple[torch.Tensor, list[torch.Tensor]]: + block_outputs = [] + for block_idx in range(start_block, len(self.blocks)): + query = self.blocks[block_idx](query, tokens) + if return_block_outputs: + block_outputs.append(query) + return query, block_outputs + + def collect_block_query_tokens( + self, + past: torch.Tensor, + images, + intent: torch.Tensor, + ) -> tuple[list[torch.Tensor], torch.Tensor, torch.Tensor]: + tokens, output_depth = self.prepare_visual_tokens(images) + query = self.prepare_initial_query(past, intent) + _, block_outputs = self.forward_transformer_blocks( + query, + tokens, + return_block_outputs=True, + ) + return block_outputs, tokens, output_depth - # predict (acceleration, angular velocity) for each timestep - # and roll it out using the kinematic bicycle model - control_pred = self.traj_decoder(query.squeeze(1)).view( + def forward_from_planner_query_tok( + self, + planner_query_tok: torch.Tensor, + past: torch.Tensor, + ) -> dict[str, torch.Tensor]: + query = planner_query_tok + if query.ndim == 3: + if query.size(1) != 1: + raise ValueError(f"Expected query shape (B, 1, C), got {query.shape}") + query = query.squeeze(1) + elif query.ndim != 2: + raise ValueError(f"Expected query shape (B, C) or (B, 1, C), got {query.shape}") + + control_pred = self.traj_decoder(query).view( query.size(0), self.n_proposals, self.horizon, 2 ) # (B, K, T, 2) traj_xy, traj_pred, accel, omega = self.bicycle_model(control_pred, past) # (B, K, T*2) traj_pred_flat = traj_xy.reshape(traj_xy.size(0), self.n_proposals, -1) # (B, K, T*2) traj_feat: torch.Tensor = self.traj_features(traj_pred_flat.detach()) # (B, K, C) - query_for_score = query.squeeze(1).detach()[:, torch.newaxis, :].expand(-1, self.n_proposals, -1) # (B, K, C) + query_for_score = query.detach()[:, torch.newaxis, :].expand(-1, self.n_proposals, -1) # (B, K, C) score_in = torch.cat([query_for_score, traj_feat], dim=-1) # (B, K, 2C) score_pred = self.score_decoder(score_in).squeeze(-1) # (B, K) return { "trajectory": traj_pred, "scores": score_pred, - "depth": output_depth, "controls": torch.stack([accel, omega], dim=-1).reshape(query.size(0), -1), + "planner_query_tok": query, } + + def forward_from_block_query_tok( + self, + planner_query_tok: torch.Tensor, + past: torch.Tensor, + tokens: torch.Tensor, + *, + start_block: int, + ) -> dict[str, torch.Tensor]: + query = planner_query_tok + if query.ndim == 2: + query = query.unsqueeze(1) + elif query.ndim != 3: + raise ValueError( + f"Expected planner_query_tok shape (B, C) or (B, 1, C), got {query.shape}" + ) + + if start_block < 0 or start_block >= len(self.blocks): + raise ValueError(f"start_block must be in [0, {len(self.blocks) - 1}], got {start_block}") + + query, _ = self.forward_transformer_blocks( + query, + tokens, + start_block=start_block + 1, + return_block_outputs=False, + ) + return self.forward_from_planner_query_tok(query, past) + + def forward(self, x, return_block_tokens: bool = False): + past, images, intent = x['PAST'], x['IMAGES'], x['INTENT'] + + block_outputs, tokens, output_depth = self.collect_block_query_tokens( + past, + images, + intent, + ) + query = block_outputs[-1] + out = self.forward_from_planner_query_tok(query, past) + out["depth"] = output_depth + if return_block_tokens: + for block_idx, block_query in enumerate(block_outputs): + out[f"planner_query_tok_block_{block_idx}"] = block_query.squeeze(1) + return out diff --git a/src/camera-based-e2e/models/sae.py b/src/camera-based-e2e/models/sae.py new file mode 100644 index 0000000..cf245c1 --- /dev/null +++ b/src/camera-based-e2e/models/sae.py @@ -0,0 +1,279 @@ +from __future__ import annotations + +import math + +import torch +import torch.nn as nn +import torch.nn.functional as F + + +LEGACY_SAE_TYPE = "legacy_relu" +TOPK_AUX_SAE_TYPE = "topk_aux" + + +class SparseAutoencoder(nn.Module): + def __init__( + self, + input_dim: int, + latent_dim: int, + *, + sae_type: str = TOPK_AUX_SAE_TYPE, + k: int | None = None, + k_aux: int = 512, + aux_alpha: float = 1.0 / 32.0, + dead_steps_threshold: int = 500, + use_encoder_bias: bool = False, + ) -> None: + super().__init__() + self.input_dim = input_dim + self.latent_dim = latent_dim + self.sae_type = sae_type + self.k = min(k if k is not None else latent_dim, latent_dim) + self.k_aux = k_aux + self.aux_alpha = aux_alpha + self.dead_steps_threshold = dead_steps_threshold + self.use_encoder_bias = use_encoder_bias + + self.encoder = nn.Linear(input_dim, latent_dim, bias=use_encoder_bias) + self.decoder = nn.Linear(latent_dim, input_dim, bias=False) + self.register_buffer("legacy_mean", torch.zeros(input_dim, dtype=torch.float32)) + self.register_buffer("legacy_std", torch.ones(input_dim, dtype=torch.float32)) + self.register_buffer("legacy_norm_enabled", torch.tensor(False)) + + if self.sae_type == TOPK_AUX_SAE_TYPE: + self.b_pre = nn.Parameter(torch.zeros(input_dim)) + self.register_buffer("cmse", torch.tensor(1.0, dtype=torch.float32)) + self.register_buffer( + "steps_since_active", + torch.zeros(latent_dim, dtype=torch.long), + ) + self.reset_topk_parameters() + elif self.sae_type == LEGACY_SAE_TYPE: + self.register_parameter("b_pre", None) + self.register_buffer("cmse", torch.tensor(1.0, dtype=torch.float32)) + self.register_buffer( + "steps_since_active", + torch.zeros(latent_dim, dtype=torch.long), + ) + self.reset_legacy_parameters() + else: + raise ValueError(f"Unsupported sae_type={sae_type}") + + def reset_topk_parameters(self) -> None: + with torch.no_grad(): + nn.init.normal_(self.decoder.weight, mean=0.0, std=1.0 / math.sqrt(self.input_dim)) + self.normalize_decoder_columns() + scale = math.sqrt(max(self.k, 1) / float(self.input_dim)) + self.encoder.weight.copy_(self.decoder.weight.t() * scale) + if self.encoder.bias is not None: + self.encoder.bias.zero_() + self.b_pre.zero_() + self.cmse.fill_(1.0) + self.steps_since_active.zero_() + + def reset_legacy_parameters(self) -> None: + self.encoder.reset_parameters() + self.decoder.reset_parameters() + if self.encoder.bias is not None: + nn.init.zeros_(self.encoder.bias) + self.cmse.fill_(1.0) + self.steps_since_active.zero_() + + def set_preprocessing_state( + self, + *, + b_pre: torch.Tensor, + cmse: float | torch.Tensor, + ) -> None: + if self.sae_type != TOPK_AUX_SAE_TYPE: + return + with torch.no_grad(): + self.b_pre.copy_(b_pre.to(self.b_pre.device, dtype=self.b_pre.dtype)) + self.cmse.fill_(float(cmse)) + + def set_legacy_normalization( + self, + *, + mean: torch.Tensor, + std: torch.Tensor, + ) -> None: + with torch.no_grad(): + self.legacy_mean.copy_(mean.to(self.legacy_mean.device, dtype=self.legacy_mean.dtype)) + self.legacy_std.copy_(std.to(self.legacy_std.device, dtype=self.legacy_std.dtype)) + self.legacy_norm_enabled.fill_(True) + + def preprocess_input(self, x: torch.Tensor) -> tuple[torch.Tensor, dict[str, torch.Tensor] | None]: + if self.sae_type == LEGACY_SAE_TYPE: + if bool(self.legacy_norm_enabled.item()): + return (x - self.legacy_mean) / self.legacy_std.clamp_min(1e-6), None + return x, None + if self.sae_type != TOPK_AUX_SAE_TYPE: + return x, None + + x_shifted = x - self.b_pre + sample_mean = x_shifted.mean(dim=-1, keepdim=True) + centered = x_shifted - sample_mean + sample_norm = torch.norm(centered, dim=-1, keepdim=True).clamp_min(1e-8) + x_norm = centered / sample_norm + stats = { + "sample_mean": sample_mean, + "sample_norm": sample_norm, + } + return x_norm, stats + + def restore_input(self, x_norm: torch.Tensor, stats: dict[str, torch.Tensor] | None) -> torch.Tensor: + if self.sae_type == LEGACY_SAE_TYPE: + if bool(self.legacy_norm_enabled.item()): + return x_norm * self.legacy_std + self.legacy_mean + return x_norm + if self.sae_type != TOPK_AUX_SAE_TYPE: + return x_norm + if stats is None: + raise ValueError("TopK SAE restore_input requires preprocessing stats.") + return x_norm * stats["sample_norm"] + stats["sample_mean"] + self.b_pre + + def topk_masked_relu(self, values: torch.Tensor, k: int) -> torch.Tensor: + if k <= 0 or values.size(-1) == 0: + return torch.zeros_like(values) + + k_eff = min(k, values.size(-1)) + top_values, top_indices = torch.topk(values, k=k_eff, dim=-1) + masked = torch.zeros_like(values) + masked.scatter_(-1, top_indices, top_values) + return F.relu(masked) + + def encode_preprocessed( + self, + x_preprocessed: torch.Tensor, + *, + return_pre_activations: bool = False, + ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]: + pre_activations = self.encoder(x_preprocessed) + if self.sae_type == LEGACY_SAE_TYPE: + z = F.relu(pre_activations) + else: + z = self.topk_masked_relu(pre_activations, self.k) + if return_pre_activations: + return z, pre_activations + return z + + def encode( + self, + x: torch.Tensor, + *, + return_pre_activations: bool = False, + return_preprocess_stats: bool = False, + ): + x_preprocessed, stats = self.preprocess_input(x) + encoded = self.encode_preprocessed( + x_preprocessed, + return_pre_activations=return_pre_activations, + ) + if return_pre_activations: + z, pre_activations = encoded + if return_preprocess_stats: + return z, pre_activations, stats + return z, pre_activations + + if return_preprocess_stats: + return encoded, stats + return encoded + + def decode_to_preprocessed(self, z: torch.Tensor) -> torch.Tensor: + return self.decoder(z) + + def decode_to_input( + self, + z: torch.Tensor, + *, + reference_x: torch.Tensor | None = None, + preprocess_stats: dict[str, torch.Tensor] | None = None, + ) -> torch.Tensor: + x_preprocessed_hat = self.decode_to_preprocessed(z) + if self.sae_type == LEGACY_SAE_TYPE: + return x_preprocessed_hat + + if preprocess_stats is None: + if reference_x is None: + raise ValueError("TopK SAE decode_to_input requires reference_x or preprocess_stats.") + _, preprocess_stats = self.preprocess_input(reference_x) + return self.restore_input(x_preprocessed_hat, preprocess_stats) + + def decode(self, z: torch.Tensor) -> torch.Tensor: + return self.decode_to_preprocessed(z) + + def forward(self, x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + z, stats = self.encode(x, return_preprocess_stats=True) + x_hat = self.decode_to_input(z, preprocess_stats=stats) + return x_hat, z + + def compute_auxiliary_reconstruction( + self, + residual_preprocessed: torch.Tensor, + dead_mask: torch.Tensor, + ) -> torch.Tensor: + if self.sae_type != TOPK_AUX_SAE_TYPE or not dead_mask.any(): + return torch.zeros_like(residual_preprocessed) + + aux_pre = self.encoder(residual_preprocessed) + masked_aux_pre = torch.full_like(aux_pre, float("-inf")) + masked_aux_pre[:, dead_mask] = aux_pre[:, dead_mask] + aux_z = self.topk_masked_relu(masked_aux_pre, min(self.k_aux, int(dead_mask.sum().item()))) + return self.decode_to_preprocessed(aux_z) + + def compute_loss(self, x: torch.Tensor) -> dict[str, torch.Tensor]: + x_preprocessed, stats = self.preprocess_input(x) + z, pre_activations = self.encode( + x, + return_pre_activations=True, + ) + x_preprocessed_hat = self.decode_to_preprocessed(z) + x_hat = self.decode_to_input(z, preprocess_stats=stats) + + recon_loss = F.mse_loss(x_hat, x) / self.cmse.clamp_min(1e-8) + aux_loss = torch.tensor(0.0, device=x.device, dtype=x.dtype) + + if self.sae_type == TOPK_AUX_SAE_TYPE: + dead_mask = self.steps_since_active >= self.dead_steps_threshold + residual_preprocessed = x_preprocessed - x_preprocessed_hat + if dead_mask.any(): + aux_hat = self.compute_auxiliary_reconstruction( + residual_preprocessed=residual_preprocessed, + dead_mask=dead_mask, + ) + aux_loss = F.mse_loss(aux_hat, residual_preprocessed) / self.cmse.clamp_min(1e-8) + + total_loss = recon_loss + self.aux_alpha * aux_loss + return { + "loss": total_loss, + "recon_loss": recon_loss, + "aux_loss": aux_loss, + "x_hat": x_hat, + "z": z, + "pre_activations": pre_activations, + } + + def update_activation_history(self, z: torch.Tensor) -> None: + with torch.no_grad(): + fired = (z > 0).any(dim=0) + self.steps_since_active.add_(1) + self.steps_since_active[fired] = 0 + + def project_decoder_gradients(self) -> None: + if self.decoder.weight.grad is None or self.sae_type != TOPK_AUX_SAE_TYPE: + return + with torch.no_grad(): + decoder_weight = self.decoder.weight + grad = self.decoder.weight.grad + projection = (grad * decoder_weight).sum(dim=0, keepdim=True) + grad.sub_(decoder_weight * projection) + + def normalize_decoder_columns(self) -> None: + with torch.no_grad(): + norms = torch.norm(self.decoder.weight, dim=0, keepdim=True).clamp_min(1e-8) + self.decoder.weight.div_(norms) + + def dead_fraction(self) -> float: + if self.latent_dim == 0: + return 0.0 + return float((self.steps_since_active >= self.dead_steps_threshold).float().mean().item()) diff --git a/src/camera-based-e2e/sae_utils.py b/src/camera-based-e2e/sae_utils.py new file mode 100644 index 0000000..8108c69 --- /dev/null +++ b/src/camera-based-e2e/sae_utils.py @@ -0,0 +1,211 @@ +from __future__ import annotations + +from pathlib import Path +from typing import TYPE_CHECKING + +import torch + +from models.sae import LEGACY_SAE_TYPE, SparseAutoencoder, TOPK_AUX_SAE_TYPE + +if TYPE_CHECKING: + from loader import WaymoE2E + + +DEFAULT_SAE_BLOCK = 3 +DEFAULT_TOPK_RATIO = 16 +DRVLA_SAE_VERSION = "topk_aux_drvla_v1" +DRVLA_SOURCE_NOTE = "Dr.VLA SAE recipe aligned to Table 4 / Appendix A.2." +DRVLA_SOURCE_URLS = ( + "https://drvla.github.io/", + "https://drvla.github.io/drvla.pdf", +) + + +def planner_token_key(block_idx: int) -> str: + return f"planner_query_tok_block_{block_idx}" + + +def default_device() -> str: + return "cuda" if torch.cuda.is_available() else "cpu" + + +def resolve_topk(input_dim: int, requested_k: int | None = None) -> int: + if requested_k is not None and requested_k > 0: + return max(1, min(requested_k, input_dim)) + return max(1, min(input_dim, round(input_dim / DEFAULT_TOPK_RATIO))) + + +def parse_blocks(text: str, *, default_blocks: list[int] | None = None) -> list[int]: + raw = text.strip().lower() + if raw in {"all", "*"}: + return list(default_blocks if default_blocks is not None else range(DEFAULT_SAE_BLOCK + 1)) + return [int(part.strip()) for part in text.split(",") if part.strip()] + + +def resolve_token_tensor(token_blob: dict, sae_block: int) -> tuple[torch.Tensor, str]: + block_key = planner_token_key(sae_block) + if block_key in token_blob: + return token_blob[block_key].float(), block_key + if sae_block == DEFAULT_SAE_BLOCK and "planner_query_tok" in token_blob: + return token_blob["planner_query_tok"].float(), "planner_query_tok" + raise KeyError( + f"Could not find block token key for sae_block={sae_block}. " + f"Expected {block_key} or planner_query_tok for block {DEFAULT_SAE_BLOCK}." + ) + + +def infer_sae_model_dir(run_root: Path, sae_block: int) -> Path: + block_dir = run_root / "model" / f"block_{sae_block}" + if block_dir.exists(): + return block_dir + if sae_block == DEFAULT_SAE_BLOCK: + return run_root / "model" + return block_dir + + +def infer_sae_paths(run_root: Path, split: str, sae_block: int) -> tuple[Path, Path, Path | None]: + model_dir = infer_sae_model_dir(run_root, sae_block) + ckpt_path = model_dir / "sae_checkpoint.pt" + token_path = run_root / "tokens" / f"planner_tokens_{split}.pt" + legacy_norm_path = model_dir / "sae_normalization.pt" + if legacy_norm_path.exists(): + return ckpt_path, token_path, legacy_norm_path + return ckpt_path, token_path, None + + +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: + model = SparseAutoencoder( + input_dim=ckpt["input_dim"], + latent_dim=ckpt["latent_dim"], + sae_type=TOPK_AUX_SAE_TYPE, + k=ckpt["k"], + k_aux=ckpt["k_aux"], + aux_alpha=ckpt["lambda_aux"], + dead_steps_threshold=ckpt["dead_steps_threshold"], + use_encoder_bias=False, + ) + else: + model = SparseAutoencoder( + input_dim=ckpt["input_dim"], + latent_dim=ckpt["latent_dim"], + sae_type=LEGACY_SAE_TYPE, + use_encoder_bias=True, + ) + model.load_state_dict(ckpt["state_dict"], strict=False) + if legacy_norm is not None: + model.set_legacy_normalization( + mean=legacy_norm["mean"], + std=legacy_norm["std"], + ) + return model + + +def load_sae_bundle( + run_root: Path, + split: str, + sae_block: int, + *, + map_location: str | torch.device = "cpu", +) -> dict: + ckpt_path, token_path, legacy_norm_path = infer_sae_paths(run_root, split, sae_block) + ckpt = torch.load(ckpt_path, map_location=map_location) + token_blob = torch.load(token_path, map_location=map_location) + legacy_norm = torch.load(legacy_norm_path, map_location=map_location) if legacy_norm_path else None + return { + "ckpt": ckpt, + "token_blob": token_blob, + "legacy_norm": legacy_norm, + "ckpt_path": ckpt_path, + "token_path": token_path, + "model_dir": ckpt_path.parent, + } + + +def default_analysis_dir(run_root: Path, sae_block: int, output_dir: str | None) -> Path: + if output_dir is not None: + return Path(output_dir) + return run_root / "analysis" / f"block_{sae_block}" + + +def encode_tensor_batchwise( + sae: SparseAutoencoder, + token_tensor: torch.Tensor, + *, + batch_size: int, + device: torch.device, +) -> torch.Tensor: + latents = [] + sae.eval() + with torch.no_grad(): + for start in range(0, len(token_tensor), batch_size): + batch_x = token_tensor[start : start + batch_size].to(device) + latents.append(sae.encode(batch_x).cpu()) + return torch.cat(latents, dim=0) + + +def decode_latents_to_input( + sae: SparseAutoencoder, + z: torch.Tensor, + reference_x: torch.Tensor, +) -> torch.Tensor: + return sae.decode_to_input(z, reference_x=reference_x) + + +def dataset_from_token_blob( + token_blob: dict, + *, + data_dir: str | None = None, + index_file: str | None = None, +) -> "WaymoE2E": + from loader import WaymoE2E + + meta = token_blob.get("meta", {}) + resolved_data_dir = data_dir or meta.get("data_dir") + resolved_index_file = index_file or meta.get("index_file") + n_items = meta.get("n_items") + if resolved_data_dir is None or resolved_index_file is None: + raise ValueError( + "Need data_dir and index_file to replay planner batches for early-block interventions." + ) + return WaymoE2E( + indexFile=resolved_index_file, + data_dir=resolved_data_dir, + n_items=n_items, + ) + + +def collate_dataset_indices(dataset: WaymoE2E, indices: torch.Tensor) -> dict: + from models.base_model import collate_with_images + + samples = [dataset[int(idx)] for idx in indices.tolist()] + return collate_with_images(samples) + + +def planner_inputs_from_collated_batch( + batch: dict, + *, + lit_model, + device: torch.device, +) -> dict[str, torch.Tensor]: + return { + "PAST": batch["PAST"].to(device, non_blocking=True), + "IMAGES": lit_model.decode_batch_jpeg(batch["IMAGES_JPEG"], device=device), + "INTENT": batch["INTENT"].to(device, non_blocking=True), + } + + +def prepare_replay_context( + planner_model, + lit_model, + batch: dict, + *, + device: torch.device, +) -> dict[str, torch.Tensor]: + model_inputs = planner_inputs_from_collated_batch(batch, lit_model=lit_model, device=device) + tokens, _ = planner_model.prepare_visual_tokens(model_inputs["IMAGES"]) + return { + "past": model_inputs["PAST"], + "tokens": tokens, + } diff --git a/src/camera-based-e2e/train.py b/src/camera-based-e2e/train.py index 3c35363..833f799 100644 --- a/src/camera-based-e2e/train.py +++ b/src/camera-based-e2e/train.py @@ -190,10 +190,10 @@ def n_batches(length: int): from loader import WaymoE2E train_dataset = WaymoE2E( - indexFile="index_train.pkl", data_dir=args.data_dir, n_items=250_000 + indexFile="index_train.pkl", data_dir=args.data_dir, n_items=25_000 ) test_dataset = WaymoE2E( - indexFile="index_val.pkl", data_dir=args.data_dir, n_items=25_000 + indexFile="index_val.pkl", data_dir=args.data_dir, n_items=2_500 ) nw = 0 elif args.dataset == "nuscenes": @@ -313,20 +313,20 @@ def n_batches(length: int): # 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") + #wandb_logger = WandbLogger( + # name=timestamp, + # save_dir=base_path + "/logs", + # project="robotvision", + # log_model=True, + #) + #wandb_logger.watch(lit_model, log="all") 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=[CSVLogger(base_path + "/logs", name=timestamp), wandb_logger], strategy=strategy, use_distributed_sampler=use_distributed_sampler, precision="bf16-mixed" if torch.cuda.is_bf16_supported() else 16, diff --git a/src/camera-based-e2e/train_sae.py b/src/camera-based-e2e/train_sae.py new file mode 100644 index 0000000..6cd3f23 --- /dev/null +++ b/src/camera-based-e2e/train_sae.py @@ -0,0 +1,330 @@ +import argparse +from pathlib import Path + +import torch +from torch.utils.data import DataLoader, TensorDataset + +from models.sae import SparseAutoencoder, TOPK_AUX_SAE_TYPE +from sae_utils import ( + DRVLA_SAE_VERSION, + DRVLA_SOURCE_NOTE, + DRVLA_SOURCE_URLS, + parse_blocks, + resolve_token_tensor, + resolve_topk, +) + + +def compute_centered_mse(x: torch.Tensor) -> float: + centered = x - x.mean(dim=0, keepdim=True) + return float(centered.pow(2).mean().clamp_min(1e-8).item()) + + +def approximate_geometric_median( + x: torch.Tensor, + *, + max_samples: int, + max_iters: int = 100, + tol: float = 1e-5, +) -> torch.Tensor: + if x.size(0) > max_samples: + generator = torch.Generator(device="cpu") + generator.manual_seed(0) + indices = torch.randperm(x.size(0), generator=generator)[:max_samples] + x = x[indices] + + median = x.mean(dim=0) + for _ in range(max_iters): + distances = torch.norm(x - median.unsqueeze(0), dim=1).clamp_min(1e-8) + weights = 1.0 / distances + next_median = (weights[:, None] * x).sum(dim=0) / weights.sum() + if torch.norm(next_median - median).item() < tol: + median = next_median + break + median = next_median + return median + + +def evaluate_epoch( + model: SparseAutoencoder, + loader: DataLoader, + *, + device: torch.device, +) -> dict[str, float]: + totals = { + "loss": 0.0, + "recon_loss": 0.0, + "aux_loss": 0.0, + } + total_items = 0 + + model.eval() + with torch.no_grad(): + for (batch_x,) in loader: + batch_x = batch_x.to(device, non_blocking=True) + loss_dict = model.compute_loss(batch_x) + batch_size = batch_x.size(0) + total_items += batch_size + for key in totals: + totals[key] += float(loss_dict[key].item()) * batch_size + + return {key: value / max(total_items, 1) for key, value in totals.items()} + + +def train_one_block( + *, + block_idx: int, + train_tensor: torch.Tensor, + val_tensor: torch.Tensor, + output_dir: Path, + device: torch.device, + batch_size: int, + lr: float, + beta1: float, + beta2: float, + max_epochs: int, + k: int, + k_aux: int, + lambda_aux: float, + dead_steps_threshold: int, + geometric_median_samples: int, + grad_clip: float, + seed: int, + train_dataset_path: str, + val_dataset_path: str, + token_key: str, + train_blob_meta: dict, + val_blob_meta: dict, +) -> None: + torch.manual_seed(seed + block_idx) + + train_loader = DataLoader( + TensorDataset(train_tensor), + batch_size=batch_size, + shuffle=True, + num_workers=0, + pin_memory=(device.type == "cuda"), + ) + val_loader = DataLoader( + TensorDataset(val_tensor), + batch_size=batch_size, + shuffle=False, + num_workers=0, + pin_memory=(device.type == "cuda"), + ) + + b_pre = approximate_geometric_median( + train_tensor, + max_samples=geometric_median_samples, + ) + cmse = compute_centered_mse(train_tensor) + + model = SparseAutoencoder( + input_dim=train_tensor.shape[1], + latent_dim=train_tensor.shape[1], + sae_type=TOPK_AUX_SAE_TYPE, + k=k, + k_aux=k_aux, + aux_alpha=lambda_aux, + dead_steps_threshold=dead_steps_threshold, + use_encoder_bias=False, + ).to(device) + model.set_preprocessing_state(b_pre=b_pre, cmse=cmse) + + optimizer = torch.optim.Adam(model.parameters(), lr=lr, betas=(beta1, beta2)) + + best_score = float("inf") + best_state = None + history = [] + + for epoch in range(max_epochs): + model.train() + running = { + "loss": 0.0, + "recon_loss": 0.0, + "aux_loss": 0.0, + } + total_items = 0 + + for (batch_x,) in train_loader: + batch_x = batch_x.to(device, non_blocking=True) + loss_dict = model.compute_loss(batch_x) + + optimizer.zero_grad(set_to_none=True) + loss_dict["loss"].backward() + model.project_decoder_gradients() + torch.nn.utils.clip_grad_norm_(model.parameters(), grad_clip) + optimizer.step() + model.normalize_decoder_columns() + model.update_activation_history(loss_dict["z"].detach()) + + batch_size_actual = batch_x.size(0) + total_items += batch_size_actual + for key in running: + running[key] += float(loss_dict[key].item()) * batch_size_actual + + train_metrics = {key: value / max(total_items, 1) for key, value in running.items()} + val_metrics = evaluate_epoch(model, val_loader, device=device) + dead_fraction = model.dead_fraction() + history.append( + { + "epoch": epoch + 1, + "train_loss": train_metrics["loss"], + "train_recon_loss": train_metrics["recon_loss"], + "train_aux_loss": train_metrics["aux_loss"], + "val_loss": val_metrics["loss"], + "val_recon_loss": val_metrics["recon_loss"], + "val_aux_loss": val_metrics["aux_loss"], + "dead_fraction": dead_fraction, + } + ) + + print( + f"[block {block_idx}] epoch {epoch + 1}/{max_epochs} " + f"train={train_metrics['loss']:.6f} " + f"val={val_metrics['loss']:.6f} " + f"dead={dead_fraction:.3f}", + flush=True, + ) + + if val_metrics["loss"] < best_score: + best_score = val_metrics["loss"] + best_state = {key: value.detach().cpu() for key, value in model.state_dict().items()} + + output_dir.mkdir(parents=True, exist_ok=True) + torch.save( + { + "state_dict": best_state, + "best_score": best_score, + "input_dim": train_tensor.shape[1], + "latent_dim": train_tensor.shape[1], + "sae_type": TOPK_AUX_SAE_TYPE, + "block_index": block_idx, + "token_key": token_key, + "expansion_ratio": 1.0, + "k": k, + "k_aux": k_aux, + "lambda_aux": lambda_aux, + "dead_steps_threshold": dead_steps_threshold, + "lr": lr, + "betas": (beta1, beta2), + "batch_size": batch_size, + "max_epochs": max_epochs, + "geometric_median_samples": geometric_median_samples, + "grad_clip": grad_clip, + "seed": seed, + "train_tokens_path": str(Path(train_dataset_path).resolve()), + "val_tokens_path": str(Path(val_dataset_path).resolve()), + "cmse": cmse, + "history": history, + "metadata": { + "sae_version": DRVLA_SAE_VERSION, + "source_note": DRVLA_SOURCE_NOTE, + "source_urls": list(DRVLA_SOURCE_URLS), + "activation_kind": "post_transformer_block_query", + "activation_key": token_key, + "block_index": block_idx, + "resolved_hyperparameters": { + "input_dim": train_tensor.shape[1], + "latent_dim": train_tensor.shape[1], + "expansion_ratio": 1.0, + "k": k, + "k_aux": k_aux, + "lambda_aux": lambda_aux, + "lr": lr, + "beta1": beta1, + "beta2": beta2, + "batch_size": batch_size, + "max_epochs": max_epochs, + "dead_steps_threshold": dead_steps_threshold, + "geometric_median_samples": geometric_median_samples, + "grad_clip": grad_clip, + }, + "train_token_meta": train_blob_meta, + "val_token_meta": val_blob_meta, + }, + }, + output_dir / "sae_checkpoint.pt", + ) + + +def parse_requested_k(raw_k: str) -> int | None: + text = raw_k.strip().lower() + if text in {"", "auto"}: + return None + return int(text) + + +def build_parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser() + parser.add_argument("--output_dir", type=str, required=True) + parser.add_argument("--train_dataset", type=str, required=True) + parser.add_argument("--val_dataset", type=str, required=True) + parser.add_argument("--blocks", type=str, default="all") + parser.add_argument("--batch_size", type=int, default=4096) + parser.add_argument("--lr", type=float, default=1e-4) + parser.add_argument("--beta1", type=float, default=0.9) + parser.add_argument("--beta2", type=float, default=0.999) + parser.add_argument("--max_epochs", type=int, default=100) + parser.add_argument("--device", type=str, default="cuda" if torch.cuda.is_available() else "cpu") + parser.add_argument("--seed", type=int, default=42) + parser.add_argument("--k", type=str, default="auto") + parser.add_argument("--k_aux", type=int, default=512) + parser.add_argument("--lambda_aux", type=float, default=1.0 / 32.0) + parser.add_argument("--dead_steps_threshold", type=int, default=500) + parser.add_argument("--geometric_median_samples", type=int, default=10000) + parser.add_argument("--grad_clip", type=float, default=1.0) + return parser + + +def main(argv: list[str] | None = None) -> None: + parser = build_parser() + args = parser.parse_args(argv) + + device = torch.device(args.device) + output_dir = Path(args.output_dir) + + train_blob = torch.load(args.train_dataset, map_location="cpu") + val_blob = torch.load(args.val_dataset, map_location="cpu") + n_blocks = int(train_blob.get("meta", {}).get("n_blocks", 4)) + blocks = parse_blocks(args.blocks, default_blocks=list(range(n_blocks))) + requested_k = parse_requested_k(args.k) + + for block_idx in blocks: + train_tensor, token_key = resolve_token_tensor(train_blob, block_idx) + val_tensor, _ = resolve_token_tensor(val_blob, block_idx) + k = resolve_topk(train_tensor.shape[1], requested_k=requested_k) + block_output_dir = output_dir / f"block_{block_idx}" + print( + f"Training SAE for block {block_idx} using {token_key}: " + f"train={tuple(train_tensor.shape)} val={tuple(val_tensor.shape)} k={k}", + flush=True, + ) + train_one_block( + block_idx=block_idx, + train_tensor=train_tensor, + val_tensor=val_tensor, + output_dir=block_output_dir, + device=device, + batch_size=args.batch_size, + lr=args.lr, + beta1=args.beta1, + beta2=args.beta2, + max_epochs=args.max_epochs, + k=k, + k_aux=args.k_aux, + lambda_aux=args.lambda_aux, + dead_steps_threshold=args.dead_steps_threshold, + geometric_median_samples=args.geometric_median_samples, + grad_clip=args.grad_clip, + seed=args.seed, + train_dataset_path=args.train_dataset, + val_dataset_path=args.val_dataset, + token_key=token_key, + train_blob_meta=dict(train_blob.get("meta", {})), + val_blob_meta=dict(val_blob.get("meta", {})), + ) + + +if __name__ == "__main__": + main()