Download scripts/run_conditional_probe_offline.py from Cccccz/Self-Forcing: direct link, hf CLI and curl.
- Browser
- Download file 18.7 kB
-
https://huggingface.co/Cccccz/Self-Forcing/resolve/main/scripts/run_conditional_probe_offline.py
- Command line
-
hf download hf://Cccccz/Self-Forcing/scripts/run_conditional_probe_offline.py
-
curl -L -o run_conditional_probe_offline.py https://huggingface.co/Cccccz/Self-Forcing/resolve/main/scripts/run_conditional_probe_offline.py
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() | |