#!/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()