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