Download scripts/generate_vbench8_extended_spanrisk.py from Cccccz/Self-Forcing: direct link, hf CLI and curl.
- Browser
- Download file 10 kB
-
https://huggingface.co/Cccccz/Self-Forcing/resolve/main/scripts/generate_vbench8_extended_spanrisk.py
- Command line
-
hf download hf://Cccccz/Self-Forcing/scripts/generate_vbench8_extended_spanrisk.py
-
curl -L -o generate_vbench8_extended_spanrisk.py https://huggingface.co/Cccccz/Self-Forcing/resolve/main/scripts/generate_vbench8_extended_spanrisk.py
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() | |