| |
| """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 ( |
| CANDIDATE_BLOCK_IDS, |
| CHUNK_FRAMES, |
| NUM_CHUNKS, |
| NUM_STEPS, |
| PredictorV4DatasetWriter, |
| PredictorV4TeacherCapture, |
| ) |
| from predictor_data.capture import clone_bf16_cpu |
| from predictor_data.kv_validation import validate_against_live_cache |
| from predictor_data.writer import atomic_write_text |
|
|
|
|
| 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() |
|
|