#!/usr/bin/env python3 """Run grouped, channel-wise Linear/Ridge conditional probes offline. The input is the normalized per-prompt dataset produced by ``build_conditional_probe_dataset.py``. All controls are derived in memory from the same canonical features, so no model forward is repeated and every probe uses identical target tokens. """ from __future__ import annotations import argparse import csv import json import math import os import sys from pathlib import Path from typing import Any def _preparse_gpu() -> str: parser = argparse.ArgumentParser(add_help=False) parser.add_argument("--gpu", default="0") args, _ = parser.parse_known_args() os.environ["CUDA_VISIBLE_DEVICES"] = str(args.gpu) return str(args.gpu) _preparse_gpu() import numpy as np import torch import torch.nn.functional as F FAMILIES = ("self_forcing", "causal_forcing", "hy_worldplay") ROLES = ("early", "middle", "late", "final") LAYER_INDICES = { "self_forcing": {"early": 7, "middle": 14, "late": 22, "final": 29}, "causal_forcing": {"early": 7, "middle": 14, "late": 22, "final": 29}, "hy_worldplay": {"early": 13, "middle": 26, "late": 40, "final": 53}, } PROBES = ( "within_affine", "within_quadratic", "cross_affine", "fusion_same", "fusion_step_duplicate", "fusion_distant", "fusion_wrong_step", "fusion_token_shuffle", "fusion_batch_shuffle", "fusion_zero", "fusion_noise", ) def write_csv(path: Path, rows: list[dict[str, Any]]) -> None: if not rows: return path.parent.mkdir(parents=True, exist_ok=True) fields: list[str] = [] for row in rows: for key in row: if key not in fields: fields.append(key) with path.open("w", newline="", encoding="utf-8") as handle: writer = csv.DictWriter(handle, fieldnames=fields, extrasaction="ignore") writer.writeheader() writer.writerows(rows) def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser() parser.add_argument("--dataset_root", type=Path, required=True) parser.add_argument("--output_dir", type=Path, required=True) parser.add_argument("--num_prompts", type=int, default=10) parser.add_argument("--chunks", type=int, default=4) parser.add_argument("--steps", type=int, default=4) parser.add_argument("--ridge", type=float, default=1e-4) parser.add_argument("--seed", type=int, default=0) parser.add_argument( "--chunk_pairing", choices=("matched_slot", "boundary_to_all"), default="matched_slot", help=( "How the previous-chunk auxiliary feature is paired with the current chunk. " "boundary_to_all broadcasts the previous chunk's final temporal slot to " "all current temporal slots at matched spatial coordinates." ), ) return parser.parse_args() def load_family(root: Path, family: str, count: int) -> list[dict[str, Any]]: runs = [] for prompt_id in range(count): pt_path = root / family / f"prompt_{prompt_id:04d}.pt" npz_path = root / family / f"prompt_{prompt_id:04d}.npz" path = npz_path if family == "hy_worldplay" and not pt_path.exists() else pt_path if not path.exists(): raise FileNotFoundError(path) if path.suffix == ".npz": data = np.load(path, allow_pickle=False) raw: dict[int, dict[tuple[int, int], torch.Tensor]] = {} for index, stage_value in enumerate(data["stages"]): stage = str(stage_value) if not stage.startswith("block_"): continue layer = int(stage.split("_")[-1]) chunk = int(data["chunks"][index]) step = int(data["steps"][index]) raw.setdefault(layer, {})[(chunk, step)] = torch.from_numpy(data["features"][index]) layer_roles = {13: "early", 26: "middle", 40: "late", 53: "final"} features = {} for layer, role in layer_roles.items(): rows = [] for chunk in range(4): rows.append(torch.stack([raw[layer][(chunk, step)] for step in range(4)], dim=0)) features[role] = torch.stack(rows, dim=0).contiguous() run = { "prompt_id": prompt_id, "features": features, "timesteps": np.asarray(data["timesteps"], dtype=np.float32), "coords": np.asarray(data["coords"], dtype=np.int64), "grid_shape": np.asarray(data["grid_shape"], dtype=np.int64), } else: run = torch.load(path, map_location="cpu", weights_only=False) if int(run.get("prompt_id", prompt_id)) != prompt_id: raise ValueError(f"Prompt id mismatch in {path}") runs.append(run) return runs def boundary_to_all(reference: torch.Tensor, run: dict[str, Any]) -> torch.Tensor: """Broadcast the last temporal slot while preserving target token order.""" coords = np.asarray(run.get("coords")) if coords.ndim != 2 or coords.shape[1] != 3 or len(coords) != reference.shape[0]: raise ValueError( "boundary_to_all requires one (temporal,y,x) coordinate per feature token; " f"coords={coords.shape}, features={tuple(reference.shape)}" ) slots = sorted(int(value) for value in np.unique(coords[:, 0])) if not slots: raise ValueError("No temporal slots in coordinates") last_mask = coords[:, 0] == slots[-1] source_coords = coords[last_mask, 1:] source = reference[torch.from_numpy(last_mask)] result = torch.empty_like(reference) for slot in slots: target_mask = coords[:, 0] == slot target_coords = coords[target_mask, 1:] if not np.array_equal(target_coords, source_coords): raise ValueError( f"Temporal slot {slot} does not share the boundary slot's spatial grid" ) result[torch.from_numpy(target_mask)] = source return result def samples( run: dict[str, Any], role: str, chunks: int, step: int, chunk_pairing: str = "matched_slot", ) -> dict[str, torch.Tensor]: values = run["features"][role].float() if values.ndim != 4: raise ValueError(f"Expected [chunk,step,token,channel], got {values.shape}") if values.shape[0] < chunks or values.shape[1] < 4: raise ValueError(f"Insufficient feature grid for {role}: {values.shape}") targets, within, cross, distant, wrong = [], [], [], [], [] # c>=2 is required for the distant-chunk control. Pool c=2 and c=3 for # the common four-chunk protocol, while retaining chunk_id downstream. for chunk in range(2, chunks): targets.append(values[chunk, step]) within.append(values[chunk, step - 1]) if chunk_pairing == "boundary_to_all": cross.append(boundary_to_all(values[chunk - 1, step], run)) distant.append(boundary_to_all(values[chunk - 2, step], run)) wrong.append(boundary_to_all(values[chunk - 1, step - 1], run)) else: cross.append(values[chunk - 1, step]) distant.append(values[chunk - 2, step]) wrong.append(values[chunk - 1, step - 1]) return { "target": torch.cat(targets, dim=0), "within": torch.cat(within, dim=0), "cross": torch.cat(cross, dim=0), "distant": torch.cat(distant, dim=0), "wrong": torch.cat(wrong, dim=0), } def token_shuffle(value: torch.Tensor, spatial_count: int) -> torch.Tensor: # Roll spatial tokens independently inside every chunk/temporal slice, # preserving the marginal feature distribution while breaking coordinate # correspondence. This supports both 3-slot Self/Causal and 4-slot HY. if spatial_count <= 0 or value.shape[0] % spatial_count: return value.roll(shifts=max(1, value.shape[0] // 2), dims=0) frames = value.reshape(-1, spatial_count, value.shape[1]) return frames.roll(shifts=1, dims=1).reshape_as(value) def noise_like(value: torch.Tensor, seed: int) -> torch.Tensor: generator = torch.Generator(device="cpu").manual_seed(int(seed)) noise = torch.randn(value.shape, generator=generator, dtype=value.dtype) return noise * value.std(dim=0, keepdim=True).clamp_min(1e-6) + value.mean(dim=0, keepdim=True) def columns( name: str, data: dict[str, torch.Tensor], batch: torch.Tensor, seed: int, spatial_count: int, ): within, cross = data["within"], data["cross"] ones = torch.ones_like(within) mapping = { "within_affine": [within, ones], "within_quadratic": [within, within.square(), ones], "cross_affine": [cross, ones], "fusion_same": [within, cross, ones], "fusion_step_duplicate": [within, within, ones], "fusion_distant": [within, data["distant"], ones], "fusion_wrong_step": [within, data["wrong"], ones], "fusion_token_shuffle": [within, token_shuffle(cross, spatial_count), ones], "fusion_batch_shuffle": [within, batch, ones], "fusion_zero": [within, torch.zeros_like(cross), ones], "fusion_noise": [within, noise_like(cross, seed), ones], } return mapping[name] def fit_ridge(features: list[torch.Tensor], target: torch.Tensor, ridge: float) -> torch.Tensor: design = torch.stack(features, dim=-1).double() # [N,D,P] y = target.double() gram = torch.einsum("ndp,ndq->dpq", design, design) rhs = torch.einsum("ndp,nd->dp", design, y) p = gram.shape[-1] scale = gram.diagonal(dim1=-2, dim2=-1).mean(dim=-1).clamp_min(1e-8) reg = torch.eye(p, dtype=gram.dtype).unsqueeze(0) * (float(ridge) * scale[:, None, None]) # The last column is the explicit bias and is not regularized. reg[:, -1, -1] = 0.0 try: return torch.linalg.solve(gram + reg, rhs.unsqueeze(-1)).squeeze(-1).float() except torch.linalg.LinAlgError: return (torch.linalg.pinv(gram + reg) @ rhs.unsqueeze(-1)).squeeze(-1).float() def predict(features: list[torch.Tensor], weights: torch.Tensor) -> torch.Tensor: return torch.einsum("ndp,dp->nd", torch.stack(features, dim=-1).float(), weights) def metrics(pred: torch.Tensor, target: torch.Tensor) -> dict[str, float]: pred, target = pred.float(), target.float() error = pred - target mse = error.square().mean() centered = target - target.mean() variance = centered.square().mean().clamp_min(1e-12) nmse = mse / variance cosine = F.cosine_similarity(pred.reshape(1, -1), target.reshape(1, -1), dim=1, eps=1e-8)[0] return { "mse": float(mse), "nMSE": float(nmse), "nRMSE": float(torch.sqrt(nmse)), "r2": float(1.0 - nmse), "cosine": float(cosine), } def bootstrap(values: list[float], seed: int, rounds: int = 4000): values = np.asarray(values, dtype=np.float64) rng = np.random.default_rng(seed) if values.size == 0: return float("nan"), float("nan"), float("nan") draws = rng.integers(0, values.size, size=(rounds, values.size)) means = values[draws].mean(axis=1) return float(values.mean()), float(np.quantile(means, 0.025)), float(np.quantile(means, 0.975)) def main() -> None: args = parse_args() args.dataset_root = args.dataset_root.resolve() args.output_dir = args.output_dir.resolve() args.output_dir.mkdir(parents=True, exist_ok=True) all_rows: list[dict[str, Any]] = [] config = { "dataset_root": str(args.dataset_root), "num_prompts": args.num_prompts, "chunks": args.chunks, "steps": args.steps, "target_chunks": list(range(2, args.chunks)), "ridge": args.ridge, "chunk_pairing": args.chunk_pairing, "outer_split": "test prompt p; donor (p+1)%N also excluded from training", "probes": list(PROBES), } for family in FAMILIES: print(f"[load] {family}", flush=True) runs = load_family(args.dataset_root, family, args.num_prompts) for layer_index, role in enumerate(ROLES): coords = np.asarray(runs[0].get("coords")) slots = sorted(int(value) for value in np.unique(coords[:, 0])) if not slots or len(coords) != runs[0]["features"][role].shape[2]: raise ValueError( f"Invalid temporal coordinates for {family}/{role}: " f"coords={coords.shape}, tokens={runs[0]['features'][role].shape[2]}" ) spatial_count = int((coords[:, 0] == slots[0]).sum()) for step in range(1, args.steps): prepared = [ samples(run, role, args.chunks, step, args.chunk_pairing) for run in runs ] for held_out in range(args.num_prompts): donor = (held_out + 1) % args.num_prompts train_ids = [i for i in range(args.num_prompts) if i not in {held_out, donor}] train_data = {key: torch.cat([prepared[i][key] for i in train_ids], dim=0) for key in prepared[0]} test_data = prepared[held_out] donor_data = prepared[donor] train_batch = torch.cat( [prepared[train_ids[(position + 1) % len(train_ids)]]["cross"] for position in range(len(train_ids))], dim=0, ) for probe in PROBES: train_cols = columns( probe, train_data, train_batch, seed=args.seed + held_out * 100 + step, spatial_count=spatial_count, ) test_cols = columns( probe, test_data, donor_data["cross"], seed=args.seed + 10000 + held_out * 100 + step, spatial_count=spatial_count, ) weights = fit_ridge(train_cols, train_data["target"], args.ridge) pred = predict(test_cols, weights) row = { "model_family": family, "layer_role": role, "layer_index": LAYER_INDICES[family][role], "target_step": step, "held_out_prompt": held_out, "other_video_prompt": donor, "train_prompts": len(train_ids), "test_tokens": int(test_data["target"].shape[0]), "probe": probe, **metrics(pred, test_data["target"]), } all_rows.append(row) if held_out % 2 == 0: print( f"[progress] {family} {role} step={step} heldout={held_out}", flush=True, ) write_csv(args.output_dir / "linear_probe_folds.csv", all_rows) summary_rows = [] for family in FAMILIES: for role in ROLES: for step in range(1, args.steps): for probe in PROBES: selected = [ row for row in all_rows if row["model_family"] == family and row["layer_role"] == role and row["target_step"] == step and row["probe"] == probe ] if not selected: continue item = { "model_family": family, "layer_role": role, "target_step": step, "probe": probe, "prompt_count": len(selected), } for metric in ("mse", "nMSE", "nRMSE", "r2", "cosine"): stable_seed = ( args.seed + 100000 * FAMILIES.index(family) + 10000 * ROLES.index(role) + 100 * int(step) + sum(ord(ch) for ch in probe) + sum(ord(ch) for ch in metric) ) mean, low, high = bootstrap( [float(row[metric]) for row in selected], stable_seed, ) item[f"{metric}_mean"] = mean item[f"{metric}_ci95_low"] = low item[f"{metric}_ci95_high"] = high baseline = [ row for row in all_rows if row["model_family"] == family and row["layer_role"] == role and row["target_step"] == step and row["probe"] == "within_affine" ] if baseline: gains = [ (float(base["mse"]) - float(cur["mse"])) / max(float(base["mse"]), 1e-12) for base, cur in zip( sorted(baseline, key=lambda row: row["held_out_prompt"]), sorted(selected, key=lambda row: row["held_out_prompt"]), ) ] mean, low, high = bootstrap(gains, args.seed + 700000 + step) item.update({ "gain_vs_within_affine_mean": mean, "gain_vs_within_affine_ci95_low": low, "gain_vs_within_affine_ci95_high": high, "gain_vs_within_affine_wins": sum(value > 0 for value in gains), }) summary_rows.append(item) write_csv(args.output_dir / "linear_probe_summary.csv", summary_rows) config["fold_rows"] = len(all_rows) config["summary_rows"] = len(summary_rows) (args.output_dir / "config.json").write_text(json.dumps(config, indent=2) + "\n", encoding="utf-8") print(f"[complete] {args.output_dir} rows={len(all_rows)}", flush=True) if __name__ == "__main__": main()