Download scripts/run_aligned_conditional_probe_3models.py from Cccccz/Self-Forcing: direct link, hf CLI and curl.
- Browser
- Download file 21.4 kB
-
https://huggingface.co/Cccccz/Self-Forcing/resolve/main/scripts/run_aligned_conditional_probe_3models.py
- Command line
-
hf download hf://Cccccz/Self-Forcing/scripts/run_aligned_conditional_probe_3models.py
-
curl -L -o run_aligned_conditional_probe_3models.py https://huggingface.co/Cccccz/Self-Forcing/resolve/main/scripts/run_aligned_conditional_probe_3models.py
21.4 kB
| #!/usr/bin/env python3 | |
| """Oracle flow-aligned conditional Ridge probes on three AR4 backbones. | |
| The target is the current chunk/current denoising-step full-grid feature. The | |
| first input is the current chunk/previous-step feature. The second input is | |
| the previous chunk/same-step boundary map, either raw or warped by global, | |
| correct, negated, or spatially shuffled target-to-source flow. | |
| This is an oracle diagnostic because the flow is computed from the generated | |
| current RGB frame. It tests whether alignment makes the previous-chunk route | |
| more predictive; it is not an inference-time implementation. | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import csv | |
| import json | |
| import os | |
| from collections import defaultdict | |
| 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) | |
| PHYSICAL_GPU = preparse_gpu() | |
| import numpy as np | |
| import torch | |
| import torch.nn.functional as F | |
| from analyze_fullgrid_bilinear_3models import ( | |
| GridRun, | |
| farneback, | |
| load_causal_runs, | |
| load_hy_runs, | |
| load_self_runs, | |
| resize_flow, | |
| shuffled_flow, | |
| warp, | |
| ) | |
| PROBES = ( | |
| "step_only", | |
| "both_raw", | |
| "both_global", | |
| "both_flow", | |
| "both_negated_flow", | |
| "both_shuffled_flow", | |
| ) | |
| LAYER_ROLES = {7: "early", 14: "middle", 22: "late", 29: "final"} | |
| def parse_args() -> argparse.Namespace: | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument("--gpu", default=PHYSICAL_GPU) | |
| parser.add_argument("--self_root", type=Path, required=True) | |
| parser.add_argument("--causal_root", type=Path, required=True) | |
| parser.add_argument("--hy_root", type=Path) | |
| parser.add_argument("--hy_cache_root", type=Path) | |
| parser.add_argument("--output_root", type=Path, required=True) | |
| parser.add_argument("--projection_dim", type=int, default=64) | |
| parser.add_argument("--ridge", type=float, default=1e-4) | |
| parser.add_argument("--seed", type=int, default=20260828) | |
| parser.add_argument( | |
| "--multilayer_self_causal", | |
| action="store_true", | |
| help="Use four-layer projected full grids for Self/Causal and skip HY.", | |
| ) | |
| return parser.parse_args() | |
| def write_csv(path: Path, rows: list[dict[str, Any]]) -> None: | |
| if not rows: | |
| return | |
| fields: list[str] = [] | |
| for row in rows: | |
| for key in row: | |
| if key not in fields: | |
| fields.append(key) | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| with path.open("w", newline="", encoding="utf-8") as handle: | |
| writer = csv.DictWriter(handle, fieldnames=fields) | |
| writer.writeheader() | |
| writer.writerows(rows) | |
| def load_multilayer_self_runs(root: Path) -> dict[str, list[GridRun]]: | |
| result = {role: [] for role in LAYER_ROLES.values()} | |
| for path in sorted((root / "runs").glob("prompt_*.pt")): | |
| state = torch.load(path, map_location="cpu", weights_only=False) | |
| projected = state.get("projected_by_layer", {}) | |
| if not projected: | |
| raise ValueError(f"No multilayer projected features in {path}") | |
| anchors = np.load(path.with_suffix(".anchors.npz"), allow_pickle=False)["frames"] | |
| by_layer: dict[int, dict[tuple[int, int], torch.Tensor]] = defaultdict(dict) | |
| for key, tensor in projected.items(): | |
| layer, chunk, step = (int(value) for value in key.split(":")) | |
| by_layer[layer][(chunk, step)] = tensor.float() | |
| for layer, role in LAYER_ROLES.items(): | |
| if layer not in by_layer: | |
| raise ValueError(f"Missing layer {layer} in {path}") | |
| result[role].append( | |
| GridRun( | |
| "self_forcing", | |
| "none", | |
| int(state.get("run_index", len(result[role]))), | |
| anchors, | |
| int(state["num_frame_per_block"]), | |
| by_layer[layer], | |
| path, | |
| ) | |
| ) | |
| return result | |
| def load_multilayer_causal_runs(root: Path) -> dict[str, list[GridRun]]: | |
| result = {role: [] for role in LAYER_ROLES.values()} | |
| for run_dir in sorted((root / "runs").glob("prompt_*")): | |
| path = run_dir / "feature_snapshots.pt" | |
| anchor_path = run_dir / "rgb_anchor_frames.npz" | |
| if not path.exists() or not anchor_path.exists(): | |
| continue | |
| state = torch.load(path, map_location="cpu", weights_only=False) | |
| projected = state.get("projected", {}) | |
| if not projected: | |
| raise ValueError(f"No projected features in {path}") | |
| anchors = np.load(anchor_path, allow_pickle=False)["frames"] | |
| by_layer: dict[int, dict[tuple[int, int], torch.Tensor]] = defaultdict(dict) | |
| for key, tensor in projected.items(): | |
| layer, chunk, step = (int(value) for value in key.split(":")) | |
| by_layer[layer][(chunk, step)] = tensor.float() | |
| for layer, role in LAYER_ROLES.items(): | |
| if layer not in by_layer: | |
| raise ValueError(f"Missing layer {layer} in {path}") | |
| result[role].append( | |
| GridRun( | |
| "causal_forcing", | |
| "none", | |
| int(state["prompt_id"]), | |
| anchors, | |
| 3, | |
| by_layer[layer], | |
| path, | |
| ) | |
| ) | |
| return result | |
| def columns(name: str, data: dict[str, torch.Tensor]) -> list[torch.Tensor]: | |
| ones = torch.ones_like(data["step"]) | |
| mapping = { | |
| "step_only": [data["step"], ones], | |
| "both_raw": [data["step"], data["raw"], ones], | |
| "both_global": [data["step"], data["global"], ones], | |
| "both_flow": [data["step"], data["flow"], ones], | |
| "both_negated_flow": [data["step"], data["negated"], ones], | |
| "both_shuffled_flow": [data["step"], data["shuffled"], ones], | |
| } | |
| return mapping[name] | |
| def collect_prompt_step(run: GridRun, target_step: int) -> dict[str, torch.Tensor]: | |
| collected: dict[str, list[torch.Tensor]] = defaultdict(list) | |
| for chunk in range(1, run.chunks): | |
| source_frame_index = chunk * run.chunk_size - 1 | |
| source_frame = run.anchors[source_frame_index] | |
| source_map = run.features[(chunk - 1, target_step)][-1].float() | |
| target_maps = run.features[(chunk, target_step)].float() | |
| step_maps = run.features[(chunk, target_step - 1)].float() | |
| for slot in range(run.chunk_size): | |
| target_frame = run.anchors[chunk * run.chunk_size + slot] | |
| flow = resize_flow(farneback(target_frame, source_frame)) | |
| global_flow = torch.zeros_like(flow) | |
| global_flow[0].fill_(float(torch.median(flow[0]))) | |
| global_flow[1].fill_(float(torch.median(flow[1]))) | |
| control_flows = { | |
| "global": global_flow, | |
| "flow": flow, | |
| "negated": -flow, | |
| "shuffled": shuffled_flow( | |
| flow, | |
| seed=(run.prompt_id + 1) * 100000 | |
| + chunk * 1000 | |
| + slot * 10 | |
| + target_step, | |
| ), | |
| } | |
| aligned: dict[str, torch.Tensor] = {} | |
| masks: list[torch.Tensor] = [] | |
| for name, control_flow in control_flows.items(): | |
| aligned[name], mask = warp(source_map, control_flow) | |
| masks.append(mask) | |
| common_mask = torch.stack(masks).all(dim=0) | |
| if not bool(common_mask.any()): | |
| continue | |
| collected["target"].append(target_maps[slot][common_mask]) | |
| collected["step"].append(step_maps[slot][common_mask]) | |
| collected["raw"].append(source_map[common_mask]) | |
| for name in control_flows: | |
| collected[name].append(aligned[name][common_mask]) | |
| result = {key: torch.cat(values, dim=0).contiguous() for key, values in collected.items()} | |
| expected = {"target", "step", "raw", "global", "flow", "negated", "shuffled"} | |
| if set(result) != expected: | |
| raise ValueError(f"Incomplete aligned sample for {run.source}: {set(result)}") | |
| return result | |
| def ridge_sufficient_statistics( | |
| data: dict[str, torch.Tensor], | |
| probe: str, | |
| device: torch.device, | |
| ) -> tuple[torch.Tensor, torch.Tensor]: | |
| design = torch.stack(columns(probe, data), dim=-1).to(device=device, dtype=torch.float64) | |
| target = data["target"].to(device=device, dtype=torch.float64) | |
| gram = torch.einsum("ndp,ndq->dpq", design, design) | |
| rhs = torch.einsum("ndp,nd->dp", design, target) | |
| return gram, rhs | |
| def solve_ridge( | |
| gram: torch.Tensor, | |
| rhs: torch.Tensor, | |
| ridge: float, | |
| ) -> torch.Tensor: | |
| parameter_count = gram.shape[-1] | |
| device = gram.device | |
| scale = gram.diagonal(dim1=-2, dim2=-1).mean(dim=-1).clamp_min(1e-8) | |
| regularizer = torch.eye(parameter_count, dtype=torch.float64, device=device)[None] | |
| regularizer = regularizer * (float(ridge) * scale[:, None, None]) | |
| regularizer[:, -1, -1] = 0.0 | |
| try: | |
| weights = torch.linalg.solve(gram + regularizer, rhs.unsqueeze(-1)).squeeze(-1) | |
| except torch.linalg.LinAlgError: | |
| weights = (torch.linalg.pinv(gram + regularizer) @ rhs.unsqueeze(-1)).squeeze(-1) | |
| return weights.float() | |
| def evaluate( | |
| data: dict[str, torch.Tensor], | |
| probe: str, | |
| weights: torch.Tensor, | |
| device: torch.device, | |
| ) -> dict[str, float]: | |
| design = torch.stack(columns(probe, data), dim=-1).to(device=device, dtype=torch.float32) | |
| target = data["target"].to(device=device, dtype=torch.float32) | |
| prediction = torch.einsum("ndp,dp->nd", design, weights) | |
| error = prediction - target | |
| mse = error.square().mean() | |
| variance = (target - target.mean()).square().mean().clamp_min(1e-12) | |
| nmse = mse / variance | |
| cosine = F.cosine_similarity(prediction, target, dim=-1, eps=1e-8).mean() | |
| 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 = 10000) -> tuple[float, float, float]: | |
| array = np.asarray(values, dtype=np.float64) | |
| generator = np.random.default_rng(seed) | |
| indices = generator.integers(0, len(array), size=(rounds, len(array))) | |
| means = array[indices].mean(axis=1) | |
| return float(array.mean()), float(np.quantile(means, 0.025)), float(np.quantile(means, 0.975)) | |
| def summarize(rows: list[dict[str, Any]], seed: int) -> list[dict[str, Any]]: | |
| groups: dict[tuple[str, str, str], list[dict[str, Any]]] = defaultdict(list) | |
| for row in rows: | |
| groups[(row["model"], row["layer_role"], row["probe"])].append(row) | |
| output = [] | |
| for (model, role, probe), selected in sorted(groups.items()): | |
| by_prompt: dict[int, list[dict[str, Any]]] = defaultdict(list) | |
| for row in selected: | |
| by_prompt[int(row["held_out_prompt"])].append(row) | |
| prompt_rows = [] | |
| for prompt_id, values in sorted(by_prompt.items()): | |
| item = {"prompt_id": prompt_id} | |
| for metric in ("mse", "nMSE", "nRMSE", "r2", "cosine"): | |
| item[metric] = float(np.mean([float(row[metric]) for row in values])) | |
| prompt_rows.append(item) | |
| item: dict[str, Any] = { | |
| "model": model, | |
| "layer_role": role, | |
| "probe": probe, | |
| "prompt_count": len(prompt_rows), | |
| "fold_count": len(selected), | |
| } | |
| for metric in ("mse", "nMSE", "nRMSE", "r2", "cosine"): | |
| stable = seed + sum(map(ord, model + role + probe + metric)) | |
| avg, low, high = bootstrap([row[metric] for row in prompt_rows], stable) | |
| item[f"{metric}_mean"] = avg | |
| item[f"{metric}_ci95_low"] = low | |
| item[f"{metric}_ci95_high"] = high | |
| output.append(item) | |
| prompt_metric: dict[tuple[str, str, str, int], dict[str, float]] = {} | |
| grouped: dict[tuple[str, str, str, int], list[dict[str, Any]]] = defaultdict(list) | |
| for row in rows: | |
| grouped[ | |
| (row["model"], row["layer_role"], row["probe"], int(row["held_out_prompt"])) | |
| ].append(row) | |
| for key, values in grouped.items(): | |
| prompt_metric[key] = { | |
| metric: float(np.mean([float(row[metric]) for row in values])) | |
| for metric in ("mse", "nMSE", "nRMSE", "r2", "cosine") | |
| } | |
| for item in output: | |
| model, role, probe = item["model"], item["layer_role"], item["probe"] | |
| if probe == "step_only": | |
| continue | |
| gain_step, gain_raw = [], [] | |
| prompt_ids = sorted( | |
| prompt_id | |
| for candidate_model, candidate_role, candidate_probe, prompt_id in prompt_metric | |
| if candidate_model == model and candidate_role == role and candidate_probe == probe | |
| ) | |
| for prompt_id in prompt_ids: | |
| current = prompt_metric[(model, role, probe, prompt_id)]["mse"] | |
| step = prompt_metric[(model, role, "step_only", prompt_id)]["mse"] | |
| raw = prompt_metric[(model, role, "both_raw", prompt_id)]["mse"] | |
| gain_step.append((step - current) / max(step, 1e-12)) | |
| gain_raw.append((raw - current) / max(raw, 1e-12)) | |
| avg, low, high = bootstrap( | |
| gain_step, seed + 300000 + sum(map(ord, model + role + probe)) | |
| ) | |
| item.update({ | |
| "mse_gain_vs_step_mean": avg, | |
| "mse_gain_vs_step_ci95_low": low, | |
| "mse_gain_vs_step_ci95_high": high, | |
| "mse_gain_vs_step_wins": int(sum(value > 0 for value in gain_step)), | |
| }) | |
| avg, low, high = bootstrap( | |
| gain_raw, seed + 600000 + sum(map(ord, model + role + probe)) | |
| ) | |
| item.update({ | |
| "mse_gain_vs_raw_mean": avg, | |
| "mse_gain_vs_raw_ci95_low": low, | |
| "mse_gain_vs_raw_ci95_high": high, | |
| "mse_gain_vs_raw_wins": int(sum(value > 0 for value in gain_raw)), | |
| }) | |
| if probe == "both_flow": | |
| for baseline_probe, label in ( | |
| ("both_global", "global"), | |
| ("both_negated_flow", "negated_flow"), | |
| ("both_shuffled_flow", "shuffled_flow"), | |
| ): | |
| gains = [] | |
| for prompt_id in prompt_ids: | |
| current = prompt_metric[(model, role, probe, prompt_id)]["mse"] | |
| baseline = prompt_metric[(model, role, baseline_probe, prompt_id)]["mse"] | |
| gains.append((baseline - current) / max(baseline, 1e-12)) | |
| avg, low, high = bootstrap( | |
| gains, | |
| seed + 900000 + sum(map(ord, model + role + baseline_probe)), | |
| ) | |
| item.update({ | |
| f"mse_gain_vs_{label}_mean": avg, | |
| f"mse_gain_vs_{label}_ci95_low": low, | |
| f"mse_gain_vs_{label}_ci95_high": high, | |
| f"mse_gain_vs_{label}_wins": int(sum(value > 0 for value in gains)), | |
| }) | |
| return output | |
| def main() -> None: | |
| args = parse_args() | |
| output = args.output_root.resolve() | |
| output.mkdir(parents=True, exist_ok=True) | |
| device = torch.device("cuda:0") | |
| if not torch.cuda.is_available(): | |
| raise RuntimeError("CUDA is required for this experiment") | |
| if args.multilayer_self_causal: | |
| model_runs = { | |
| "self_forcing": load_multilayer_self_runs(args.self_root.resolve()), | |
| "causal_forcing": load_multilayer_causal_runs(args.causal_root.resolve()), | |
| } | |
| else: | |
| if args.hy_root is None or args.hy_cache_root is None: | |
| raise ValueError("--hy_root and --hy_cache_root are required without multilayer mode") | |
| model_runs = { | |
| "self_forcing": {"final": load_self_runs(args.self_root.resolve())}, | |
| "causal_forcing": {"final": load_causal_runs(args.causal_root.resolve())}, | |
| "hy_static": { | |
| "final": load_hy_runs( | |
| args.hy_root.resolve(), | |
| "static", | |
| args.hy_cache_root.resolve(), | |
| args.projection_dim, | |
| device, | |
| False, | |
| ) | |
| }, | |
| } | |
| fold_rows: list[dict[str, Any]] = [] | |
| for model, role_runs in model_runs.items(): | |
| for role, runs in role_runs.items(): | |
| if len(runs) != 10: | |
| raise ValueError(f"Expected 10 runs for {model}/{role}, found {len(runs)}") | |
| for target_step in range(1, 4): | |
| prepared = [collect_prompt_step(run, target_step) for run in runs] | |
| token_counts = [int(data["target"].shape[0]) for data in prepared] | |
| print( | |
| f"[prepare] {model}/{role} step={target_step} tokens={token_counts}", | |
| flush=True, | |
| ) | |
| statistics = { | |
| probe: [ridge_sufficient_statistics(data, probe, device) for data in prepared] | |
| for probe in PROBES | |
| } | |
| for held_out in range(10): | |
| test = prepared[held_out] | |
| for probe in PROBES: | |
| grams, right_sides = zip(*statistics[probe]) | |
| train_gram = torch.stack(grams).sum(dim=0) - grams[held_out] | |
| train_rhs = torch.stack(right_sides).sum(dim=0) - right_sides[held_out] | |
| weights = solve_ridge(train_gram, train_rhs, args.ridge) | |
| values = evaluate(test, probe, weights, device) | |
| fold_rows.append({ | |
| "model": model, | |
| "layer_role": role, | |
| "target_step": target_step, | |
| "held_out_prompt": held_out, | |
| "train_prompts": 9, | |
| "test_tokens": int(test["target"].shape[0]), | |
| "probe": probe, | |
| **values, | |
| }) | |
| print( | |
| f"[fold] {model}/{role} step={target_step} heldout={held_out}", | |
| flush=True, | |
| ) | |
| del prepared | |
| torch.cuda.empty_cache() | |
| summary = summarize(fold_rows, args.seed) | |
| write_csv(output / "aligned_probe_folds.csv", fold_rows) | |
| write_csv(output / "aligned_probe_summary.csv", summary) | |
| summary_lookup = { | |
| (row["model"], row["layer_role"], row["probe"]): row for row in summary | |
| } | |
| report = [ | |
| "# Oracle flow-aligned conditional Ridge probe", | |
| "", | |
| "All gains are prompt-wise relative MSE reductions averaged over 10 held-out prompts and three target denoising steps. Flow is computed from the generated current RGB frame and is therefore an oracle diagnostic.", | |
| "", | |
| "| model | layer | raw chunk vs step-only | flow-aligned vs step-only | flow-aligned vs raw chunk | flow-aligned vs shuffled flow | flow-vs-raw wins |", | |
| "|---|---|---:|---:|---:|---:|---:|", | |
| ] | |
| for model, role_runs in model_runs.items(): | |
| for role in role_runs: | |
| raw = summary_lookup[(model, role, "both_raw")] | |
| flow = summary_lookup[(model, role, "both_flow")] | |
| report.append( | |
| f"| {model} | {role} | {100 * raw['mse_gain_vs_step_mean']:.2f}% | " | |
| f"{100 * flow['mse_gain_vs_step_mean']:.2f}% | " | |
| f"{100 * flow['mse_gain_vs_raw_mean']:.2f}% | " | |
| f"{100 * flow['mse_gain_vs_shuffled_flow_mean']:.2f}% | " | |
| f"{flow['mse_gain_vs_raw_wins']}/10 |" | |
| ) | |
| report.extend([ | |
| "", | |
| "Self-Forcing and Causal-Forcing are evaluated at early, middle, late, and final layers under one identical feature space, mask, split, and Ridge capacity.", | |
| "", | |
| "The experiment uses the common intersection of in-bounds masks for every warp, so all predictor variants see identical target tokens. Full-grid features are fixed 64-D signed random projections; conclusions concern within-model paired gains rather than native-space or cross-model absolute errors.", | |
| ]) | |
| (output / "REPORT.md").write_text("\n".join(report) + "\n", encoding="utf-8") | |
| config = { | |
| "gpu": str(args.gpu), | |
| "models": list(model_runs), | |
| "layer_roles": {model: list(role_runs) for model, role_runs in model_runs.items()}, | |
| "prompt_count": 10, | |
| "target_steps": [1, 2, 3], | |
| "probes": list(PROBES), | |
| "ridge": args.ridge, | |
| "projection_dim": args.projection_dim, | |
| "grid": [30, 52], | |
| "split": "leave-one-prompt-out (9 train, 1 test)", | |
| "support": "intersection of in-bounds masks for global/correct/negated/shuffled warps", | |
| "flow": "Farneback target RGB to previous-chunk boundary RGB; oracle diagnostic", | |
| "row_count": len(fold_rows), | |
| } | |
| (output / "config.json").write_text(json.dumps(config, indent=2) + "\n", encoding="utf-8") | |
| print(f"[complete] {output} rows={len(fold_rows)}", flush=True) | |
| if __name__ == "__main__": | |
| main() | |