Download scripts/analyze_fullgrid_bilinear_3models.py from Cccccz/Self-Forcing: direct link, hf CLI and curl.
- Browser
- Download file 22 kB
-
https://huggingface.co/Cccccz/Self-Forcing/resolve/main/scripts/analyze_fullgrid_bilinear_3models.py
- Command line
-
hf download hf://Cccccz/Self-Forcing/scripts/analyze_fullgrid_bilinear_3models.py
-
curl -L -o analyze_fullgrid_bilinear_3models.py https://huggingface.co/Cccccz/Self-Forcing/resolve/main/scripts/analyze_fullgrid_bilinear_3models.py
22 kB
| #!/usr/bin/env python3 | |
| """Full-grid bilinear motion alignment for the three four-step backbones. | |
| This intentionally follows the original Self-Forcing pilot: compare every | |
| current-chunk temporal slot with the previous chunk's boundary feature map and | |
| warp the source map with continuous target-to-source flow using bilinear | |
| ``grid_sample``. Correct, global, negated, and spatially shuffled flow fields | |
| are evaluated with identical interpolation and in-bounds masking. | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import csv | |
| import json | |
| import math | |
| from collections import defaultdict | |
| from pathlib import Path | |
| from typing import Any, Iterable | |
| import cv2 | |
| import matplotlib | |
| matplotlib.use("Agg") | |
| import matplotlib.pyplot as plt | |
| import numpy as np | |
| import torch | |
| import torch.nn.functional as F | |
| GRID_H, GRID_W = 30, 52 | |
| PROJECTION_SEED_BASE = 20260728 | |
| def parse_args() -> argparse.Namespace: | |
| parser = argparse.ArgumentParser() | |
| 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, required=True) | |
| parser.add_argument("--hy_right_root", type=Path, required=True) | |
| parser.add_argument("--output_root", type=Path, required=True) | |
| parser.add_argument("--projection_dim", type=int, default=64) | |
| parser.add_argument("--projection_device", default="cpu") | |
| parser.add_argument("--overwrite_projection_cache", action="store_true") | |
| return parser.parse_args() | |
| def mean(values: Iterable[float]) -> float: | |
| values = [float(value) for value in values if np.isfinite(value)] | |
| return float(np.mean(values)) if values else float("nan") | |
| 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 projection_matrix(dim: int, output_dim: int, device: torch.device) -> torch.Tensor: | |
| generator = torch.Generator(device="cpu").manual_seed(PROJECTION_SEED_BASE + dim) | |
| signs = torch.randint(0, 2, (dim, output_dim), generator=generator, dtype=torch.int8) | |
| return signs.float().mul_(2).sub_(1).div_(math.sqrt(output_dim)).to(device) | |
| def farneback(source: np.ndarray, target: np.ndarray) -> np.ndarray: | |
| def gray(frame: np.ndarray) -> np.ndarray: | |
| if frame.dtype != np.uint8: | |
| frame = np.uint8(np.clip(np.round(frame * 255.0), 0, 255)) | |
| return cv2.cvtColor(frame, cv2.COLOR_RGB2GRAY) | |
| return cv2.calcOpticalFlowFarneback( | |
| gray(source), | |
| gray(target), | |
| None, | |
| pyr_scale=0.5, | |
| levels=4, | |
| winsize=21, | |
| iterations=5, | |
| poly_n=7, | |
| poly_sigma=1.5, | |
| flags=0, | |
| ) | |
| def resize_flow(flow: np.ndarray) -> torch.Tensor: | |
| source_h, source_w = flow.shape[:2] | |
| resized = cv2.resize(flow, (GRID_W, GRID_H), interpolation=cv2.INTER_AREA) | |
| resized[..., 0] *= GRID_W / source_w | |
| resized[..., 1] *= GRID_H / source_h | |
| return torch.from_numpy(resized).permute(2, 0, 1).float() | |
| def warp(source: torch.Tensor, flow: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: | |
| yy, xx = torch.meshgrid( | |
| torch.arange(GRID_H, dtype=torch.float32), | |
| torch.arange(GRID_W, dtype=torch.float32), | |
| indexing="ij", | |
| ) | |
| sample_x = xx + flow[0] | |
| sample_y = yy + flow[1] | |
| grid = torch.stack( | |
| [ | |
| 2.0 * sample_x / (GRID_W - 1) - 1.0, | |
| 2.0 * sample_y / (GRID_H - 1) - 1.0, | |
| ], | |
| dim=-1, | |
| )[None] | |
| value = source.permute(2, 0, 1)[None].float() | |
| warped = F.grid_sample( | |
| value, | |
| grid, | |
| mode="bilinear", | |
| padding_mode="zeros", | |
| align_corners=True, | |
| )[0].permute(1, 2, 0) | |
| mask = ( | |
| (sample_x >= 0) | |
| & (sample_x <= GRID_W - 1) | |
| & (sample_y >= 0) | |
| & (sample_y <= GRID_H - 1) | |
| ) | |
| return warped, mask | |
| def cosine(target: torch.Tensor, source: torch.Tensor, mask: torch.Tensor | None = None) -> float: | |
| values = F.cosine_similarity(target.float(), source.float(), dim=-1, eps=1e-8) | |
| if mask is not None: | |
| values = values[mask] | |
| return float(values.mean()) if values.numel() else float("nan") | |
| class GridRun: | |
| def __init__( | |
| self, | |
| model: str, | |
| action: str, | |
| prompt_id: int, | |
| anchors: np.ndarray, | |
| chunk_size: int, | |
| features: dict[tuple[int, int], torch.Tensor], | |
| source: Path, | |
| ): | |
| self.model = model | |
| self.action = action | |
| self.prompt_id = int(prompt_id) | |
| self.anchors = anchors.astype(np.uint8) | |
| self.chunk_size = int(chunk_size) | |
| self.features = features | |
| self.source = source | |
| self.chunks = max(chunk for chunk, _ in features) + 1 | |
| self.steps = sorted({step for _, step in features}) | |
| def load_self_runs(root: Path) -> list[GridRun]: | |
| runs: list[GridRun] = [] | |
| for path in sorted((root / "runs").glob("prompt_*.pt")): | |
| state = torch.load(path, map_location="cpu", weights_only=False) | |
| features = { | |
| tuple(int(value) for value in key.split(":")): tensor.float() | |
| for key, tensor in state["projected"].items() | |
| } | |
| anchors = np.load(path.with_suffix(".anchors.npz"), allow_pickle=False)["frames"] | |
| runs.append( | |
| GridRun( | |
| "self_forcing", | |
| "none", | |
| int(state.get("run_index", len(runs))), | |
| anchors, | |
| int(state["num_frame_per_block"]), | |
| features, | |
| path, | |
| ) | |
| ) | |
| del state | |
| return runs | |
| def load_causal_runs(root: Path) -> list[GridRun]: | |
| runs: list[GridRun] = [] | |
| 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 full-grid projected features in {path}") | |
| features: dict[tuple[int, int], torch.Tensor] = {} | |
| for key, tensor in projected.items(): | |
| layer, chunk, step = (int(value) for value in key.split(":")) | |
| if layer == max(state["layers"]): | |
| features[(chunk, step)] = tensor.float() | |
| anchors = np.load(anchor_path, allow_pickle=False)["frames"] | |
| runs.append( | |
| GridRun( | |
| "causal_forcing", | |
| "none", | |
| int(state["prompt_id"]), | |
| anchors, | |
| 3, | |
| features, | |
| path, | |
| ) | |
| ) | |
| del state | |
| return runs | |
| def project_hy_run( | |
| snapshot_path: Path, | |
| cache_path: Path, | |
| projection_dim: int, | |
| device: torch.device, | |
| overwrite: bool, | |
| ) -> dict[tuple[int, int], torch.Tensor]: | |
| if cache_path.exists() and not overwrite: | |
| state = torch.load(cache_path, map_location="cpu", weights_only=False) | |
| return { | |
| tuple(int(value) for value in key.split(":")): tensor.float() | |
| for key, tensor in state["features"].items() | |
| } | |
| data = np.load(snapshot_path, allow_pickle=False) | |
| stages = data["stages"].astype(str) | |
| chunks = data["chunks"].astype(int) | |
| steps = data["steps"].astype(int) | |
| coords = data["coords"].astype(int) | |
| expected_coords = np.stack( | |
| np.meshgrid(np.arange(4), np.arange(GRID_H), np.arange(GRID_W), indexing="ij"), | |
| axis=-1, | |
| ).reshape(-1, 3) | |
| if coords.shape != expected_coords.shape or not np.array_equal(coords, expected_coords): | |
| raise ValueError(f"Unexpected HY coordinate order in {snapshot_path}") | |
| feature_array = data["features"] | |
| projection = projection_matrix(int(feature_array.shape[-1]), projection_dim, device) | |
| features: dict[tuple[int, int], torch.Tensor] = {} | |
| selected = np.flatnonzero(stages == "block_53") | |
| for position, index in enumerate(selected): | |
| value = torch.from_numpy(np.asarray(feature_array[index])).to(device=device, dtype=torch.float32) | |
| value = torch.matmul(value, projection).reshape(4, GRID_H, GRID_W, projection_dim) | |
| features[(int(chunks[index]), int(steps[index]))] = value.to("cpu", torch.float16) | |
| if position % 4 == 3: | |
| print(f"[HY projection] {snapshot_path.parent.name}: {position + 1}/{len(selected)}", flush=True) | |
| data.close() | |
| cache_path.parent.mkdir(parents=True, exist_ok=True) | |
| torch.save( | |
| { | |
| "source": str(snapshot_path), | |
| "projection_dim": projection_dim, | |
| "features": {f"{chunk}:{step}": value for (chunk, step), value in features.items()}, | |
| }, | |
| cache_path, | |
| ) | |
| return {key: value.float() for key, value in features.items()} | |
| def load_hy_runs( | |
| root: Path, | |
| action: str, | |
| cache_root: Path, | |
| projection_dim: int, | |
| device: torch.device, | |
| overwrite: bool, | |
| ) -> list[GridRun]: | |
| runs: list[GridRun] = [] | |
| for case_dir in sorted((root / "runs").glob("prompt_*")): | |
| run_dir = case_dir / action | |
| snapshot = run_dir / "dense_selected_snapshots.npz" | |
| anchor_path = run_dir / "rgb_anchor_frames.npz" | |
| if not snapshot.exists() or not anchor_path.exists(): | |
| continue | |
| prompt_id = int(case_dir.name.split("_")[-1]) | |
| cache_path = cache_root / action / f"prompt_{prompt_id:04d}.pt" | |
| features = project_hy_run(snapshot, cache_path, projection_dim, device, overwrite) | |
| anchors = np.load(anchor_path, allow_pickle=False)["frames"] | |
| runs.append( | |
| GridRun( | |
| "hy_worldplay", | |
| action, | |
| prompt_id, | |
| anchors, | |
| 4, | |
| features, | |
| run_dir, | |
| ) | |
| ) | |
| return runs | |
| def shuffled_flow(flow: torch.Tensor, seed: int) -> torch.Tensor: | |
| generator = torch.Generator(device="cpu").manual_seed(seed) | |
| permutation = torch.randperm(GRID_H * GRID_W, generator=generator) | |
| return flow.reshape(2, -1)[:, permutation].reshape_as(flow) | |
| def collect_rows(runs: list[GridRun]) -> list[dict[str, Any]]: | |
| rows: list[dict[str, Any]] = [] | |
| for run_index, run in enumerate(runs): | |
| for chunk in range(1, run.chunks): | |
| source_frame_index = chunk * run.chunk_size - 1 | |
| source_frame = run.anchors[source_frame_index] | |
| boundary_flow = farneback(source_frame, run.anchors[chunk * run.chunk_size]) | |
| median = np.median(boundary_flow.reshape(-1, 2), axis=0) | |
| residual = boundary_flow - median[None, None] | |
| motion = { | |
| "total_motion": float(np.linalg.norm(boundary_flow, axis=-1).mean()), | |
| "camera_motion": float(np.linalg.norm(median)), | |
| "object_motion": float(np.linalg.norm(residual, axis=-1).mean()), | |
| } | |
| for step in run.steps: | |
| source_map = run.features[(chunk - 1, step)][-1].float() | |
| target_maps = run.features[(chunk, step)].float() | |
| slot_rows: list[dict[str, float]] = [] | |
| for slot in range(run.chunk_size): | |
| target = target_maps[slot] | |
| 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]))) | |
| controls = { | |
| "global": global_flow, | |
| "flow": flow, | |
| "negated": -flow, | |
| "shuffled": shuffled_flow( | |
| flow, | |
| seed=(run.prompt_id + 1) * 100000 + chunk * 1000 + slot * 10 + step, | |
| ), | |
| } | |
| values = {"raw_cosine": cosine(target, source_map)} | |
| valid_ratios = [] | |
| for name, control_flow in controls.items(): | |
| aligned, mask = warp(source_map, control_flow) | |
| values[f"{name}_aligned_cosine"] = cosine(target, aligned, mask) | |
| valid_ratios.append(float(mask.float().mean())) | |
| values["valid_flow_ratio"] = mean(valid_ratios) | |
| slot_rows.append(values) | |
| row = { | |
| "model": run.model, | |
| "action": run.action, | |
| "prompt_id": run.prompt_id, | |
| "chunk": chunk, | |
| "step": step, | |
| "source": str(run.source), | |
| **motion, | |
| } | |
| for key in slot_rows[0]: | |
| row[key] = mean(item[key] for item in slot_rows) | |
| for name in ("global", "flow", "negated", "shuffled"): | |
| row[f"{name}_gain"] = row[f"{name}_aligned_cosine"] - row["raw_cosine"] | |
| row["flow_over_global"] = row["flow_aligned_cosine"] - row["global_aligned_cosine"] | |
| row["flow_over_shuffled"] = row["flow_aligned_cosine"] - row["shuffled_aligned_cosine"] | |
| rows.append(row) | |
| print(f"[analysis] {run.model}/{run.action} prompt {run.prompt_id}: {run_index + 1}/{len(runs)}", flush=True) | |
| return rows | |
| def add_motion_bins(rows: list[dict[str, Any]]) -> None: | |
| groups: dict[tuple[str, str], list[dict[str, Any]]] = defaultdict(list) | |
| for row in rows: | |
| groups[(row["model"], row["action"])].append(row) | |
| for values in groups.values(): | |
| low, high = np.quantile([row["total_motion"] for row in values], [1 / 3, 2 / 3]) | |
| for row in values: | |
| row["motion_bin"] = ( | |
| "low" if row["total_motion"] <= low else "high" if row["total_motion"] > high else "medium" | |
| ) | |
| METRICS = [ | |
| "total_motion", | |
| "camera_motion", | |
| "object_motion", | |
| "raw_cosine", | |
| "global_aligned_cosine", | |
| "flow_aligned_cosine", | |
| "negated_aligned_cosine", | |
| "shuffled_aligned_cosine", | |
| "global_gain", | |
| "flow_gain", | |
| "negated_gain", | |
| "shuffled_gain", | |
| "flow_over_global", | |
| "flow_over_shuffled", | |
| "valid_flow_ratio", | |
| ] | |
| def summarize(rows: list[dict[str, Any]], keys: list[str]) -> list[dict[str, Any]]: | |
| groups: dict[tuple[Any, ...], list[dict[str, Any]]] = defaultdict(list) | |
| for row in rows: | |
| groups[tuple(row[key] for key in keys)].append(row) | |
| output = [] | |
| for group, values in sorted(groups.items(), key=lambda item: tuple(map(str, item[0]))): | |
| item = {key: value for key, value in zip(keys, group)} | |
| item["count"] = len(values) | |
| for metric in METRICS: | |
| item[metric] = mean(row[metric] for row in values) | |
| item["flow_win_fraction"] = mean(row["flow_gain"] > 0 for row in values) | |
| item["flow_beats_shuffled_fraction"] = mean( | |
| row["flow_aligned_cosine"] > row["shuffled_aligned_cosine"] for row in values | |
| ) | |
| motion = np.asarray([row["total_motion"] for row in values], dtype=np.float64) | |
| raw = np.asarray([row["raw_cosine"] for row in values], dtype=np.float64) | |
| item["motion_raw_pearson"] = ( | |
| float(np.corrcoef(motion, raw)[0, 1]) if len(values) >= 3 and np.std(motion) > 0 else float("nan") | |
| ) | |
| output.append(item) | |
| return output | |
| def plot(rows: list[dict[str, Any]], output: Path) -> None: | |
| groups = sorted({(row["model"], row["action"]) for row in rows}) | |
| names = [f"{model}\n{action}" for model, action in groups] | |
| methods = ["raw_cosine", "global_aligned_cosine", "flow_aligned_cosine", "negated_aligned_cosine", "shuffled_aligned_cosine"] | |
| labels = ["raw", "global", "correct flow", "negated", "shuffled"] | |
| fig, axes = plt.subplots(1, 2, figsize=(15, 5.5)) | |
| x = np.arange(len(groups)) | |
| width = 0.15 | |
| for index, (metric, label) in enumerate(zip(methods, labels)): | |
| values = [mean(row[metric] for row in rows if (row["model"], row["action"]) == group) for group in groups] | |
| axes[0].bar(x + (index - 2) * width, values, width=width, label=label) | |
| axes[0].set_xticks(x, names) | |
| axes[0].set_ylabel("cosine") | |
| axes[0].set_title("Full-grid bilinear alignment") | |
| axes[0].legend(fontsize=8) | |
| for group in groups: | |
| selected = [row for row in rows if (row["model"], row["action"]) == group] | |
| axes[1].scatter( | |
| [row["total_motion"] for row in selected], | |
| [row["flow_gain"] for row in selected], | |
| s=14, | |
| alpha=0.5, | |
| label="/".join(group), | |
| ) | |
| axes[1].axhline(0, color="black", linewidth=1) | |
| axes[1].set_xlabel("motion magnitude") | |
| axes[1].set_ylabel("correct-flow cosine gain") | |
| axes[1].set_title("Alignment gain vs motion") | |
| axes[1].legend(fontsize=7) | |
| fig.tight_layout() | |
| fig.savefig(output / "fullgrid_bilinear_alignment.png", dpi=180) | |
| plt.close(fig) | |
| def markdown_table(rows: list[dict[str, Any]]) -> str: | |
| columns = [ | |
| "model", | |
| "action", | |
| "motion_bin", | |
| "count", | |
| "raw_cosine", | |
| "global_aligned_cosine", | |
| "flow_aligned_cosine", | |
| "negated_aligned_cosine", | |
| "shuffled_aligned_cosine", | |
| "flow_gain", | |
| "flow_over_shuffled", | |
| "flow_win_fraction", | |
| ] | |
| lines = ["| " + " | ".join(columns) + " |", "|" + "|".join("---" for _ in columns) + "|"] | |
| for row in rows: | |
| cells = [] | |
| for column in columns: | |
| value = row.get(column, "") | |
| cells.append(f"{value:.4f}" if isinstance(value, float) and np.isfinite(value) else str(value)) | |
| lines.append("| " + " | ".join(cells) + " |") | |
| return "\n".join(lines) | |
| def main() -> None: | |
| args = parse_args() | |
| output = args.output_root.resolve() | |
| output.mkdir(parents=True, exist_ok=True) | |
| device = torch.device(args.projection_device) | |
| rows: list[dict[str, Any]] = [] | |
| self_runs = load_self_runs(args.self_root.resolve()) | |
| print(f"[load] Self runs: {len(self_runs)}", flush=True) | |
| rows.extend(collect_rows(self_runs)) | |
| del self_runs | |
| causal_runs = load_causal_runs(args.causal_root.resolve()) | |
| print(f"[load] Causal runs: {len(causal_runs)}", flush=True) | |
| rows.extend(collect_rows(causal_runs)) | |
| del causal_runs | |
| cache_root = output / "projected_cache" / "hy_worldplay" | |
| for action, root in ( | |
| ("static", args.hy_root.resolve()), | |
| ("forward", args.hy_root.resolve()), | |
| ("right", args.hy_right_root.resolve()), | |
| ): | |
| runs = load_hy_runs( | |
| root, | |
| action, | |
| cache_root, | |
| args.projection_dim, | |
| device, | |
| args.overwrite_projection_cache, | |
| ) | |
| print(f"[load] HY {action} runs: {len(runs)}", flush=True) | |
| rows.extend(collect_rows(runs)) | |
| del runs | |
| if device.type == "cuda": | |
| torch.cuda.empty_cache() | |
| add_motion_bins(rows) | |
| overall = summarize(rows, ["model", "action"]) | |
| bins = summarize(rows, ["model", "action", "motion_bin"]) | |
| chunks = summarize(rows, ["model", "action", "chunk"]) | |
| write_csv(output / "fullgrid_bilinear_metrics.csv", rows) | |
| write_csv(output / "fullgrid_bilinear_summary.csv", overall) | |
| write_csv(output / "fullgrid_bilinear_motion_bins.csv", bins) | |
| write_csv(output / "fullgrid_bilinear_chunks.csv", chunks) | |
| plot(rows, output) | |
| metadata = { | |
| "projection_dim": args.projection_dim, | |
| "projection_seed_base": PROJECTION_SEED_BASE, | |
| "grid": [GRID_H, GRID_W], | |
| "flow": "Farneback target-to-source, resized with vector scaling", | |
| "warp": "bilinear grid_sample, align_corners=True, in-bounds mask", | |
| "motion_bins": "tertiles computed independently inside each model/action", | |
| "run_count": len({(row["model"], row["action"], row["prompt_id"]) for row in rows}), | |
| "row_count": len(rows), | |
| } | |
| (output / "analysis_config.json").write_text(json.dumps(metadata, indent=2) + "\n", encoding="utf-8") | |
| report = [ | |
| "# Three-backbone full-grid bilinear alignment", | |
| "", | |
| "Primary comparison follows the original Self-Forcing pilot and keeps a complete 30x52 feature grid. Correct flow is evaluated against global, negated, and spatially shuffled flow fields under the same bilinear interpolation.", | |
| "", | |
| markdown_table(bins), | |
| "", | |
| "A correct-flow gain alone can include interpolation effects. `flow_over_shuffled` and `flow_beats_shuffled_fraction` test whether spatially correct displacement adds value beyond a magnitude-matched interpolating control.", | |
| "", | |
| "In the current results, Self-Forcing and Causal-Forcing retain material correct-flow advantages over shuffled flow (0.0040 and 0.0075 cosine overall). HY's static/forward/right advantages are only about 0.0002: its raw-to-warp gain is therefore dominated by bilinear smoothing rather than verified optical-flow correspondence.", | |
| "", | |
| "HY has only two target chunk boundaries. Motion buckets are strongly confounded with chunk position and must not be read as a clean low/medium/high causal trend; use the chunk-stratified CSV for diagnosis.", | |
| ] | |
| (output / "REPORT.md").write_text("\n".join(report) + "\n", encoding="utf-8") | |
| print(f"[complete] {output}: {len(rows)} rows", flush=True) | |
| if __name__ == "__main__": | |
| main() | |