#!/usr/bin/env python3 """Build one resumable worker shard of Self-Forcing Predictor v4 data.""" from __future__ import annotations import argparse import gc import hashlib import json import os import sys from pathlib import Path from typing import Any import torch ROOT = Path(__file__).resolve().parents[1] if str(ROOT) not in sys.path: sys.path.insert(0, str(ROOT)) from predictor_data import ( # noqa: E402 CANDIDATE_BLOCK_IDS, CHUNK_FRAMES, NUM_CHUNKS, NUM_STEPS, PredictorV4DatasetWriter, PredictorV4TeacherCapture, ) from predictor_data.capture import clone_bf16_cpu # noqa: E402 from predictor_data.kv_validation import validate_against_live_cache # noqa: E402 from predictor_data.writer import atomic_write_text # noqa: E402 DEFAULT_ROOT = Path( "/mnt/local_nvme/zoubin/cz/self_forcing_predictor_v4_1000_seed0" ) DEFAULT_CONFIG = ROOT / "configs" / "self_forcing_dmd.yaml" DEFAULT_CHECKPOINT = ROOT / "checkpoints" / "self_forcing_dmd.pt" def load_cases(path: Path) -> list[dict[str, Any]]: result: list[dict[str, Any]] = [] with path.open("r", encoding="utf-8") as handle: for line_number, line in enumerate(handle, start=1): if not line.strip(): continue try: item = json.loads(line) except json.JSONDecodeError as exc: raise ValueError(f"invalid JSON at {path}:{line_number}") from exc if int(item["case_id"]) != len(result): raise ValueError("cases.jsonl case_id values must be dense and ordered") if int(item.get("seed", -1)) != 0: raise ValueError(f"case {item['case_id']} does not use seed 0") result.append(item) return result def parse_ids(value: str | None) -> set[int] | None: if value is None: return None result: set[int] = set() for component in value.split(","): component = component.strip() if not component: continue if "-" in component: start, end = (int(item) for item in component.split("-", 1)) if end < start: raise ValueError(f"invalid case range: {component}") result.update(range(start, end + 1)) else: result.add(int(component)) return result def load_teacher(args, device: torch.device): from omegaconf import OmegaConf from pipeline.causal_inference import CausalInferencePipeline config = OmegaConf.merge( OmegaConf.load(str(ROOT / "configs" / "default_config.yaml")), OmegaConf.load(str(args.config)), ) if bool(getattr(config, "reuse_first_step_velocity", False)): raise ValueError("offline teacher data must use the F-F-F-F schedule") if int(config.num_frame_per_block) != CHUNK_FRAMES: raise ValueError(f"num_frame_per_block must equal {CHUNK_FRAMES}") if bool(config.independent_first_frame): raise ValueError("Predictor v4 T2V data requires independent_first_frame=false") pipeline = CausalInferencePipeline( config, device=device, vae=torch.nn.Identity(), ) checkpoint = torch.load(args.checkpoint, map_location="cpu", weights_only=False) if "generator_ema" not in checkpoint: raise KeyError(f"{args.checkpoint} does not contain generator_ema") pipeline.generator.load_state_dict(checkpoint["generator_ema"], strict=True) del checkpoint pipeline.generator.eval().requires_grad_(False) pipeline.text_encoder.eval().requires_grad_(False) pipeline.generator.to(device=device, dtype=torch.bfloat16) pipeline.text_encoder.to(device=device, dtype=torch.bfloat16) pipeline._initialize_kv_cache(batch_size=1, dtype=torch.bfloat16, device=device) pipeline._initialize_crossattn_cache( batch_size=1, dtype=torch.bfloat16, device=device ) if len(pipeline.generator.model.blocks) != 30: raise ValueError("Predictor v4 builder expects the 30-block Wan 1.3B generator") if pipeline.generator.model.dim != 1536: raise ValueError("Predictor v4 builder expects Wan hidden size 1536") if tuple(int(value) for value in args.blocks) != CANDIDATE_BLOCK_IDS: raise ValueError( f"formal v4 construction requires candidate blocks {CANDIDATE_BLOCK_IDS}" ) return pipeline, config def reset_case_state(pipeline) -> None: for cache in pipeline.kv_cache1: cache["global_end_index"].zero_() cache["local_end_index"].zero_() for cache in pipeline.crossattn_cache: cache["is_init"] = False def case_cross_kv( pipeline, block_ids: tuple[int, ...], ) -> dict[str, torch.Tensor]: result: dict[str, torch.Tensor] = {} for block_id in block_ids: cache = pipeline.crossattn_cache[block_id] if not bool(cache["is_init"]): raise RuntimeError(f"text cross-attention cache for block {block_id} is empty") result[f"block_{block_id:02d}_cross_k"] = clone_bf16_cpu(cache["k"]) result[f"block_{block_id:02d}_cross_v"] = clone_bf16_cpu(cache["v"]) return result def append_jsonl_fsync(path: Path, item: dict[str, Any]) -> None: path.parent.mkdir(parents=True, exist_ok=True) with path.open("a", encoding="utf-8") as handle: handle.write(json.dumps(item, ensure_ascii=False, sort_keys=True) + "\n") handle.flush() os.fsync(handle.fileno()) @torch.inference_mode() def build_case( *, case: dict[str, Any], pipeline, writer: PredictorV4DatasetWriter, capture: PredictorV4TeacherCapture, device: torch.device, validate_kv_rebuild: bool, checkpoint_path: Path, ) -> None: from utils.misc import set_seed case_id = int(case["case_id"]) prompt = str(case["prompt"]) set_seed(0) reset_case_state(pipeline) conditional_dict = pipeline.text_encoder(text_prompts=[prompt]) noise = torch.randn( [1, NUM_CHUNKS * CHUNK_FRAMES, 16, 60, 104], device=device, dtype=torch.bfloat16, ) actual_timesteps = [int(value.item()) for value in pipeline.denoising_step_list] if len(actual_timesteps) != NUM_STEPS: raise ValueError(f"expected {NUM_STEPS} denoising timesteps, got {actual_timesteps}") for chunk_id in range(NUM_CHUNKS): start_frame = chunk_id * CHUNK_FRAMES noisy_input = noise[:, start_frame : start_frame + CHUNK_FRAMES] step_tensors: dict[str, torch.Tensor] = {} denoised_pred: torch.Tensor | None = None timestep: torch.Tensor | None = None for step_id, current_timestep in enumerate(pipeline.denoising_step_list): timestep = torch.full( (1, CHUNK_FRAMES), int(current_timestep.item()), device=device, dtype=torch.int64, ) step_tensors[f"step_{step_id}_noisy_latent"] = clone_bf16_cpu(noisy_input) step_tensors[f"step_{step_id}_timestep"] = timestep.cpu().contiguous() capture.begin_denoise(step_id) try: flow, denoised_pred = pipeline.generator( noisy_image_or_video=noisy_input, conditional_dict=conditional_dict, timestep=timestep, kv_cache=pipeline.kv_cache1, crossattn_cache=pipeline.crossattn_cache, current_start=start_frame * pipeline.frame_seq_length, ) final_hidden = capture.finish_denoise() except BaseException: capture.abort_active_call() raise step_tensors[f"step_{step_id}_final_hidden"] = final_hidden step_tensors[f"step_{step_id}_flow"] = clone_bf16_cpu(flow) if step_id < NUM_STEPS - 1: next_timestep = int(pipeline.denoising_step_list[step_id + 1].item()) noisy_input = pipeline.scheduler.add_noise( denoised_pred.flatten(0, 1), torch.randn_like(denoised_pred.flatten(0, 1)), torch.full( (CHUNK_FRAMES,), next_timestep, device=device, dtype=torch.long, ), ).unflatten(0, denoised_pred.shape[:2]) if denoised_pred is None or timestep is None: raise RuntimeError("denoising loop did not produce a clean prediction") context_timestep = torch.full_like(timestep, int(pipeline.args.context_noise)) capture.begin_clean() try: pipeline.generator( noisy_image_or_video=denoised_pred, conditional_dict=conditional_dict, timestep=context_timestep, kv_cache=pipeline.kv_cache1, crossattn_cache=pipeline.crossattn_cache, current_start=start_frame * pipeline.frame_seq_length, ) clean_features = capture.finish_clean() except BaseException: capture.abort_active_call() raise if not writer.case_path(case_id).is_file(): writer.save_case( case_id, case_cross_kv(pipeline, tuple(int(value) for value in writer.block_ids)), ) if validate_kv_rebuild: metrics = validate_against_live_cache( pipeline.generator.model, pipeline.kv_cache1, clean_features, start_frame=start_frame, num_frames=CHUNK_FRAMES, ) append_jsonl_fsync( writer.log_dir / f"kv_rebuild_worker_{writer.worker_id:02d}.jsonl", {"case_id": case_id, "chunk_id": chunk_id, **metrics}, ) step_path = writer.save_chunk( case_id=case_id, chunk_id=chunk_id, step_tensors=step_tensors, clean_features=clean_features, start_frame=start_frame, metadata={ "prompt": prompt, "prompt_sha256": case.get( "prompt_sha256", hashlib.sha256(prompt.encode("utf-8")).hexdigest(), ), "source_line_number": int(case["source_line_number"]), "seed": 0, "seed_reset_per_case": True, "num_chunks": NUM_CHUNKS, "frames_per_chunk": CHUNK_FRAMES, "num_steps": NUM_STEPS, "denoising_timesteps": actual_timesteps, "context_timestep": int(pipeline.args.context_noise), "teacher_checkpoint": str(checkpoint_path.resolve()), "teacher_checkpoint_key": "generator_ema", "schedule": "F-F-F-F", "supervision_pairs": [[0, 1], [1, 2], [2, 3]], "clean_prefeature_semantics": ( "this chunk's self_attn.k projection input from the final clean pass" ), }, ) print( f"[saved] worker={writer.worker_id} case={case_id} " f"chunk={chunk_id}/{NUM_CHUNKS - 1} path={step_path}", flush=True, ) del step_tensors, clean_features del conditional_dict, noise def main() -> None: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--dataset_root", type=Path, default=DEFAULT_ROOT) parser.add_argument("--config", type=Path, default=DEFAULT_CONFIG) parser.add_argument("--checkpoint", type=Path, default=DEFAULT_CHECKPOINT) parser.add_argument("--worker_id", type=int, required=True) parser.add_argument("--num_workers", type=int, default=8) parser.add_argument("--case_ids", default=None) parser.add_argument("--seed", type=int, default=0) parser.add_argument("--blocks", type=int, nargs="+", default=list(CANDIDATE_BLOCK_IDS)) parser.add_argument("--validate_kv_rebuild", action="store_true") parser.add_argument( "--dry_run", action="store_true", help="validate case assignment and paths without importing or loading Wan", ) args = parser.parse_args() root = args.dataset_root.resolve() cases_path = root / "cases.jsonl" for path in (cases_path, args.config.resolve(), args.checkpoint.resolve()): if not path.is_file(): raise FileNotFoundError(path) if args.seed != 0: raise ValueError("every case must reset and use inference seed 0") if not 0 <= args.worker_id < args.num_workers: raise ValueError("--worker_id must be in [0, num_workers)") selected_ids = parse_ids(args.case_ids) cases = [ case for case in load_cases(cases_path) if int(case["case_id"]) % args.num_workers == args.worker_id and (selected_ids is None or int(case["case_id"]) in selected_ids) ] print( f"[worker {args.worker_id}] assigned {len(cases)} cases " f"using case_id % {args.num_workers}", flush=True, ) if args.dry_run or not cases: return if not torch.cuda.is_available(): raise RuntimeError("CUDA is required to build Predictor v4 teacher data") torch.cuda.set_device(0) device = torch.device("cuda:0") pipeline, _ = load_teacher(args, device) writer = PredictorV4DatasetWriter( root, worker_id=args.worker_id, block_ids=tuple(args.blocks), ) with PredictorV4TeacherCapture( pipeline.generator.model, block_ids=tuple(args.blocks), ) as capture: for position, case in enumerate(cases, start=1): case_id = int(case["case_id"]) if writer.case_path(case_id).is_file() and all( writer.is_chunk_complete(case_id, chunk_id) for chunk_id in range(NUM_CHUNKS) ): print(f"[skip] worker={args.worker_id} case={case_id} complete", flush=True) continue print( f"[run] worker={args.worker_id} case={case_id} " f"({position}/{len(cases)})", flush=True, ) build_case( case=case, pipeline=pipeline, writer=writer, capture=capture, device=device, validate_kv_rebuild=args.validate_kv_rebuild, checkpoint_path=args.checkpoint, ) gc.collect() torch.cuda.empty_cache() summary = { "worker_id": args.worker_id, "num_workers": args.num_workers, "assigned_cases": len(cases), "completed_cases": sum( all(writer.is_chunk_complete(int(case["case_id"]), chunk_id) for chunk_id in range(NUM_CHUNKS)) for case in cases ), } atomic_write_text( writer.log_dir / f"worker_{args.worker_id:02d}_summary.json", json.dumps(summary, indent=2, sort_keys=True) + "\n", ) if __name__ == "__main__": main()