Download scripts/build_predictor_offline_data.py from Cccccz/Self-Forcing: direct link, hf CLI and curl.
- Browser
- Download file 29.9 kB
-
https://huggingface.co/Cccccz/Self-Forcing/resolve/main/scripts/build_predictor_offline_data.py
- Command line
-
hf download hf://Cccccz/Self-Forcing/scripts/build_predictor_offline_data.py
-
curl -L -o build_predictor_offline_data.py https://huggingface.co/Cccccz/Self-Forcing/resolve/main/scripts/build_predictor_offline_data.py
29.9 kB
| #!/usr/bin/env python3 | |
| """Build offline FFFF trajectories for the lightweight Self-Forcing predictor. | |
| The dataset is prompt-sharded and resumable. Common denoising tensors are | |
| stored once per prompt, while clean self-attention prefeatures are stored in | |
| one sidecar per Teacher block. This layout avoids duplicating history tensors | |
| across the 18 adjacent-step training samples produced by each prompt. | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import hashlib | |
| import json | |
| import os | |
| import random | |
| import shutil | |
| import sys | |
| import time | |
| from pathlib import Path | |
| from typing import Any | |
| def _preparse_gpu() -> str: | |
| parser = argparse.ArgumentParser(add_help=False) | |
| parser.add_argument("--gpu", default="2") | |
| args, _ = parser.parse_known_args() | |
| os.environ["CUDA_VISIBLE_DEVICES"] = str(args.gpu) | |
| return str(args.gpu) | |
| PHYSICAL_GPU = _preparse_gpu() | |
| import torch | |
| from omegaconf import OmegaConf | |
| from safetensors.torch import save_file | |
| REPO_ROOT = Path(__file__).resolve().parents[1] | |
| if str(REPO_ROOT) not in sys.path: | |
| sys.path.insert(0, str(REPO_ROOT)) | |
| from pipeline import CausalInferencePipeline | |
| from utils.misc import set_seed | |
| from utils.wan_wrapper import WanDiffusionWrapper, WanTextEncoder | |
| DATASET_VERSION = 2 | |
| EXCLUDED_CHUNKS = (0,) | |
| LATENT_CHANNELS = 16 | |
| LATENT_HEIGHT = 60 | |
| LATENT_WIDTH = 104 | |
| def parse_args() -> argparse.Namespace: | |
| parser = argparse.ArgumentParser( | |
| description="Build full-step Self-Forcing predictor trajectories" | |
| ) | |
| parser.add_argument("--gpu", default=PHYSICAL_GPU) | |
| parser.add_argument( | |
| "--config_path", | |
| type=Path, | |
| default=Path("configs/self_forcing_sid.yaml"), | |
| ) | |
| parser.add_argument( | |
| "--checkpoint_path", | |
| type=Path, | |
| default=Path("checkpoints/self_forcing_dmd.pt"), | |
| ) | |
| parser.add_argument( | |
| "--prompt_path", | |
| type=Path, | |
| default=Path("prompts/vidprom_filtered_extended.txt"), | |
| ) | |
| parser.add_argument( | |
| "--validation_prompt_path", | |
| type=Path, | |
| default=Path("prompts/MovieGenVideoBench_extended.txt"), | |
| ) | |
| parser.add_argument("--output_dir", type=Path, required=True) | |
| parser.add_argument("--num_prompts", type=int, default=100) | |
| parser.add_argument( | |
| "--prompt_ids", | |
| type=int, | |
| nargs="*", | |
| default=None, | |
| help=( | |
| "Only materialize these selected prompt IDs. The manifest still " | |
| "records the full deterministic prompt selection." | |
| ), | |
| ) | |
| parser.add_argument("--num_frames", type=int, default=21) | |
| parser.add_argument("--selection_seed", type=int, default=0) | |
| parser.add_argument("--generation_seed", type=int, default=0) | |
| parser.add_argument( | |
| "--layers", | |
| type=int, | |
| nargs="*", | |
| default=None, | |
| help="Teacher blocks to cache. Omit to cache every block.", | |
| ) | |
| parser.add_argument( | |
| "--max_new_prompts", | |
| type=int, | |
| default=None, | |
| help="Stop after this many new prompt shards; use 1 for the dry run.", | |
| ) | |
| parser.add_argument( | |
| "--min_free_gib", | |
| type=float, | |
| default=50.0, | |
| help="Stop before a new prompt if free disk space falls below this value.", | |
| ) | |
| parser.add_argument("--overwrite", action="store_true") | |
| args = parser.parse_args() | |
| if args.num_prompts < 1: | |
| parser.error("--num_prompts must be positive") | |
| if args.prompt_ids is not None and any( | |
| value < 0 or value >= args.num_prompts for value in args.prompt_ids | |
| ): | |
| parser.error("--prompt_ids must be within [0, --num_prompts)") | |
| if args.num_frames < 1 or args.num_frames % 3: | |
| parser.error("--num_frames must be a positive multiple of 3") | |
| if args.max_new_prompts is not None and args.max_new_prompts < 0: | |
| parser.error("--max_new_prompts must be non-negative") | |
| return args | |
| def resolve_path(path: Path) -> Path: | |
| path = path.expanduser() | |
| return path.resolve() if path.is_absolute() else (REPO_ROOT / path).resolve() | |
| def atomic_write_json(path: Path, value: Any) -> None: | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| temporary = path.with_suffix(path.suffix + ".tmp") | |
| temporary.write_text( | |
| json.dumps(value, indent=2, ensure_ascii=False) + "\n", | |
| encoding="utf-8", | |
| ) | |
| os.replace(temporary, path) | |
| def file_sha256(path: Path) -> str: | |
| digest = hashlib.sha256() | |
| with path.open("rb") as handle: | |
| while chunk := handle.read(8 * 1024 * 1024): | |
| digest.update(chunk) | |
| return digest.hexdigest() | |
| def read_nonempty_lines(path: Path) -> list[str]: | |
| with path.open("r", encoding="utf-8") as handle: | |
| return [line.strip() for line in handle if line.strip()] | |
| def select_prompts( | |
| prompt_path: Path, | |
| validation_prompt_path: Path, | |
| num_prompts: int, | |
| seed: int, | |
| ) -> list[dict[str, Any]]: | |
| source = read_nonempty_lines(prompt_path) | |
| validation = set(read_nonempty_lines(validation_prompt_path)[:100]) | |
| eligible = [ | |
| {"source_index": index, "prompt": prompt} | |
| for index, prompt in enumerate(source) | |
| if prompt not in validation | |
| ] | |
| if len(eligible) < num_prompts: | |
| raise ValueError( | |
| f"Only {len(eligible)} eligible prompts remain after excluding " | |
| f"the first 100 validation prompts; requested {num_prompts}" | |
| ) | |
| return random.Random(seed).sample(eligible, num_prompts) | |
| def tensor_to_bf16_cpu(value: torch.Tensor) -> torch.Tensor: | |
| return value.detach().to(device="cpu", dtype=torch.bfloat16).contiguous() | |
| def atomic_save_safetensors( | |
| tensors: dict[str, torch.Tensor], | |
| path: Path, | |
| metadata: dict[str, str], | |
| ) -> None: | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| temporary = path.with_suffix(path.suffix + ".tmp") | |
| save_file(tensors, temporary, metadata=metadata) | |
| os.replace(temporary, path) | |
| def directory_size(path: Path) -> int: | |
| return sum(item.stat().st_size for item in path.rglob("*") if item.is_file()) | |
| class TrajectoryRecorder: | |
| """Capture final hidden states and clean K-projection inputs.""" | |
| def __init__(self, model: torch.nn.Module, layers: list[int]) -> None: | |
| self.model = model | |
| self.layers = layers | |
| self.mode: str | None = None | |
| self.final_hidden: torch.Tensor | None = None | |
| self.current_clean: dict[int, torch.Tensor] = {} | |
| self.clean_prefeatures: dict[int, list[torch.Tensor]] = { | |
| layer: [] for layer in layers | |
| } | |
| self.handles: list[Any] = [] | |
| self.handles.append( | |
| model.head.register_forward_pre_hook(self._head_pre_hook) | |
| ) | |
| for layer in layers: | |
| self.handles.append( | |
| model.blocks[layer].self_attn.k.register_forward_pre_hook( | |
| self._make_clean_prefeature_hook(layer) | |
| ) | |
| ) | |
| def close(self) -> None: | |
| for handle in self.handles: | |
| handle.remove() | |
| self.handles.clear() | |
| def start_denoising_step(self) -> None: | |
| self.mode = "denoise" | |
| self.final_hidden = None | |
| def finish_denoising_step(self) -> torch.Tensor: | |
| if self.final_hidden is None: | |
| raise RuntimeError("The Teacher head hook did not capture final_hidden") | |
| value = self.final_hidden | |
| self.final_hidden = None | |
| self.mode = None | |
| return value | |
| def start_clean_pass(self) -> None: | |
| self.mode = "clean" | |
| self.current_clean = {} | |
| def finish_clean_pass(self, *, store: bool = True) -> dict[int, torch.Tensor]: | |
| missing = sorted(set(self.layers) - set(self.current_clean)) | |
| if missing: | |
| raise RuntimeError( | |
| f"Clean pass did not capture prefeatures for blocks {missing}" | |
| ) | |
| captured = self.current_clean | |
| if store: | |
| for layer in self.layers: | |
| self.clean_prefeatures[layer].append(captured[layer]) | |
| self.current_clean = {} | |
| self.mode = None | |
| return captured | |
| def _head_pre_hook( | |
| self, _module: torch.nn.Module, inputs: tuple[torch.Tensor, ...] | |
| ) -> None: | |
| if self.mode != "denoise": | |
| return | |
| if self.final_hidden is not None: | |
| raise RuntimeError("Captured final_hidden more than once in one step") | |
| if not inputs or not isinstance(inputs[0], torch.Tensor): | |
| raise RuntimeError("Unexpected Teacher head inputs") | |
| self.final_hidden = tensor_to_bf16_cpu(inputs[0]) | |
| def _make_clean_prefeature_hook(self, layer: int): | |
| def hook( | |
| _module: torch.nn.Module, inputs: tuple[torch.Tensor, ...] | |
| ) -> None: | |
| if self.mode != "clean": | |
| return | |
| if layer in self.current_clean: | |
| raise RuntimeError( | |
| f"Captured block {layer} clean prefeature more than once" | |
| ) | |
| if not inputs or not isinstance(inputs[0], torch.Tensor): | |
| raise RuntimeError(f"Unexpected block {layer} K inputs") | |
| self.current_clean[layer] = tensor_to_bf16_cpu(inputs[0]) | |
| return hook | |
| def build_pipeline( | |
| config: Any, checkpoint_path: Path, device: torch.device | |
| ) -> CausalInferencePipeline: | |
| generator = WanDiffusionWrapper( | |
| **getattr(config, "model_kwargs", {}), is_causal=True | |
| ) | |
| text_encoder = WanTextEncoder() | |
| pipeline = CausalInferencePipeline( | |
| config, | |
| device=device, | |
| generator=generator, | |
| text_encoder=text_encoder, | |
| vae=torch.nn.Identity(), | |
| ) | |
| checkpoint = torch.load( | |
| checkpoint_path, map_location="cpu", weights_only=False, mmap=True | |
| ) | |
| if set(checkpoint) != {"generator_ema"}: | |
| raise KeyError( | |
| f"Expected checkpoint key generator_ema, found {sorted(checkpoint)}" | |
| ) | |
| pipeline.generator.load_state_dict(checkpoint["generator_ema"], strict=True) | |
| del checkpoint | |
| pipeline.to(dtype=torch.bfloat16) | |
| pipeline.text_encoder.to(device=device) | |
| pipeline.generator.to(device=device) | |
| pipeline.eval() | |
| pipeline.generator.model.requires_grad_(False) | |
| pipeline.text_encoder.requires_grad_(False) | |
| return pipeline | |
| def reset_caches( | |
| pipeline: CausalInferencePipeline, | |
| batch_size: int, | |
| dtype: torch.dtype, | |
| device: torch.device, | |
| ) -> None: | |
| if pipeline.kv_cache1 is None: | |
| pipeline._initialize_kv_cache(batch_size, dtype, device) | |
| pipeline._initialize_crossattn_cache(batch_size, dtype, device) | |
| return | |
| 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 collect_cross_attention_cache( | |
| pipeline: CausalInferencePipeline, | |
| layers: list[int], | |
| ) -> dict[str, torch.Tensor]: | |
| output: dict[str, torch.Tensor] = {} | |
| for layer in layers: | |
| cache = pipeline.crossattn_cache[layer] | |
| if not cache["is_init"]: | |
| raise RuntimeError(f"Cross-attention cache for block {layer} is empty") | |
| output[f"block_{layer:02d}_k"] = tensor_to_bf16_cpu(cache["k"]) | |
| output[f"block_{layer:02d}_v"] = tensor_to_bf16_cpu(cache["v"]) | |
| return output | |
| def generate_prompt( | |
| pipeline: CausalInferencePipeline, | |
| recorder: TrajectoryRecorder, | |
| prompt: str, | |
| num_frames: int, | |
| generation_seed: int, | |
| device: torch.device, | |
| ) -> tuple[ | |
| dict[str, torch.Tensor], | |
| dict[int, list[torch.Tensor]], | |
| dict[str, torch.Tensor], | |
| dict[int, torch.Tensor], | |
| dict[str, torch.Tensor], | |
| float, | |
| float, | |
| ]: | |
| set_seed(generation_seed) | |
| reset_caches(pipeline, 1, torch.bfloat16, device) | |
| recorder.clean_prefeatures = {layer: [] for layer in recorder.layers} | |
| conditional_dict = pipeline.text_encoder(text_prompts=[prompt]) | |
| noise = torch.randn( | |
| 1, | |
| num_frames, | |
| LATENT_CHANNELS, | |
| LATENT_HEIGHT, | |
| LATENT_WIDTH, | |
| dtype=torch.bfloat16, | |
| device=device, | |
| ) | |
| trajectory: dict[str, torch.Tensor] = {} | |
| chunk0_trajectory: dict[str, torch.Tensor] = {} | |
| chunk0_prefeatures: dict[int, torch.Tensor] = {} | |
| chunk_size = pipeline.num_frame_per_block | |
| num_chunks = num_frames // chunk_size | |
| timesteps = pipeline.denoising_step_list.to(device=device) | |
| torch.cuda.reset_peak_memory_stats() | |
| torch.cuda.synchronize() | |
| start_time = time.perf_counter() | |
| current_start_frame = 0 | |
| for chunk in range(num_chunks): | |
| noisy_input = noise[ | |
| :, current_start_frame : current_start_frame + chunk_size | |
| ] | |
| timestep: torch.Tensor | None = None | |
| denoised_pred: torch.Tensor | None = None | |
| for step, current_timestep in enumerate(timesteps): | |
| timestep = ( | |
| torch.ones( | |
| [1, chunk_size], | |
| device=device, | |
| dtype=torch.int64, | |
| ) | |
| * current_timestep | |
| ) | |
| prefix = f"chunk_{chunk:02d}_step_{step:02d}" | |
| if chunk not in EXCLUDED_CHUNKS: | |
| trajectory[f"{prefix}_noisy_latent"] = tensor_to_bf16_cpu( | |
| noisy_input | |
| ) | |
| trajectory[f"{prefix}_timestep"] = ( | |
| timestep.detach() | |
| .to(device="cpu", dtype=torch.float32) | |
| .contiguous() | |
| ) | |
| recorder.start_denoising_step() | |
| 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=current_start_frame * pipeline.frame_seq_length, | |
| ) | |
| final_hidden = recorder.finish_denoising_step() | |
| if chunk not in EXCLUDED_CHUNKS: | |
| trajectory[f"{prefix}_final_hidden"] = final_hidden | |
| trajectory[f"{prefix}_flow"] = tensor_to_bf16_cpu(flow) | |
| else: | |
| chunk0_trajectory[f"{prefix}_final_hidden"] = final_hidden | |
| if step < len(timesteps) - 1: | |
| next_timestep = timesteps[step + 1] | |
| denoised_flat = denoised_pred.flatten(0, 1) | |
| noisy_input = pipeline.scheduler.add_noise( | |
| denoised_flat, | |
| torch.randn_like(denoised_flat), | |
| next_timestep | |
| * torch.ones( | |
| [chunk_size], device=device, dtype=torch.long | |
| ), | |
| ).unflatten(0, denoised_pred.shape[:2]) | |
| if denoised_pred is None or timestep is None: | |
| raise RuntimeError("Denoising loop produced no output") | |
| if chunk not in EXCLUDED_CHUNKS: | |
| trajectory[f"chunk_{chunk:02d}_clean_latent"] = ( | |
| tensor_to_bf16_cpu(denoised_pred) | |
| ) | |
| recorder.start_clean_pass() | |
| context_timestep = torch.ones_like(timestep) * pipeline.args.context_noise | |
| 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=current_start_frame * pipeline.frame_seq_length, | |
| ) | |
| captured_clean = recorder.finish_clean_pass( | |
| store=chunk not in EXCLUDED_CHUNKS | |
| ) | |
| if chunk in EXCLUDED_CHUNKS: | |
| if chunk != 0: | |
| raise RuntimeError(f"Unsupported excluded context chunk {chunk}") | |
| chunk0_prefeatures = captured_clean | |
| current_start_frame += chunk_size | |
| cross_attention = collect_cross_attention_cache(pipeline, recorder.layers) | |
| torch.cuda.synchronize() | |
| elapsed = time.perf_counter() - start_time | |
| peak_gib = torch.cuda.max_memory_allocated() / (1024**3) | |
| del conditional_dict, noise | |
| return ( | |
| trajectory, | |
| recorder.clean_prefeatures, | |
| chunk0_trajectory, | |
| chunk0_prefeatures, | |
| cross_attention, | |
| elapsed, | |
| peak_gib, | |
| ) | |
| def save_prompt_shard( | |
| output_dir: Path, | |
| prompt_index: int, | |
| selection: dict[str, Any], | |
| trajectory: dict[str, torch.Tensor], | |
| clean_prefeatures: dict[int, list[torch.Tensor]], | |
| chunk0_trajectory: dict[str, torch.Tensor], | |
| chunk0_prefeatures: dict[int, torch.Tensor], | |
| cross_attention: dict[str, torch.Tensor], | |
| elapsed_s: float, | |
| peak_gpu_gib: float, | |
| layers: list[int], | |
| generation_seed: int, | |
| num_chunks: int, | |
| ) -> Path: | |
| destination = output_dir / f"prompt_{prompt_index:04d}" | |
| partial = output_dir / f"prompt_{prompt_index:04d}.partial" | |
| if partial.exists(): | |
| shutil.rmtree(partial) | |
| partial.mkdir(parents=True) | |
| shared_metadata = { | |
| "dataset_version": str(DATASET_VERSION), | |
| "dtype": "bfloat16", | |
| "prompt_index": str(prompt_index), | |
| } | |
| atomic_save_safetensors( | |
| trajectory, | |
| partial / "trajectory.safetensors", | |
| {**shared_metadata, "kind": "trajectory"}, | |
| ) | |
| atomic_save_safetensors( | |
| cross_attention, | |
| partial / "cross_attention.safetensors", | |
| {**shared_metadata, "kind": "cross_attention_kv"}, | |
| ) | |
| atomic_save_safetensors( | |
| chunk0_trajectory, | |
| partial / "chunk0_context" / "trajectory.safetensors", | |
| {**shared_metadata, "kind": "chunk0_context_final_hidden"}, | |
| ) | |
| for layer in layers: | |
| atomic_save_safetensors( | |
| {"chunk_00": chunk0_prefeatures[layer]}, | |
| partial | |
| / "chunk0_context" | |
| / "clean_prefeatures" | |
| / f"block_{layer:02d}.safetensors", | |
| { | |
| **shared_metadata, | |
| "kind": "chunk0_context_clean_self_attention_k_input", | |
| "block_id": str(layer), | |
| }, | |
| ) | |
| atomic_write_json( | |
| partial / "chunk0_context" / "metadata.json", | |
| { | |
| "kind": "context_only", | |
| "chunk": 0, | |
| "is_training_target": False, | |
| "hidden_steps": [0, 1, 2, 3], | |
| "layers": layers, | |
| }, | |
| ) | |
| (partial / "chunk0_context" / "_SUCCESS").write_text( | |
| "ok\n", encoding="utf-8" | |
| ) | |
| prefeature_shapes: dict[str, list[int]] = {} | |
| for layer in layers: | |
| values = clean_prefeatures[layer] | |
| stored_chunks = [ | |
| chunk for chunk in range(num_chunks) if chunk not in EXCLUDED_CHUNKS | |
| ] | |
| if len(values) != len(stored_chunks): | |
| raise RuntimeError( | |
| f"Expected {len(stored_chunks)} stored chunks, got {len(values)}" | |
| ) | |
| tensors = { | |
| f"chunk_{chunk:02d}": value | |
| for chunk, value in zip(stored_chunks, values) | |
| } | |
| atomic_save_safetensors( | |
| tensors, | |
| partial / "clean_prefeatures" / f"block_{layer:02d}.safetensors", | |
| { | |
| **shared_metadata, | |
| "kind": "clean_self_attention_k_input", | |
| "block_id": str(layer), | |
| }, | |
| ) | |
| if values: | |
| prefeature_shapes[str(layer)] = list(values[0].shape) | |
| metadata = { | |
| "dataset_version": DATASET_VERSION, | |
| "prompt_index": prompt_index, | |
| "source_index": selection["source_index"], | |
| "prompt": selection["prompt"], | |
| "generation_seed": generation_seed, | |
| "dtype": "bfloat16", | |
| "layers": layers, | |
| "num_clean_chunks": len(clean_prefeatures[layers[0]]), | |
| "excluded_chunks": list(EXCLUDED_CHUNKS), | |
| "stored_chunks": [ | |
| chunk for chunk in range(num_chunks) if chunk not in EXCLUDED_CHUNKS | |
| ], | |
| "prefeature_shapes": prefeature_shapes, | |
| "elapsed_s": elapsed_s, | |
| "peak_gpu_gib": peak_gpu_gib, | |
| } | |
| atomic_write_json(partial / "metadata.json", metadata) | |
| (partial / "_SUCCESS").write_text("ok\n", encoding="utf-8") | |
| os.replace(partial, destination) | |
| return destination | |
| def prepare_manifest( | |
| args: argparse.Namespace, | |
| config: Any, | |
| prompt_path: Path, | |
| validation_prompt_path: Path, | |
| checkpoint_path: Path, | |
| output_dir: Path, | |
| ) -> tuple[dict[str, Any], list[dict[str, Any]]]: | |
| output_dir.mkdir(parents=True, exist_ok=True) | |
| prompt_selection_path = output_dir / "prompt_selection.json" | |
| selected = select_prompts( | |
| prompt_path, | |
| validation_prompt_path, | |
| args.num_prompts, | |
| args.selection_seed, | |
| ) | |
| selection_document = { | |
| "selection_seed": args.selection_seed, | |
| "num_prompts": args.num_prompts, | |
| "prompt_source": str(prompt_path), | |
| "prompt_source_sha256": file_sha256(prompt_path), | |
| "validation_source": str(validation_prompt_path), | |
| "validation_source_sha256": file_sha256(validation_prompt_path), | |
| "excluded_validation_count": 100, | |
| "prompts": selected, | |
| } | |
| if prompt_selection_path.exists() and not args.overwrite: | |
| existing = json.loads(prompt_selection_path.read_text(encoding="utf-8")) | |
| if existing != selection_document: | |
| raise RuntimeError( | |
| "Existing prompt_selection.json differs from the requested " | |
| "selection. Use another output directory or --overwrite." | |
| ) | |
| else: | |
| atomic_write_json(prompt_selection_path, selection_document) | |
| manifest = { | |
| "dataset_version": DATASET_VERSION, | |
| "config_path": str(resolve_path(args.config_path)), | |
| "checkpoint_path": str(checkpoint_path), | |
| "checkpoint_key": "generator_ema", | |
| "checkpoint_sha256": file_sha256(checkpoint_path), | |
| "model": "Wan2.1-T2V-1.3B causal generator_ema", | |
| "model_hidden_dim": 1536, | |
| "num_teacher_blocks": 30, | |
| "cached_layers": args.layers, | |
| "storage_dtype": "bfloat16", | |
| "num_prompts": args.num_prompts, | |
| "num_frames": args.num_frames, | |
| "num_chunks": args.num_frames // int(config.num_frame_per_block), | |
| "excluded_chunks": list(EXCLUDED_CHUNKS), | |
| "stored_chunks": [ | |
| chunk | |
| for chunk in range( | |
| args.num_frames // int(config.num_frame_per_block) | |
| ) | |
| if chunk not in EXCLUDED_CHUNKS | |
| ], | |
| "num_frame_per_block": int(config.num_frame_per_block), | |
| "local_attention_latents": ( | |
| int(config.model_kwargs.local_attn_size) | |
| if args.num_frames > 21 else None | |
| ), | |
| "denoising_step_source": list(config.denoising_step_list), | |
| "selection_seed": args.selection_seed, | |
| "generation_seed_reset_per_prompt": args.generation_seed, | |
| "prompt_selection_file": str(prompt_selection_path), | |
| "schema": { | |
| "trajectory": "prompt_NNNN/trajectory.safetensors", | |
| "cross_attention": "prompt_NNNN/cross_attention.safetensors", | |
| "clean_prefeature": ( | |
| "prompt_NNNN/clean_prefeatures/block_XX.safetensors" | |
| ), | |
| "chunk0_context": "prompt_NNNN/chunk0_context/", | |
| }, | |
| } | |
| atomic_write_json(output_dir / "manifest.json", manifest) | |
| return manifest, selected | |
| def update_progress(output_dir: Path, num_prompts: int) -> None: | |
| completed = [] | |
| total_bytes = 0 | |
| for index in range(num_prompts): | |
| prompt_dir = output_dir / f"prompt_{index:04d}" | |
| if (prompt_dir / "_SUCCESS").exists(): | |
| completed.append(index) | |
| total_bytes += directory_size(prompt_dir) | |
| atomic_write_json( | |
| output_dir / "progress.json", | |
| { | |
| "completed_prompts": completed, | |
| "completed_count": len(completed), | |
| "num_prompts": num_prompts, | |
| "stored_bytes": total_bytes, | |
| "stored_gib": total_bytes / (1024**3), | |
| }, | |
| ) | |
| def main() -> None: | |
| args = parse_args() | |
| args.config_path = resolve_path(args.config_path) | |
| args.checkpoint_path = resolve_path(args.checkpoint_path) | |
| args.prompt_path = resolve_path(args.prompt_path) | |
| args.validation_prompt_path = resolve_path(args.validation_prompt_path) | |
| args.output_dir = resolve_path(args.output_dir) | |
| config = OmegaConf.merge( | |
| OmegaConf.load(REPO_ROOT / "configs/default_config.yaml"), | |
| OmegaConf.load(args.config_path), | |
| ) | |
| if int(config.num_frame_per_block) != 3: | |
| raise ValueError("This dataset builder currently expects 3-frame chunks") | |
| if args.num_frames > 21: | |
| # Keep the released model's 21-latent training horizon as a rolling | |
| # attention window while global RoPE positions continue increasing. | |
| config.model_kwargs.local_attn_size = 21 | |
| checkpoint = torch.load( | |
| args.checkpoint_path, map_location="cpu", weights_only=False, mmap=True | |
| ) | |
| state_dict = checkpoint.get("generator_ema") | |
| if state_dict is None: | |
| raise KeyError("Checkpoint does not contain generator_ema") | |
| checkpoint_layers = sorted( | |
| { | |
| int(key.split(".")[2]) | |
| for key in state_dict | |
| if key.startswith("model.blocks.") | |
| } | |
| ) | |
| del checkpoint, state_dict | |
| if checkpoint_layers != list(range(30)): | |
| raise ValueError( | |
| f"Expected checkpoint blocks 0..29, found {checkpoint_layers}" | |
| ) | |
| layers = ( | |
| list(range(30)) | |
| if args.layers is None or len(args.layers) == 0 | |
| else sorted(set(args.layers)) | |
| ) | |
| invalid = [layer for layer in layers if layer not in checkpoint_layers] | |
| if invalid: | |
| raise ValueError(f"Invalid requested block IDs: {invalid}") | |
| args.layers = layers | |
| _, selected = prepare_manifest( | |
| args, | |
| config, | |
| args.prompt_path, | |
| args.validation_prompt_path, | |
| args.checkpoint_path, | |
| args.output_dir, | |
| ) | |
| update_progress(args.output_dir, args.num_prompts) | |
| if args.max_new_prompts == 0: | |
| print("[prepare] prompt selection and manifest are ready", flush=True) | |
| return | |
| requested_prompt_ids = ( | |
| set(range(args.num_prompts)) | |
| if args.prompt_ids is None | |
| else set(args.prompt_ids) | |
| ) | |
| pending = [] | |
| for index, selection in enumerate(selected): | |
| if index not in requested_prompt_ids: | |
| continue | |
| destination = args.output_dir / f"prompt_{index:04d}" | |
| if (destination / "_SUCCESS").exists() and not args.overwrite: | |
| continue | |
| pending.append((index, selection)) | |
| if not pending: | |
| print("[dataset] all prompt shards already exist", flush=True) | |
| return | |
| device = torch.device("cuda") | |
| pipeline = build_pipeline(config, args.checkpoint_path, device) | |
| if len(pipeline.generator.model.blocks) != 30: | |
| raise ValueError( | |
| f"Loaded generator has {len(pipeline.generator.model.blocks)} blocks" | |
| ) | |
| recorder = TrajectoryRecorder(pipeline.generator.model, layers) | |
| generated = 0 | |
| try: | |
| for index, selection in pending: | |
| if ( | |
| args.max_new_prompts is not None | |
| and generated >= args.max_new_prompts | |
| ): | |
| break | |
| free_gib = shutil.disk_usage(args.output_dir).free / (1024**3) | |
| if free_gib < args.min_free_gib: | |
| raise RuntimeError( | |
| f"Only {free_gib:.1f} GiB free, below --min_free_gib " | |
| f"{args.min_free_gib:.1f}" | |
| ) | |
| destination = args.output_dir / f"prompt_{index:04d}" | |
| if destination.exists(): | |
| if not args.overwrite: | |
| raise RuntimeError( | |
| f"Incomplete destination exists: {destination}" | |
| ) | |
| shutil.rmtree(destination) | |
| print( | |
| f"[dataset] prompt {index + 1}/{args.num_prompts}, " | |
| f"free={free_gib:.1f} GiB", | |
| flush=True, | |
| ) | |
| ( | |
| trajectory, | |
| clean_prefeatures, | |
| chunk0_trajectory, | |
| chunk0_prefeatures, | |
| cross_attention, | |
| elapsed_s, | |
| peak_gpu_gib, | |
| ) = generate_prompt( | |
| pipeline, | |
| recorder, | |
| selection["prompt"], | |
| args.num_frames, | |
| args.generation_seed, | |
| device, | |
| ) | |
| destination = save_prompt_shard( | |
| args.output_dir, | |
| index, | |
| selection, | |
| trajectory, | |
| clean_prefeatures, | |
| chunk0_trajectory, | |
| chunk0_prefeatures, | |
| cross_attention, | |
| elapsed_s, | |
| peak_gpu_gib, | |
| layers, | |
| args.generation_seed, | |
| args.num_frames // int(config.num_frame_per_block), | |
| ) | |
| shard_gib = directory_size(destination) / (1024**3) | |
| print( | |
| f"[dataset] saved {destination.name}: {shard_gib:.3f} GiB, " | |
| f"{elapsed_s:.1f}s, peak={peak_gpu_gib:.1f} GiB", | |
| flush=True, | |
| ) | |
| generated += 1 | |
| update_progress(args.output_dir, args.num_prompts) | |
| del ( | |
| trajectory, | |
| clean_prefeatures, | |
| chunk0_trajectory, | |
| chunk0_prefeatures, | |
| cross_attention, | |
| ) | |
| torch.cuda.empty_cache() | |
| finally: | |
| recorder.close() | |
| print(f"[dataset] generated {generated} new prompt shards", flush=True) | |
| if __name__ == "__main__": | |
| main() | |