comparison / eval /generate_predictor_eval.py
Cccccz's picture
Add files using upload-large-folder tool
f70ac4f verified
Raw History Blame Contribute Delete
28.6 kB
#!/usr/bin/env python
"""Extended-251 videos for the Causal-Forcing-a single-block Predictors.
Same protocol as generate_eval.py: work is sharded by prompt, FFFF is generated
first in the same process and kept as the pixel-metric reference, every strategy
sees the identical noise draw, and the per-prompt record has the layout
aggregate.py / run_vbench.py expect.
Strategies (prefix ``cfa`` = Causal-Forcing-a stack):
cfa_ffff reference, 28 full DiT forwards
cfa_<model>_fpff chunks 1..6: F P F F (22 full + 6 Predictor)
cfa_<model>_fppf chunks 1..6: F P P F (16 full + 12 Predictor)
cfa_<model>_fppp chunks 1..6: F P P P (10 full + 18 Predictor)
where <model> is one of atc_chunk, atc_last_frame, disca (Stage-1 Layer-17
Predictors) or atc_chunk_s2 (Stage-2 DMD EMA). Chunk 0 is always fully denoised.
``<model>@<step>`` selects another checkpoint of that run and names the strategy
cfa_<model>_s<step>..., e.g. ``disca@500`` -> cfa_disca_s500_fppf; a bare model
uses --checkpoint-step (default 2000) and the unsuffixed name.
cfa_naive2step chunks 1..: the 2-step subset of the grid (t = 1000 / 250,
the same subset as cf_naive2step_c0FFFF), no Predictor
CUDA_VISIBLE_DEVICES=3 python eval/generate_predictor_eval.py --shard 0 --num-shards 4
"""
import argparse
import json
import os
import sys
import time
ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
REPO = os.environ.get("CFA_REPO", "/local/zoubin/cz/projects/Causal-Forcing-a")
MAPPING = os.path.join(ROOT, "assets/vbench8_extended_subset_mapping.json")
BASE = {
"repo": REPO,
"config": "configs/causal_forcing_dmd_chunkwise.yaml",
"checkpoint": "checkpoints/chunkwise/causal_forcing.pt",
}
MODELS = {
"atc_chunk": "output/layer17_stage1_atc_chunk_1000p_21f_seed0_4gpu_b16_acc1_2000steps",
"atc_last_frame": "output/layer17_stage1_atc_last_frame_1000p_21f_seed0_4gpu_b16_acc1_2000steps",
"disca": "output/layer17_stage1_disca_1000p_21f_seed0_4gpu_b16_acc1_2000steps",
# Stage-2 random-exit DMD (LoRA critic), EMA weights (a (run dir, file name) pair;
# a plain directory gets checkpoint_step_<N>/predictor.safetensors).
"atc_chunk_s2": ("output/layer17_stage2_dmd_atc_chunk_lora128_4gpu_b2_2000steps", "predictor_ema.safetensors"),
}
NAIVE = {"naive2step": [0, 3]} # pseudo-models: reduced step grid, no Predictor
def checkpoint_weights(model, step):
"""Weights file of ``model`` at training step ``step`` (Stage-1 runs zero-pad the
checkpoint directory, Stage-2 does not)."""
spec = MODELS[model]
run_dir, fname = spec if isinstance(spec, tuple) else (spec, "predictor.safetensors")
for d in (f"checkpoint_step_{step:04d}", f"checkpoint_step_{step}"):
w = os.path.join(REPO, run_dir, d, fname)
if os.path.exists(w):
return w
raise FileNotFoundError(f"{model} step {step}: no checkpoint under {os.path.join(REPO, run_dir)}")
PATTERNS = {"FPFF": {1}, "FPPF": {1, 2}, "FPPP": {1, 2, 3}}
# Reduced schedules: which step indices of the 4-step grid a chunk (after chunk 0)
# actually visits. FPF keeps steps 0 and 3 (t = 1000 / 250) as Full and runs one
# Predictor step at step 1 (t = 750, anchored on step 0 as in training), skipping
# step 2 entirely: 3 denoising positions per chunk, 2 Full + 1 Predictor.
RUN_STEPS = {"FPF": [0, 1, 3]}
PATTERNS["FPF"] = {1}
NUM_CHUNKS, FRAMES_PER_CHUNK, NUM_STEPS = 7, 3, 4
# A Predictor step runs 1 of the 30 Teacher blocks (plus a small fusion MLP);
# its compute-equivalent cost is booked as 1/30 of a full forward.
PREDICTOR_COMPUTE_EQUIV = 1.0 / 30.0
def parse_args():
p = argparse.ArgumentParser()
p.add_argument("--shard", type=int, default=0)
p.add_argument("--num-shards", type=int, default=1)
p.add_argument("--out-root", default="eval_out")
p.add_argument("--seed", type=int, default=0)
p.add_argument("--fps", type=int, default=16)
p.add_argument("--models", default=",".join(MODELS))
p.add_argument("--patterns", default="FPPF,FPPP")
p.add_argument("--checkpoint-step", type=int, default=2000)
p.add_argument("--num-chunks", type=int, default=NUM_CHUNKS,
help="Video length in 3-latent-frame chunks (7 = the 81-frame protocol; longer "
"videos use a rolling 21-frame KV window like Self-Forcing's long-video mode)")
p.add_argument("--ref-tag", default="",
help="Inserted into the reference name, e.g. _14c -> cfa_14c_ffff (a separate FFFF "
"reference for a different video length)")
p.add_argument("--tag", default="",
help="Inserted into strategy names after the model, e.g. _s1000 -> "
"cfa_disca_s1000_fppf, so another checkpoint does not overwrite the default rows")
p.add_argument("--regenerate-ffff", action="store_true",
help="Regenerate the FFFF reference instead of reading cfa_ffff's MP4")
p.add_argument("--limit", type=int, default=None, help="Smoke-test: first N prompts")
p.add_argument("--overwrite", action="store_true")
p.add_argument("--retime-only", action="store_true",
help="Do not regenerate videos/metrics; re-time FFFF and each strategy in the same "
"process per prompt and update policy_latency_ms / matched_ffff_policy_latency_ms.")
p.add_argument("--retime-max-ffff-ms", type=float, default=3500.0,
help="With --retime-only, records whose paired FFFF exceeds this are re-timed again "
"(up to --retime-passes passes); clean records are left alone.")
p.add_argument("--retime-passes", type=int, default=3)
p.add_argument("--retime-strategy-only", action="store_true",
help="Re-time only the strategies (no FFFF run): rerun each selected record's "
"rollout on an idle GPU and replace policy_latency_ms; the generation-run "
"value is kept under generation_run_latency.")
p.add_argument("--written-before", type=float, default=None,
help="With --retime-strategy-only: only records whose file mtime is before this "
"epoch second (the window in which another job shared the GPU)")
return p.parse_args()
def build_strategies(args):
out = [{"name": f"cfa{args.ref_tag}_ffff", "model": None, "pattern": "FFFF", "is_reference": True}]
for token in args.models.split(","):
if token in NAIVE:
out.append({"name": f"cfa_{token}{args.tag}", "model": None, "pattern": token.upper(),
"naive": token, "is_reference": False})
continue
model, _, step = token.partition("@")
step = int(step) if step else args.checkpoint_step
key = f"{model}@{step}" # one loaded Predictor per (model, step)
label = f"{model}_s{step}" if "@" in token else model
for pattern in args.patterns.split(","):
out.append({"name": f"cfa_{label}{args.tag}_{pattern.lower()}", "model": key,
"base_model": model, "checkpoint_step": step,
"pattern": pattern.upper(), "is_reference": False})
return out
class FinalHiddenCapture:
"""Grab the Teacher's final hidden (input of its output head) on a Full step."""
def __init__(self, teacher):
self.enabled, self.value = False, None
self.handle = teacher.head.register_forward_pre_hook(self._hook)
def _hook(self, _module, inputs):
if self.enabled:
self.value = inputs[0].detach()
def start(self):
self.value, self.enabled = None, True
def finish(self):
self.enabled = False
if self.value is None:
raise RuntimeError("Teacher final hidden was not captured")
value, self.value = self.value, None
return value
def main():
args = parse_args()
out_root = args.out_root if os.path.isabs(args.out_root) else os.path.join(ROOT, args.out_root)
with open(MAPPING) as f:
rows = json.load(f)["rows"]
if args.limit:
rows = rows[:args.limit]
shard_rows = rows[args.shard::args.num_shards]
strategies = build_strategies(args)
ref_name = strategies[0]["name"]
print(f"cfa shard {args.shard}/{args.num_shards}: {len(shard_rows)} prompts x "
f"{len(strategies)} strategies", flush=True)
os.chdir(REPO)
sys.path.insert(0, ROOT)
sys.path.insert(0, REPO) # pipeline/, utils/, predictor_training/ from Causal-Forcing-a
import torch
from einops import rearrange
from safetensors.torch import load_file
from torchvision.io import read_video, write_video
from eval.pixel_metrics import PixelMetrics
from harness import load_pipeline
from predictor_training.metadata import build_single_block_predictor, read_predictor_config
from predictor_training.online import online_predictor_step as _online_predictor_step
from predictor_training.offline_data import TOKENS_PER_CHUNK
def online_predictor_step(*, history_cache, chunk, **kw):
"""Stage-1 contract: the Predictor sees the clean K/V of the *earlier* chunks.
With the rolling 21-frame window (long videos) the cache holds the last six
clean chunks plus the current noisy one, so the history is everything before
the current chunk's tokens rather than ``chunk * TOKENS_PER_CHUNK`` (which
the stock helper asserts and which exceeds the window from chunk 7 on).
For chunk <= 6 the two are identical."""
if chunk * TOKENS_PER_CHUNK <= int(history_cache["local_end_index"].item()):
return _online_predictor_step(history_cache=history_cache, chunk=chunk, **kw)
end = int(history_cache["local_end_index"].item()) - TOKENS_PER_CHUNK
window = {"k": history_cache["k"][:, :end], "v": history_cache["v"][:, :end],
"local_end_index": torch.tensor([end], device=history_cache["k"].device)}
# left-pad to the absolute length the stock helper slices; the attention
# window (max_attention_size = 21 frames) never reaches the padding
pad = chunk * TOKENS_PER_CHUNK - end
if pad > 0:
z = history_cache["k"].new_zeros(history_cache["k"].shape[0], pad, *history_cache["k"].shape[2:])
window = {"k": torch.cat([z, window["k"]], 1), "v": torch.cat([z, window["v"]], 1),
"local_end_index": torch.tensor([chunk * TOKENS_PER_CHUNK], device=z.device)}
return _online_predictor_step(history_cache=window, chunk=chunk, **kw)
from utils.misc import set_seed
torch.set_grad_enabled(False)
device = torch.device("cuda")
pipeline = load_pipeline(BASE)
n_chunks = args.num_chunks
if n_chunks * FRAMES_PER_CHUNK > 21:
# Longer than the 21-latent-frame training context: roll the KV cache over a
# 21-frame local window (the model's long-video mechanism; RoPE is relative
# so positions inside the window stay in the trained range).
model = pipeline.generator.model
for m in [model] + [b for b in model.blocks] + [b.self_attn for b in model.blocks]:
m.local_attn_size = 21
if hasattr(m, "max_attention_size"):
m.max_attention_size = 21 * pipeline.frame_seq_length
pipeline.local_attn_size = 21
pipeline.kv_cache1 = None
print(f"long video: {n_chunks} chunks = {n_chunks * FRAMES_PER_CHUNK} latent frames, "
f"rolling KV window 21 frames", flush=True)
teacher = pipeline.generator.model
capture = FinalHiddenCapture(teacher)
metrics = PixelMetrics(device)
timesteps = pipeline.denoising_step_list.to(device)
assert len(timesteps) == NUM_STEPS and pipeline.num_frame_per_block == FRAMES_PER_CHUNK
predictors, predictor_configs = {}, {}
for model in {s["model"] for s in strategies if s["model"]}:
base, step = model.split("@")
weights = checkpoint_weights(base, int(step))
config = read_predictor_config(weights)
predictor = build_single_block_predictor(teacher, config, gradient_checkpointing=False)
predictor.load_state_dict(load_file(weights, device="cpu"), strict=True)
predictors[model] = predictor.to(device).eval().requires_grad_(False)
predictor_configs[model] = {**config, "weights": weights}
print(f"loaded {model}: {config}", flush=True)
def reset_caches():
if pipeline.kv_cache1 is None:
pipeline._initialize_kv_cache(1, torch.bfloat16, device)
pipeline._initialize_crossattn_cache(1, torch.bfloat16, device)
for cache in pipeline.kv_cache1:
cache["global_end_index"] = torch.tensor([0], dtype=torch.long, device=device)
cache["local_end_index"] = torch.tensor([0], dtype=torch.long, device=device)
for cache in pipeline.crossattn_cache:
cache["is_init"] = False
def rollout(conditional, noise, predictor=None, predictor_steps=frozenset(), run_steps=None):
"""Mirror of cachelib.runner.cached_inference with Predictor steps spliced in."""
reset_caches()
output = torch.zeros_like(noise)
denoise_events, context_events = [], []
full_calls = predictor_calls = 0
previous_hidden = None
for chunk in range(n_chunks):
start_frame = chunk * FRAMES_PER_CHUNK
token_start = start_frame * pipeline.frame_seq_length
noisy_input = noise[:, start_frame:start_frame + FRAMES_PER_CHUNK]
current_hidden = [None] * NUM_STEPS
# chunk 0 always runs the full grid; later chunks may visit a subset
steps = list(range(NUM_STEPS)) if (run_steps is None or chunk == 0) else list(run_steps)
for si, step in enumerate(steps):
timestep = torch.ones([1, FRAMES_PER_CHUNK], device=device, dtype=torch.int64) * timesteps[step]
start_ev = torch.cuda.Event(enable_timing=True)
end_ev = torch.cuda.Event(enable_timing=True)
start_ev.record()
if predictor is not None and chunk > 0 and step in predictor_steps:
hidden, flow = online_predictor_step(
predictor=predictor, teacher=teacher, noisy_input=noisy_input,
timestep=timestep, anchor_timestep=torch.ones_like(timestep) * timesteps[step - 1],
anchor_hidden=current_hidden[step - 1], previous_hidden=previous_hidden[step],
history_cache=pipeline.kv_cache1[predictor.source_layer],
cross_cache=pipeline.crossattn_cache[predictor.source_layer], chunk=chunk)
denoised_pred = pipeline.generator._convert_flow_pred_to_x0(
flow_pred=flow.flatten(0, 1), xt=noisy_input.flatten(0, 1),
timestep=timestep.flatten(0, 1)).unflatten(0, flow.shape[:2])
current_hidden[step] = hidden
predictor_calls += 1
else:
capture.start()
_, denoised_pred = pipeline.generator(
noisy_image_or_video=noisy_input, conditional_dict=conditional,
timestep=timestep, kv_cache=pipeline.kv_cache1,
crossattn_cache=pipeline.crossattn_cache, current_start=token_start)
current_hidden[step] = capture.finish()
full_calls += 1
end_ev.record()
denoise_events.append((start_ev, end_ev))
if si < len(steps) - 1:
# Same RNG draw order as the protocol runner, so every strategy
# re-noises with identical samples.
noisy_input = pipeline.scheduler.add_noise(
denoised_pred.flatten(0, 1), torch.randn_like(denoised_pred.flatten(0, 1)),
timesteps[steps[si + 1]] * torch.ones([FRAMES_PER_CHUNK], device=device, dtype=torch.long),
).unflatten(0, denoised_pred.shape[:2])
output[:, start_frame:start_frame + FRAMES_PER_CHUNK] = denoised_pred
# Clean-context KV refresh: full DiT, timed separately, never added in.
ctx_start = torch.cuda.Event(enable_timing=True)
ctx_end = torch.cuda.Event(enable_timing=True)
ctx_start.record()
pipeline.generator(
noisy_image_or_video=denoised_pred, conditional_dict=conditional,
timestep=torch.ones_like(timestep) * pipeline.args.context_noise,
kv_cache=pipeline.kv_cache1, crossattn_cache=pipeline.crossattn_cache,
current_start=token_start)
ctx_end.record()
context_events.append((ctx_start, ctx_end))
previous_hidden = current_hidden
torch.cuda.synchronize()
timing = {
"denoise_dit_ms": sum(s.elapsed_time(e) for s, e in denoise_events),
"context_kv_dit_ms": sum(s.elapsed_time(e) for s, e in context_events),
"num_denoise_forwards": len(denoise_events),
}
video = pipeline.vae.decode_to_pixel(output, use_cache=False)
video = (video * 0.5 + 0.5).clamp(0, 1)
pipeline.vae.model.clear_cache()
return video, timing, {"full_forwards": full_calls, "predictor_forwards": predictor_calls}
def run_one(strategy, row):
set_seed(args.seed)
noise = torch.randn([1, n_chunks * FRAMES_PER_CHUNK, 16, 60, 104], device=device, dtype=torch.bfloat16)
conditional = pipeline.text_encoder(text_prompts=[row["extended_prompt"]])
wall0 = time.time()
if strategy["is_reference"]:
video, timing, counts = rollout(conditional, noise)
elif strategy.get("naive"):
video, timing, counts = rollout(conditional, noise, None, frozenset(),
NAIVE[strategy["naive"]])
else:
video, timing, counts = rollout(conditional, noise, predictors[strategy["model"]],
PATTERNS[strategy["pattern"]],
RUN_STEPS.get(strategy["pattern"]))
return video, timing, counts, time.time() - wall0
def record_path(strategy, row):
d = os.path.join(out_root, "per_prompt", strategy)
os.makedirs(d, exist_ok=True)
return os.path.join(d, f"{row['prompt_suite']}_{row['suite_index']:03d}.json")
def video_path(strategy, row):
d = os.path.join(out_root, "generated_videos", strategy, row["prompt_suite"])
os.makedirs(d, exist_ok=True)
return os.path.join(d, f"{row['suite_index']:03d}.mp4")
def is_done(strategy, row):
p = record_path(strategy, row)
if args.overwrite or not os.path.exists(p):
return False
try:
with open(p) as f:
rec = json.load(f)
return rec.get("status") == "complete" and os.path.exists(rec.get("video", ""))
except Exception:
return False
def write_json(path, payload):
tmp = path + ".tmp"
with open(tmp, "w") as f:
json.dump(payload, f, indent=2)
os.replace(tmp, path)
if shard_rows:
print("warm-up pass (not recorded)", flush=True)
for strategy in strategies:
run_one(strategy, shard_rows[0])
torch.cuda.empty_cache()
if args.retime_strategy_only:
targets = [s for s in strategies if not s["is_reference"]]
n_done = 0
for n, row in enumerate(shard_rows):
for strategy in targets:
rp = record_path(strategy["name"], row)
if not os.path.exists(rp):
continue
if args.written_before and os.path.getmtime(rp) >= args.written_before:
continue
rec = json.load(open(rp))
if rec.get("status") != "complete" or "generation_run_latency" in rec:
continue
_, timing, counts, _ = run_one(strategy, row)
assert counts["full_forwards"] == rec["cache_diagnostics"]["full_forwards"], rp
rec["generation_run_latency"] = {
"policy_latency_ms": rec["policy_latency_ms"],
"excluded_context_kv_latency_ms": rec["excluded_context_kv_latency_ms"]}
rec["policy_latency_ms"] = timing["denoise_dit_ms"]
rec["excluded_context_kv_latency_ms"] = timing["context_kv_dit_ms"]
rec["latency_retimed"] = "strategy-only rerun on an idle GPU (FFFF not re-run)"
write_json(rp, rec)
n_done += 1
if (n + 1) % 20 == 0:
print(f"[strategy retime shard {args.shard}] {n + 1}/{len(shard_rows)}", flush=True)
print(f"shard {args.shard}: {n_done} records re-timed", flush=True)
return
if args.retime_only:
# Protocol "matched FFFF": pair every strategy with an FFFF run of the same
# prompt in the same process, seconds apart, so contention cancels in the
# ratio. Only records whose pairing is missing or dirty are touched.
targets = [s for s in strategies if not s["is_reference"]]
for pass_index in range(args.retime_passes):
dirty = 0
for n, row in enumerate(shard_rows):
todo = []
for strategy in targets:
rp = record_path(strategy["name"], row)
if not os.path.exists(rp):
continue
rec = json.load(open(rp))
m = rec.get("matched_ffff_policy_latency_ms")
if m is None or m > args.retime_max_ffff_ms:
todo.append((strategy, rp, rec))
if not todo:
continue
dirty += len(todo)
_, ref_timing, _, _ = run_one(strategies[0], row)
for strategy, rp, rec in todo:
_, timing, _, _ = run_one(strategy, row)
rec["policy_latency_ms"] = timing["denoise_dit_ms"]
rec["excluded_context_kv_latency_ms"] = timing["context_kv_dit_ms"]
rec["matched_ffff_policy_latency_ms"] = ref_timing["denoise_dit_ms"]
rec["matched_ffff_context_kv_latency_ms"] = ref_timing["context_kv_dit_ms"]
rec["latency_source"] = "paired retime (same process as FFFF)"
write_json(rp, rec)
if (n + 1) % 10 == 0:
print(f"[retime pass {pass_index} shard {args.shard}] {n + 1}/{len(shard_rows)} "
f"ffff={ref_timing['denoise_dit_ms']:.0f}ms", flush=True)
print(f"[retime pass {pass_index} shard {args.shard}] re-timed {dirty} records", flush=True)
if dirty == 0:
break
return
t_start = time.time()
for n, row in enumerate(shard_rows):
todo = [s for s in strategies if not is_done(s["name"], row)]
if not todo:
continue
need_ref = any(not s["is_reference"] for s in todo)
ref_frames = None
ref_timing = None
ref_from_mp4 = False
if need_ref and os.path.exists(video_path(ref_name, row)) and not args.regenerate_ffff:
# Reference from the cfa_ffff run's MP4 (H.264, 40+ dB above the strategy
# PSNRs); FFFF is not regenerated, so there is no same-process FFFF timing
# and the speedup is taken against cfa_ffff's own records (aggregate.py).
vid, _, _ = read_video(video_path(ref_name, row), pts_unit="sec", output_format="TCHW")
ref_frames = vid.to(device, torch.float16) / 255.0
ref_from_mp4 = True
need_ref = False
for strategy in strategies:
if strategy["is_reference"]:
if strategy not in todo and not need_ref:
continue
elif strategy not in todo:
continue
video, timing, counts, wall = run_one(strategy, row)
frames = video[0]
if strategy["is_reference"]:
ref_frames = frames.detach().to(torch.float16)
ref_timing = timing
pix = PixelMetrics.identity(frames.shape[0])
else:
pix = metrics.compute(frames, ref_frames)
if strategy in todo:
vp = video_path(strategy["name"], row)
arr = (255.0 * rearrange(frames, "t c h w -> t h w c")).clamp(0, 255)
write_video(vp, arr.to(torch.uint8).cpu(), fps=args.fps)
compute_equiv = counts["full_forwards"] + PREDICTOR_COMPUTE_EQUIV * counts["predictor_forwards"]
model = strategy["model"]
naive = strategy.get("naive")
schedule = "F" * len(NAIVE[naive]) if naive else strategy["pattern"]
write_json(record_path(strategy["name"], row), {
"status": "complete",
"protocol": "Self-Forcing Extended-251 Full Evaluation",
"strategy": strategy["name"],
"base_model": "causal_forcing_a",
"method": "none" if model is None else f"predictor_{strategy['base_model']}",
"param": None if model is None else "pattern",
"param_value": None if model is None else strategy["pattern"],
"target_speedup": (n_chunks * NUM_STEPS) / compute_equiv,
"schedule": schedule,
"num_inference_steps": len(NAIVE[naive]) if naive else NUM_STEPS,
"first_chunk_schedule": "FFFF" if (naive or model) else None,
"predictor_config": None if model is None else predictor_configs[model],
"global_index": row["global_index"],
"prompt_suite": row["prompt_suite"],
"suite_index": row["suite_index"],
"prompt": row["extended_prompt"],
"original_prompt": row["original_prompt"],
"seed": args.seed,
"video": vp,
"num_frames": int(frames.shape[0]),
"height": int(frames.shape[2]),
"width": int(frames.shape[3]),
"fps": args.fps,
"policy_latency_ms": timing["denoise_dit_ms"],
"excluded_context_kv_latency_ms": timing["context_kv_dit_ms"],
# Same-process FFFF timing of this prompt (protocol "matched FFFF").
"matched_ffff_policy_latency_ms": ref_timing["denoise_dit_ms"] if ref_timing else None,
"matched_ffff_context_kv_latency_ms": ref_timing["context_kv_dit_ms"] if ref_timing else None,
"checkpoint_step": strategy.get("checkpoint_step"),
"num_chunks": n_chunks,
"wall_generation_s": wall,
"reference_strategy": ref_name,
"pixel_metrics_vs_ffff": pix,
"reference_source": "ffff_mp4" if ref_from_mp4 else "ffff_frames",
"cache_diagnostics": {
"denoise_forwards": timing["num_denoise_forwards"],
"full_forwards": counts["full_forwards"],
"predictor_forwards": counts["predictor_forwards"],
"predictor_compute_equivalent": PREDICTOR_COMPUTE_EQUIV,
"compute_equivalent_forwards": compute_equiv,
"middle_steps": n_chunks * 2,
"middle_compute_equivalent": compute_equiv - n_chunks * 2.0,
},
})
del frames, video
ref_frames = None
done = n + 1
rate = (time.time() - t_start) / done
print(f"[cfa shard {args.shard}] {done}/{len(shard_rows)} prompts "
f"({row['prompt_suite']}/{row['suite_index']:03d}) {rate:.1f}s/prompt "
f"eta {(len(shard_rows) - done) * rate / 60:.0f} min", flush=True)
print(f"shard {args.shard} done in {(time.time() - t_start) / 60:.1f} min", flush=True)
if __name__ == "__main__":
main()