diff --git a/config/dataset/qwen2.5-vl-3B-nuplan-navtest-nocot.yaml b/config/dataset/qwen2.5-vl-3B-nuplan-navtest-nocot.yaml new file mode 100644 index 0000000..d710890 --- /dev/null +++ b/config/dataset/qwen2.5-vl-3B-nuplan-navtest-nocot.yaml @@ -0,0 +1,14 @@ +name: qwen2.5-vl-3B-nuplan-navtest-nocot +description: Preprocessing dataset for AutoVLA (No-CoT) on the full navtest split, reusing ishaan.rawal's NFS-mirrored nuPlan/navsim data + +# model (only used to load the AutoProcessor/tokenizer, no CoT inference is run) +pretrained_model_path: /media/training_data/jaagat-prashar/navsim_autovla_eval/models/Qwen2.5-VL-3B-Instruct + +# training +batch_size: 1 +num_workers: 32 + +# dataset (the split should match the dataset path) +dataset_name: nuplan +dataset_path: /media/training_data/ishaan.rawal/navsim/dataset/placeholder/test +scene_filter: ./navsim/navsim/planning/script/config/common/train_test_split/scene_filter/navtest.yaml diff --git a/config/training/qwen2.5-vl-3B-nuplan-navtest-eval.yaml b/config/training/qwen2.5-vl-3B-nuplan-navtest-eval.yaml new file mode 100644 index 0000000..551afdd --- /dev/null +++ b/config/training/qwen2.5-vl-3B-nuplan-navtest-eval.yaml @@ -0,0 +1,46 @@ +name: qwen2.5-vl-3B-nuplan-navtest-eval +description: AutoVLA (Qwen2.5-VL-3B) navtest PDM-Score evaluation config, pointed at the HF AutoVLA_PDMS_89 checkpoint and NFS-mirrored navtest data + +model: + use_cot: true + pretrained_model_path: /media/training_data/jaagat-prashar/navsim_autovla_eval/models/Qwen2.5-VL-3B-Instruct + sft_model_path: /media/training_data/jaagat-prashar/navsim_autovla_eval/checkpoints/AutoVLA_PDMS_89.ckpt + codebook_cache_path: "codebook_cache/agent_vocab.pkl" + train_vision_backbone: false + train_lm_backbone: true + lora: + use: true + task_type: CAUSAL_LM + target_modules: ["q_proj", "v_proj", "k_proj", "o_proj"] + r: 8 + alpha: 8 + dropout: 0.1 + bias: "none" + trajectory: # future trajectory parameters + num_poses: 10 + interval_length: 0.5 + time_horizon: 5.0 + tokens: # action token parameters, depends on the pretrained large model + action_start_id: 151665 + ignore_index: -100 + assistant_id: [151644, 77091] + video: + min_pixels: 109760 + max_pixels: 109760 + +data: + val: + path: /media/training_data/ishaan.rawal/navsim/dataset/placeholder/test + scene_filter: ./navsim/navsim/planning/script/config/common/train_test_split/scene_filter/navtest.yaml + metric_cache_path: /media/training_data/jaagat-prashar/navsim_autovla_eval/dataset/nuplan/navtest_metric_cache + json_dataset_path: /media/training_data/jaagat-prashar/navsim_autovla_eval/dataset/nuplan/navtest_nocot + sensor_data_path: /media/training_data/ishaan.rawal/navsim/dataset/sensor_blobs/test + +inference: + batch_size: 1 + num_workers: 1 + sample: + max_length: 2048 + temperature: 0.01 + top_k: 0.0 + top_p: 1.0 diff --git a/navsim/navsim/planning/script/run_pdm_score_cot.py b/navsim/navsim/planning/script/run_pdm_score_cot.py index 1137fe1..d7094fa 100644 --- a/navsim/navsim/planning/script/run_pdm_score_cot.py +++ b/navsim/navsim/planning/script/run_pdm_score_cot.py @@ -89,18 +89,19 @@ def run_pdm_score(args: List[Dict[str, Union[List[str], DictConfig]]]) -> List[D else: trajectory, cot_results = agent.compute_trajectory(agent_input) - # scene = scene_loader.get_scene_from_token(token) - # frame_idx = scene.scene_metadata.num_history_frames - 1 - # fig, _ = plot_cameras_frame_with_bev_agent_cot(scene, frame_idx, agent_trajectory=trajectory, cot=cot_results) - # vis_dir = Path(cfg.output_dir) / "Visualization" - # vis_dir.mkdir(parents=True, exist_ok=True) - # vis_path = vis_dir / f"{token}_bevagent.png" - # fig.savefig(vis_path, bbox_inches="tight") - # plt.close(fig) - # if cot_results: - # cot_md_path = vis_dir / f"{token}_cot.md" - # with open(cot_md_path, "w", encoding="utf-8") as f: - # f.write(cot_results.strip() + "\n") + if cfg.get("save_visualization", False): + scene = scene_loader.get_scene_from_token(token) + frame_idx = scene.scene_metadata.num_history_frames - 1 + fig, _ = plot_cameras_frame_with_bev_agent_cot(scene, frame_idx, agent_trajectory=trajectory, cot=cot_results) + vis_dir = Path(cfg.output_dir) / "Visualization" + vis_dir.mkdir(parents=True, exist_ok=True) + vis_path = vis_dir / f"{token}_bevagent.png" + fig.savefig(vis_path, bbox_inches="tight") + plt.close(fig) + if cot_results: + cot_md_path = vis_dir / f"{token}_cot.md" + with open(cot_md_path, "w", encoding="utf-8") as f: + f.write(cot_results.strip() + "\n") pdm_result = pdm_score( metric_cache=metric_cache, diff --git a/navsim/scripts/evaluation/run_autovla_agent_pdm_score_evaluation_navtest_jaagat.sh b/navsim/scripts/evaluation/run_autovla_agent_pdm_score_evaluation_navtest_jaagat.sh new file mode 100755 index 0000000..93e1304 --- /dev/null +++ b/navsim/scripts/evaluation/run_autovla_agent_pdm_score_evaluation_navtest_jaagat.sh @@ -0,0 +1,37 @@ +#!/bin/bash +# navtest PDM-Score eval of the AutoVLA_PDMS_89 HF checkpoint. +# Data reused from ishaan.rawal's NFS mirror (maps/logs/sensor_blobs); metric cache +# and Qwen2.5-VL-3B base model pulled fresh into /media/training_data/jaagat-prashar. +set -euo pipefail + +export TOKENIZERS_PARALLELISM=false + +BASE=/media/training_data/jaagat-prashar/navsim_autovla_eval +OPENSCENE_ROOT=/media/training_data/ishaan.rawal/navsim/dataset + +export NUPLAN_MAP_VERSION="nuplan-maps-v1.0" +export NUPLAN_MAPS_ROOT="$OPENSCENE_ROOT/maps" +export NAVSIM_EXP_ROOT="$BASE/exp" +export NAVSIM_DEVKIT_ROOT="/home/jaagat-prashar/workspace/research-project-template-main/autovla/AutoVLA/navsim" +export OPENSCENE_DATA_ROOT="$OPENSCENE_ROOT" + +export PYTHONPATH="./navsim:${PYTHONPATH:-}" + +TRAIN_TEST_SPLIT=navtest +CHECKPOINT="$BASE/checkpoints/AutoVLA_PDMS_89.ckpt" +CACHE_PATH="$BASE/dataset/nuplan/navtest_metric_cache" +JSON_DATA_PATH="$BASE/dataset/nuplan/navtest_nocot" +SENSOR_DATA_PATH="$OPENSCENE_ROOT/sensor_blobs/test" +CONFIG_PATH="./config/training/qwen2.5-vl-3B-nuplan-navtest-eval.yaml" +LORA=false + +CUDA_VISIBLE_DEVICES=1 python $NAVSIM_DEVKIT_ROOT/navsim/planning/script/run_pdm_score_cot.py \ + train_test_split=$TRAIN_TEST_SPLIT \ + agent=autovla_agent \ + +agent.config_path="$CONFIG_PATH" \ + +agent.checkpoint_path="$CHECKPOINT" \ + +agent.sensor_data_path="$SENSOR_DATA_PATH" \ + +agent.lora_conf.use_lora=$LORA \ + metric_cache_path=$CACHE_PATH \ + json_data_path=$JSON_DATA_PATH \ + experiment_name=autovla_agent_navtest_jaagat diff --git a/navsim/scripts/evaluation/run_autovla_agent_pdm_score_evaluation_navtest_jaagat_smoke.sh b/navsim/scripts/evaluation/run_autovla_agent_pdm_score_evaluation_navtest_jaagat_smoke.sh new file mode 100755 index 0000000..cc13feb --- /dev/null +++ b/navsim/scripts/evaluation/run_autovla_agent_pdm_score_evaluation_navtest_jaagat_smoke.sh @@ -0,0 +1,37 @@ +#!/bin/bash +# Smoke test: same as run_autovla_agent_pdm_score_evaluation_navtest_jaagat.sh but +# capped to a small number of scenes, to validate the pipeline before a full navtest run. +set -euo pipefail + +export TOKENIZERS_PARALLELISM=false + +BASE=/media/training_data/jaagat-prashar/navsim_autovla_eval +OPENSCENE_ROOT=/media/training_data/ishaan.rawal/navsim/dataset + +export NUPLAN_MAP_VERSION="nuplan-maps-v1.0" +export NUPLAN_MAPS_ROOT="$OPENSCENE_ROOT/maps" +export NAVSIM_EXP_ROOT="$BASE/exp" +export NAVSIM_DEVKIT_ROOT="/home/jaagat-prashar/workspace/research-project-template-main/autovla/AutoVLA/navsim" +export OPENSCENE_DATA_ROOT="$OPENSCENE_ROOT" + +export PYTHONPATH="./navsim:${PYTHONPATH:-}" + +TRAIN_TEST_SPLIT=navtest +CHECKPOINT="$BASE/checkpoints/AutoVLA_PDMS_89.ckpt" +CACHE_PATH="$BASE/dataset/nuplan/navtest_metric_cache" +JSON_DATA_PATH="$BASE/dataset/nuplan/navtest_nocot" +SENSOR_DATA_PATH="$OPENSCENE_ROOT/sensor_blobs/test" +CONFIG_PATH="./config/training/qwen2.5-vl-3B-nuplan-navtest-eval.yaml" +LORA=false + +CUDA_VISIBLE_DEVICES=1 python $NAVSIM_DEVKIT_ROOT/navsim/planning/script/run_pdm_score_cot.py \ + train_test_split=$TRAIN_TEST_SPLIT \ + train_test_split.scene_filter.max_scenes=30 \ + agent=autovla_agent \ + +agent.config_path="$CONFIG_PATH" \ + +agent.checkpoint_path="$CHECKPOINT" \ + +agent.sensor_data_path="$SENSOR_DATA_PATH" \ + +agent.lora_conf.use_lora=$LORA \ + metric_cache_path=$CACHE_PATH \ + json_data_path=$JSON_DATA_PATH \ + experiment_name=autovla_agent_navtest_jaagat_smoke diff --git a/navsim/scripts/evaluation/run_autovla_agent_pdm_score_evaluation_navtest_jaagat_visuals.sh b/navsim/scripts/evaluation/run_autovla_agent_pdm_score_evaluation_navtest_jaagat_visuals.sh new file mode 100755 index 0000000..733658e --- /dev/null +++ b/navsim/scripts/evaluation/run_autovla_agent_pdm_score_evaluation_navtest_jaagat_visuals.sh @@ -0,0 +1,38 @@ +#!/bin/bash +# Small visualization pass: same setup as the smoke test but renders the +# per-scene camera+BEV+trajectory figure for a handful of scenes. +set -euo pipefail + +export TOKENIZERS_PARALLELISM=false + +BASE=/media/training_data/jaagat-prashar/navsim_autovla_eval +OPENSCENE_ROOT=/media/training_data/ishaan.rawal/navsim/dataset + +export NUPLAN_MAP_VERSION="nuplan-maps-v1.0" +export NUPLAN_MAPS_ROOT="$OPENSCENE_ROOT/maps" +export NAVSIM_EXP_ROOT="$BASE/exp" +export NAVSIM_DEVKIT_ROOT="/home/jaagat-prashar/workspace/research-project-template-main/autovla/AutoVLA/navsim" +export OPENSCENE_DATA_ROOT="$OPENSCENE_ROOT" + +export PYTHONPATH="./navsim:${PYTHONPATH:-}" + +TRAIN_TEST_SPLIT=navtest +CHECKPOINT="$BASE/checkpoints/AutoVLA_PDMS_89.ckpt" +CACHE_PATH="$BASE/dataset/nuplan/navtest_metric_cache" +JSON_DATA_PATH="$BASE/dataset/nuplan/navtest_nocot" +SENSOR_DATA_PATH="$OPENSCENE_ROOT/sensor_blobs/test" +CONFIG_PATH="./config/training/qwen2.5-vl-3B-nuplan-navtest-eval.yaml" +LORA=false + +CUDA_VISIBLE_DEVICES=1 python $NAVSIM_DEVKIT_ROOT/navsim/planning/script/run_pdm_score_cot.py \ + train_test_split=$TRAIN_TEST_SPLIT \ + train_test_split.scene_filter.max_scenes=6 \ + +save_visualization=true \ + agent=autovla_agent \ + +agent.config_path="$CONFIG_PATH" \ + +agent.checkpoint_path="$CHECKPOINT" \ + +agent.sensor_data_path="$SENSOR_DATA_PATH" \ + +agent.lora_conf.use_lora=$LORA \ + metric_cache_path=$CACHE_PATH \ + json_data_path=$JSON_DATA_PATH \ + experiment_name=autovla_agent_navtest_jaagat_visuals diff --git a/tools/ablation/cot_ablation.py b/tools/ablation/cot_ablation.py new file mode 100644 index 0000000..c086f00 --- /dev/null +++ b/tools/ablation/cot_ablation.py @@ -0,0 +1,461 @@ +""" +cot_ablation.py — Text-level CoT ablation probes for AutoVLA. + +AutoVLA's chain-of-thought is plain generated text sharing one autoregressive +stream with its physical action tokens (see models/autovla.py:AutoVLA.get_prompt +and models/action_tokenizer.py), not a separately-addressable set of hidden +states. That's different from the masking/ subsystem's MaskedAlpamayo1_5, which +exposes custom primitives (compare_conditions, salience_leave_one_word_out) that +mask specific attention columns in a single forward pass across several +conditions at once. AutoVLA has no equivalent, so the ablations here work at the +text/token level instead: + + 1. Generate one baseline rollout per scene and split its completion into a + reasoning-text span and an action-token span (the point where token ids + first cross `action_start_id`). + 2. Edit the reasoning text for a given condition (mask concept words, keep + only a prefix/suffix, or inject a plausible mistake). + 3. Teacher-force the edited text back in as the assistant's "reasoning so + far" (tokenized and appended after the prompt) and let the model continue + generating autoregressively -- its own words plus fresh action tokens -- + from that point. + 4. Decode the resulting action tokens to a trajectory and compare it to the + baseline trajectory. + +This measures how much the *decoded trajectory* actually moves in response to +what the model is shown as its own prior reasoning -- a faithfulness probe, not +just a text-quality check. + +Conditions (adapted from masking/training/run.py's experiments a-d): + no_cot -- experiment A: reasoning vs. no-reasoning, by toggling + `use_cot` and regenerating the full response from scratch. + concept_mask -- experiment B: strip concept-relevant words (e.g. + "pedestrian", "stop", "red") from the reasoning text. + prefix_{n}w / + suffix_{n}w -- experiment C: keep only the first/last n words of the + reasoning text, sweeping n over --threshold_words. + injected_mistake -- in the spirit of experiment D's clause-reversal probe: + swap decision-relevant words for plausible-but-wrong + opposites (stop<->accelerate, red<->green, left<->right, + pedestrian->no pedestrian, ...) and see whether the + trajectory follows the injected error. + +Metrics per condition, relative to the baseline trajectory (mirrors masking's +ade_m/endpoint_m/curvature/accel fields, computed here via simple finite +differences since AutoVLA's action_tokenizer doesn't expose a `controls` dict): + ade_m, endpoint_m, delta_xy_per_waypoint, d_curvature_mean, d_accel_mean + +This is intentionally a single-process CLI script, matching tools/eval/ +nusc_eval.py's loading/eval convention -- it is not a multi-GPU/Lilypad +launcher. Requires an SFT (or CoT-capable) checkpoint and a preprocessed +val split, same prerequisites as nusc_eval.py. + +Usage: + python tools/ablation/cot_ablation.py \ + --config config/training/qwen2.5-vl-3B-mix-sft.yaml \ + --checkpoint /path/to/sft_checkpoint.ckpt \ + --num_samples 50 \ + --output ablation_results.jsonl +""" +import argparse +import json +import logging +import re +import string +import sys +from pathlib import Path +from typing import Any, Dict, List, Optional, Tuple + +PROJECT_ROOT = Path(__file__).resolve().parent.parent.parent +sys.path.insert(0, str(PROJECT_ROOT)) +sys.path.insert(0, str(PROJECT_ROOT / "navsim")) + +import numpy as np +import torch +import yaml +from tqdm import tqdm +from transformers import AutoProcessor + +from dataset_utils.sft_dataset import SFTDataset +from models.autovla import SFTAutoVLA, AutoVLA + +logger = logging.getLogger(__name__) + +DEFAULT_CONCEPTS = "pedestrian,person,cyclist,crosswalk,vehicle,stop,red,light" +DEFAULT_THRESHOLD_WORDS = "0,5,10,20,30,50" + +# Word/phrase swaps used by the injected_mistake condition: each key, if found +# in the reasoning text (case-insensitive, whole word/phrase), is replaced by +# its value to inject a plausible-but-wrong claim or decision. +MISTAKE_SWAPS: Dict[str, str] = { + "turn left": "turn right", + "turn right": "turn left", + "change lane to left": "change lane to right", + "change lane to right": "change lane to left", + "quick acceleration": "quick deceleration", + "quick deceleration": "quick acceleration", + "acceleration": "deceleration", + "deceleration": "acceleration", + "stop": "accelerate", + "red": "green", + "green": "red", + "pedestrian": "no pedestrian", + "crossing": "clear", +} +_MISTAKE_PATTERN = re.compile( + r"\b(" + "|".join(re.escape(k) for k in sorted(MISTAKE_SWAPS, key=len, reverse=True)) + r")\b", + re.IGNORECASE, +) + + +def load_config(path: str) -> dict: + with open(path, "r") as f: + return yaml.safe_load(f) + + +def parse_args(): + parser = argparse.ArgumentParser(description="CoT ablation probes for AutoVLA") + parser.add_argument("--config", type=str, required=True, help="Path to the training/eval config file") + parser.add_argument("--checkpoint", type=str, required=True, help="Path to the model checkpoint") + parser.add_argument("--output", type=str, default="ablation_results.jsonl") + parser.add_argument("--device", type=str, default="cuda:0") + parser.add_argument("--num_samples", type=int, default=None, help="Number of val scenes to run (default: all)") + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--concepts", type=str, default=DEFAULT_CONCEPTS, + help="Comma-separated concept words for the concept_mask condition") + parser.add_argument("--threshold_words", type=str, default=DEFAULT_THRESHOLD_WORDS, + help="Comma-separated word-count thresholds for the prefix/suffix sweep") + parser.add_argument("--max_new_tokens", type=int, default=128, + help="Continuation budget for teacher-forced conditions") + parser.add_argument("--verbose", action="store_true") + return parser.parse_args() + + +# --------------------------------------------------------------------------- +# Text-editing helpers for each condition +# --------------------------------------------------------------------------- + +def concept_mask(text: str, concepts: List[str]) -> Tuple[str, int]: + """Drop any word that starts with one of `concepts` (case-insensitive, + so plural/inflected forms like "pedestrians" are also caught).""" + concept_list = [c.strip().lower() for c in concepts if c.strip()] + + def is_concept(word: str) -> bool: + bare = word.strip(string.punctuation).lower() + return any(bare.startswith(c) for c in concept_list if bare) + + words = text.split() + kept = [w for w in words if not is_concept(w)] + return " ".join(kept), len(words) - len(kept) + + +def prefix_truncate(text: str, n: int) -> str: + return " ".join(text.split()[:n]) + + +def suffix_truncate(text: str, n: int) -> str: + if n <= 0: + return "" + return " ".join(text.split()[-n:]) + + +def inject_mistakes(text: str) -> Tuple[str, List[str]]: + fired: List[str] = [] + + def _sub(m: "re.Match") -> str: + key = m.group(0).lower() + fired.append(key) + return MISTAKE_SWAPS[key] + + edited = _MISTAKE_PATTERN.sub(_sub, text) + return edited, fired + + +# --------------------------------------------------------------------------- +# Trajectory metrics +# --------------------------------------------------------------------------- + +def trajectory_deltas(baseline: np.ndarray, other: np.ndarray) -> dict: + T = min(len(baseline), len(other)) + if T == 0: + return {"ade_m": None, "endpoint_m": None, "delta_xy_per_waypoint": []} + delta_xy = np.linalg.norm(other[:T, :2] - baseline[:T, :2], axis=-1) + return { + "ade_m": float(delta_xy.mean()), + "endpoint_m": float(delta_xy[-1]), + "delta_xy_per_waypoint": delta_xy.round(4).tolist(), + } + + +def _heading_rate(traj: np.ndarray, dt: float) -> np.ndarray: + dh = np.diff(traj[:, 2]) + dh = (dh + np.pi) % (2 * np.pi) - np.pi + return dh / dt + + +def _speed(traj: np.ndarray, dt: float) -> np.ndarray: + return np.linalg.norm(np.diff(traj[:, :2], axis=0), axis=-1) / dt + + +def control_deltas(baseline: np.ndarray, other: np.ndarray, dt: float) -> dict: + if len(baseline) < 2 or len(other) < 2: + return {"d_curvature_mean": None, "d_accel_mean": None} + hb, ho = _heading_rate(baseline, dt), _heading_rate(other, dt) + sb, so = _speed(baseline, dt), _speed(other, dt) + Th = min(len(hb), len(ho)) + d_curv = np.abs(ho[:Th] - hb[:Th]) if Th > 0 else np.array([0.0]) + + ab = np.diff(sb) / dt if len(sb) > 1 else np.array([]) + ao = np.diff(so) / dt if len(so) > 1 else np.array([]) + Ta = min(len(ab), len(ao)) + d_accel = np.abs(ao[:Ta] - ab[:Ta]) if Ta > 0 else np.array([0.0]) + + return { + "d_curvature_mean": float(d_curv.mean()), + "d_accel_mean": float(d_accel.mean()), + } + + +# --------------------------------------------------------------------------- +# Generation helpers +# --------------------------------------------------------------------------- + +def _to_device(inputs, device: str) -> Dict[str, torch.Tensor]: + return {k: v.to(device) for k, v in inputs.items() if isinstance(v, torch.Tensor)} + + +def generate_full(autovla: AutoVLA, input_features: dict, seed: int, device: str) -> Optional[dict]: + """Full from-scratch generation (prompt -> reasoning text + action tokens), + following the same call pattern as AutoVLA.predict(), but also returning the + raw reasoning/action token split so callers can edit the reasoning span.""" + prompt_inputs = autovla.get_prompt(input_features) + model_inputs = _to_device(prompt_inputs, device) + + torch.manual_seed(seed) + with torch.no_grad(): + generated = autovla.vlm.generate( + **model_inputs, + max_length=autovla.gen_conf["max_length"], + do_sample=True, + temperature=autovla.gen_conf["temperature"], + top_k=autovla.gen_conf["top_k"], + top_p=autovla.gen_conf["top_p"], + ) + + prompt_len = model_inputs["input_ids"].shape[1] + completion = generated[0, prompt_len:][:-1].cpu() # drop trailing eos, mirrors AutoVLA.predict() + + action_mask = completion >= autovla.action_start_id + if not action_mask.any(): + return None + action_start_idx = int(action_mask.nonzero()[0].item()) + reasoning_ids = completion[:action_start_idx] + action_ids = completion[action_start_idx:] + + trajectory = autovla.action_tokenizer.decode_token_ids_to_trajectory(action_ids) + if len(trajectory) == 0: + return None + + return { + "model_inputs": model_inputs, + "reasoning_text": autovla.processor.decode(reasoning_ids).strip(), + "trajectory": trajectory[0, 1:].numpy(), + } + + +def continue_from_text( + autovla: AutoVLA, + prompt_model_inputs: Dict[str, torch.Tensor], + edited_text: str, + max_new_tokens: int, + seed: int, +) -> Optional[np.ndarray]: + """Teacher-force `edited_text` as the assistant's reasoning-so-far by + appending its tokens after the prompt, then let the model continue + autoregressively (fresh words + action tokens) from there.""" + if edited_text.strip(): + edited_ids = autovla.processor.tokenizer( + edited_text, add_special_tokens=False, return_tensors="pt" + ).input_ids.to(prompt_model_inputs["input_ids"].device) + else: + edited_ids = torch.empty( + (1, 0), dtype=prompt_model_inputs["input_ids"].dtype, + device=prompt_model_inputs["input_ids"].device, + ) + + forced_input_ids = torch.cat([prompt_model_inputs["input_ids"], edited_ids], dim=1) + forced_attention_mask = torch.cat( + [prompt_model_inputs["attention_mask"], torch.ones_like(edited_ids)], dim=1 + ) + + torch.manual_seed(seed) + with torch.no_grad(): + generated = autovla.vlm.generate( + input_ids=forced_input_ids, + attention_mask=forced_attention_mask, + pixel_values_videos=prompt_model_inputs["pixel_values_videos"], + video_grid_thw=prompt_model_inputs["video_grid_thw"], + max_new_tokens=max_new_tokens, + do_sample=True, + temperature=autovla.gen_conf["temperature"], + top_k=autovla.gen_conf["top_k"], + top_p=autovla.gen_conf["top_p"], + ) + + continuation = generated[0, forced_input_ids.shape[1]:].cpu() + action_ids = continuation[continuation >= autovla.action_start_id] + if len(action_ids) == 0: + return None + + trajectory = autovla.action_tokenizer.decode_token_ids_to_trajectory(action_ids) + if len(trajectory) == 0: + return None + return trajectory[0, 1:].numpy() + + +# --------------------------------------------------------------------------- +# Per-scene driver +# --------------------------------------------------------------------------- + +def run_scene( + model: SFTAutoVLA, + input_features: dict, + concepts: List[str], + threshold_words: List[int], + max_new_tokens: int, + seed: int, + device: str, +) -> Optional[dict]: + autovla = model.autovla + dt = model.cfg["model"]["trajectory"]["interval_length"] + + baseline = generate_full(autovla, input_features, seed, device) + if baseline is None: + return None + + result: Dict[str, Any] = { + "cot": baseline["reasoning_text"], + "n_words_total": len(baseline["reasoning_text"].split()), + "traj_baseline_xy": baseline["trajectory"][:, :2].round(4).tolist(), + "conditions": {}, + } + + # no_cot: full independent regeneration with use_cot toggled off (experiment A) + autovla.use_cot = False + no_cot = generate_full(autovla, input_features, seed, device) + autovla.use_cot = True + if no_cot is not None: + result["conditions"]["no_cot"] = { + **trajectory_deltas(baseline["trajectory"], no_cot["trajectory"]), + **control_deltas(baseline["trajectory"], no_cot["trajectory"], dt), + } + + reasoning_text = baseline["reasoning_text"] + prompt_model_inputs = baseline["model_inputs"] + + conditions_to_run: List[Tuple[str, str, dict]] = [] + + edited, n_masked = concept_mask(reasoning_text, concepts) + conditions_to_run.append(("concept_mask", edited, {"concepts": concepts, "n_words_masked": n_masked})) + + for n in sorted(threshold_words): + conditions_to_run.append((f"prefix_{n}w", prefix_truncate(reasoning_text, n), {"n": n})) + conditions_to_run.append((f"suffix_{n}w", suffix_truncate(reasoning_text, n), {"n": n})) + + mistake_text, fired = inject_mistakes(reasoning_text) + if fired: + conditions_to_run.append(("injected_mistake", mistake_text, {"rules_fired": fired})) + else: + result["injected_mistake_skipped"] = "no mistake-substitution rule matched this scene's reasoning text" + + for name, edited_text, meta in conditions_to_run: + traj = continue_from_text(autovla, prompt_model_inputs, edited_text, max_new_tokens, seed) + if traj is None: + result["conditions"][name] = {**meta, "failed": True} + continue + result["conditions"][name] = { + **meta, + "edited_text": edited_text, + **trajectory_deltas(baseline["trajectory"], traj), + **control_deltas(baseline["trajectory"], traj, dt), + } + + return result + + +# --------------------------------------------------------------------------- +# Entry point +# --------------------------------------------------------------------------- + +def main(): + args = parse_args() + logging.basicConfig(level=logging.INFO if args.verbose else logging.WARNING) + + config = load_config(args.config) + if not config["model"].get("use_cot", False): + logger.warning( + "config['model']['use_cot'] is False -- this checkpoint was not " + "trained to produce CoT reasoning, so the concept/prefix/suffix/" + "mistake conditions here are unlikely to be meaningful." + ) + + concepts = [c.strip() for c in args.concepts.split(",") if c.strip()] + threshold_words = [int(n) for n in args.threshold_words.split(",") if n.strip()] + + processor = AutoProcessor.from_pretrained(config["model"]["pretrained_model_path"], use_fast=True) + dataset = SFTDataset(config["data"]["val"], config["model"], processor) + + model = SFTAutoVLA(config) + model.autovla.vlm.resize_token_embeddings(len(processor.tokenizer)) + state_dict = torch.load(Path(args.checkpoint), map_location=args.device)["state_dict"] + model.autovla.load_state_dict(state_dict, strict=False) + model.to(args.device) + model.autovla.device = args.device + model.eval() + + sample_num = len(dataset.scenes) + if args.num_samples is not None: + sample_num = min(args.num_samples, sample_num) + logger.info("Running CoT ablation over %d scenes", sample_num) + + n_success, n_skipped = 0, 0 + outdir = Path(args.output).resolve().parent + outdir.mkdir(parents=True, exist_ok=True) + + with open(args.output, "a") as out_f: + for idx in tqdm(range(sample_num), desc="Scenes"): + scene_path, _ = dataset.scenes[idx] + with open(scene_path, "r") as f: + scene_data = json.load(f) + + input_features: Dict[str, Any] = {} + for builder in dataset._agent.get_feature_builders(): + input_features.update(builder.compute_features(scene_data)) + + result = run_scene( + model, input_features, concepts, threshold_words, + args.max_new_tokens, args.seed, args.device, + ) + if result is None: + n_skipped += 1 + continue + + result["scene"] = str(scene_path) + out_f.write(json.dumps(result) + "\n") + out_f.flush() + n_success += 1 + + if args.verbose: + logger.info( + "[%d/%d] %s no_cot_ade=%.4f concept_ade=%.4f", + idx + 1, sample_num, scene_path.name, + result["conditions"].get("no_cot", {}).get("ade_m") or -1.0, + result["conditions"].get("concept_mask", {}).get("ade_m") or -1.0, + ) + + logger.info("Done: %d succeeded, %d skipped (no action tokens found). Results: %s", + n_success, n_skipped, args.output) + + +if __name__ == "__main__": + main() diff --git a/tools/preprocessing/nocot_sample_generation.py b/tools/preprocessing/nocot_sample_generation.py index e60dfa2..7a67239 100644 --- a/tools/preprocessing/nocot_sample_generation.py +++ b/tools/preprocessing/nocot_sample_generation.py @@ -9,7 +9,6 @@ from torch.utils.data import DataLoader import shutil from dataset_utils.preprocessing.nuplan_dataset import NuplanCoTAnnotationDataset, DataCollator as NuplanDataCollator -from dataset_utils.preprocessing.waymo_e2e_dataset import WaymoE2ECoTAnnotationDataset, DataCollator as WaymoDataCollator CAM_LIST = ['front', 'front_left', 'front_right', @@ -111,6 +110,7 @@ def process_sample(sample, dataset_name): dataset = NuplanCoTAnnotationDataset(config, processor) collator = NuplanDataCollator(processor) elif dataset_name == "waymo": + from dataset_utils.preprocessing.waymo_e2e_dataset import WaymoE2ECoTAnnotationDataset, DataCollator as WaymoDataCollator dataset = WaymoE2ECoTAnnotationDataset(config, processor) collator = WaymoDataCollator(processor) else: