Self-Forcing / scripts /run_conditional_probe_offline.py
Cccccz's picture
Upload Python scripts
bc29ee3 verified
Raw History Blame Contribute Delete
18.7 kB
#!/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()