| |
| """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 |
|
|
|
|
| @torch.inference_mode() |
| 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: |
| |
| |
| 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() |
|
|