Self-Forcing / scripts /generate_vbench8_extended_spanrisk.py
Cccccz's picture
Upload Python scripts
bc29ee3 verified
Raw History Blame Contribute Delete
10 kB
#!/usr/bin/env python3
"""Generate one outgoing-span Confidence-Head strategy for Extended-251.
The FFFF baseline is reused from the established Extended-251 run. Pixel
metrics decode both MP4s, so reference and prediction receive identical video
encoding/decoding treatment.
"""
from __future__ import annotations
import argparse
import json
import sys
from pathlib import Path
from typing import Any
REPO_ROOT = Path(__file__).resolve().parents[1]
if str(REPO_ROOT) not in sys.path:
sys.path.insert(0, str(REPO_ROOT))
from scripts import generate_vbench8_extended_strategies as extended
import torch
from torchvision.io import read_video
from scripts import evaluate_single_block_fppf as base
MAPPING_DEFAULT = REPO_ROOT / "assets/vbench8_extended_subset_mapping.json"
EXPERIMENT_DEFAULT = (
REPO_ROOT / "confidence_experiments/layer17_stage1_step2000_20260831"
)
SELECTION_DEFAULT = (
REPO_ROOT
/ "confidence_experiments"
/ "layer17_stage1_step2000_spanrisk_beta2_threshold_vbench3_20260901"
/ "threshold_summary.json"
)
REFERENCE_DEFAULT = (
REPO_ROOT / "evaluation_runs/vbench8_extended_stage1_step2000_20260901"
)
OUTPUT_DEFAULT = (
REPO_ROOT
/ "evaluation_runs/vbench8_extended_stage1_step2000_spanrisk_beta2_20260901"
)
STRATEGIES = (
"step12_span_k06",
"step12_span_k08",
"step12_span_k10",
"step123_span_k06",
"step123_span_k09",
"step123_span_k12",
"step123_span_k15",
)
def load_strategy_configs(path: Path) -> dict[str, dict[str, Any]]:
payload = json.loads(path.read_text(encoding="utf-8"))
if payload.get("status") != "complete" or float(payload["beta"]) != 2.0:
raise ValueError(f"Threshold selection is not completed fixed-beta=2: {path}")
configs: dict[str, dict[str, Any]] = {}
for row in payload["selected"]:
head = str(row["head"])
target = int(row["target_accepts"])
name = f"{head}_span_k{target:02d}"
configs[name] = {
"name": name,
"policy": "dynamic",
"candidate_steps": [int(step) for step in row["candidate_steps"]],
"beta": float(row["beta"]),
"threshold": float(row["threshold"]),
"target_accepts": target,
"head": head,
"risk_mode": "outgoing_span",
"source_config_name": str(row["name"]),
}
if tuple(configs) != STRATEGIES:
raise ValueError(f"Unexpected selected strategies: {tuple(configs)}")
return configs
def read_u8_video(path: Path) -> torch.Tensor:
if not path.is_file():
raise FileNotFoundError(path)
frames, _, _ = read_video(str(path), pts_unit="sec", output_format="TCHW")
frames = frames.to(device="cpu", dtype=torch.uint8).contiguous()
if frames.ndim != 4 or frames.shape[0] != 81 or frames.shape[1] != 3:
raise RuntimeError(f"Unexpected decoded video shape {tuple(frames.shape)}: {path}")
return frames
def artifact_paths(
output_root: Path, mapping_row: dict[str, Any], strategy_name: str
) -> tuple[Path, Path]:
suite = str(mapping_row["prompt_suite"])
suite_index = int(mapping_row["suite_index"])
global_index = int(mapping_row["global_index"])
return (
output_root / "generated_videos" / strategy_name / suite / f"{suite_index:03d}.mp4",
output_root
/ "generation_metrics/per_prompt"
/ strategy_name
/ f"global_{global_index:04d}.json",
)
def run_prompt(
*,
mapping_row: dict[str, Any],
strategy: dict[str, Any],
reference_root: Path,
output_root: Path,
seed: int,
models: tuple[Any, ...],
device: torch.device,
) -> None:
global_index = int(mapping_row["global_index"])
suite = str(mapping_row["prompt_suite"])
suite_index = int(mapping_row["suite_index"])
prompt = str(mapping_row["extended_prompt"])
name = str(strategy["name"])
reference_path = (
reference_root / "generated_videos/ffff" / suite / f"{suite_index:03d}.mp4"
)
print(f"[prompt] global={global_index} suite={suite}/{suite_index}", flush=True)
conditional = models[3](text_prompts=[prompt])
head = models[4] if strategy["head"] == "step12" else models[5]
latent, diagnostic = extended.generate_rollout(
pipeline=models[1],
conditional_dict=conditional,
seed=seed,
device=device,
predictor=models[2],
head=head,
config=strategy,
)
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
pixels = models[0].decode_to_pixel(latent, use_cache=False)
prediction_u8 = base.pixels_to_u8(pixels)
video_path, record_path = artifact_paths(output_root, mapping_row, name)
extended.atomic_video(prediction_u8, video_path)
reference_decoded = read_u8_video(reference_path)
prediction_decoded = read_u8_video(video_path)
metrics = base.frame_metrics(
reference_u8=reference_decoded,
prediction_u8=prediction_decoded,
lpips_model=models[6],
batch_size=4,
device=device,
)
record = {
"status": "complete",
"strategy": name,
"policy": "dynamic",
"candidate_steps": strategy["candidate_steps"],
"beta": strategy["beta"],
"threshold": strategy["threshold"],
"target_accepts": strategy["target_accepts"],
"risk_mode": "outgoing_span",
"allow_chunk0_predictor": False,
"source_config_name": strategy["source_config_name"],
"global_index": global_index,
"prompt_suite": suite,
"suite_index": suite_index,
"prompt": prompt,
"seed": seed,
"generation": {
key: value for key, value in diagnostic.items() if key != "decisions"
},
"decisions": diagnostic["decisions"],
"pixel_metrics_vs_ffff": extended.compact_metrics(metrics),
"pixel_metric_input": "MP4-decoded uint8 RGB, all 81 frames on both sides",
"reference_video": str(reference_path),
"video": str(video_path.relative_to(output_root)),
}
extended.atomic_json(record_path, record)
print(
f"[result] global={global_index} strategy={name} "
f"accept={diagnostic['accepted_predictor_calls']} "
f"psnr={metrics['psnr']:.4f} ssim={metrics['ssim']:.6f} "
f"lpips={metrics['lpips']:.6f}",
flush=True,
)
if hasattr(models[0].model, "clear_cache"):
models[0].model.clear_cache()
del conditional, latent, pixels, prediction_u8, reference_decoded, prediction_decoded
torch.cuda.empty_cache()
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--gpu", default=extended.PHYSICAL_GPU)
parser.add_argument("--strategy", required=True, choices=STRATEGIES)
parser.add_argument("--mapping", type=Path, default=MAPPING_DEFAULT)
parser.add_argument("--experiment-root", type=Path, default=EXPERIMENT_DEFAULT)
parser.add_argument("--selection", type=Path, default=SELECTION_DEFAULT)
parser.add_argument("--reference-root", type=Path, default=REFERENCE_DEFAULT)
parser.add_argument("--output-root", type=Path, default=OUTPUT_DEFAULT)
parser.add_argument("--global-index", action="append", type=int, default=None)
parser.add_argument("--seed", type=int, default=0)
parser.add_argument("--overwrite", action="store_true")
return parser.parse_args()
def main() -> None:
args = parse_args()
mapping = extended.read_mapping(args.mapping.resolve())
configs = load_strategy_configs(args.selection.resolve())
strategy = configs[args.strategy]
reference_root = args.reference_root.resolve()
if len(list((reference_root / "generation_metrics/per_prompt/ffff").glob("global_*.json"))) != 251:
raise ValueError("Incomplete FFFF reference records")
if len(list((reference_root / "generated_videos/ffff").rglob("*.mp4"))) != 251:
raise ValueError("Incomplete FFFF reference videos")
if args.global_index is not None:
requested = set(args.global_index)
mapping = [row for row in mapping if int(row["global_index"]) in requested]
if len(mapping) != len(requested):
raise ValueError("At least one requested global index is unknown")
output_root = args.output_root.resolve()
output_root.mkdir(parents=True, exist_ok=True)
extended.write_shard_manifest(
output_root, str(args.gpu), 0, 1, mapping, [strategy], "running"
)
print(
f"[setup] physical_gpu={args.gpu} strategy={args.strategy} "
f"prompts={len(mapping)} reference=existing_ffff chunk0=ffff",
flush=True,
)
device = torch.device("cuda")
torch.set_grad_enabled(False)
models = extended.load_models(args.experiment_root.resolve(), device)
completed = 0
try:
for row in mapping:
video_path, record_path = artifact_paths(output_root, row, args.strategy)
if not args.overwrite and video_path.is_file() and record_path.is_file():
completed += 1
print(f"[cached] {completed}/{len(mapping)} global={row['global_index']}", flush=True)
continue
run_prompt(
mapping_row=row,
strategy=strategy,
reference_root=reference_root,
output_root=output_root,
seed=args.seed,
models=models,
device=device,
)
completed += 1
print(f"[progress] {completed}/{len(mapping)} prompts", flush=True)
finally:
extended.write_shard_manifest(
output_root,
str(args.gpu),
0,
1,
mapping,
[strategy],
"complete" if completed == len(mapping) else "failed",
)
print(f"[complete] gpu={args.gpu} strategy={args.strategy} prompts={completed}", flush=True)
if __name__ == "__main__":
main()