#!/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()