Self-Forcing / scripts /build_predictor_offline_data.py
Cccccz's picture
Upload Python scripts
bc29ee3 verified
Raw
History Blame Contribute Delete
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
@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:
# 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()