#!/usr/bin/env python3 """Add the minimal chunk-0 context required by Stage-1 Predictor training.""" from __future__ import annotations import argparse import json import os import shutil import time from pathlib import Path import build_predictor_offline_data as base torch = base.torch OmegaConf = base.OmegaConf def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--gpu", default=base.PHYSICAL_GPU) parser.add_argument("--dataset_root", type=Path, required=True) parser.add_argument( "--config_path", type=Path, default=Path("configs/causal_forcing_dmd_chunkwise.yaml"), ) parser.add_argument( "--checkpoint_path", type=Path, default=Path("checkpoints/chunkwise/causal_forcing.pt"), ) parser.add_argument("--prompt_ids", type=int, nargs="*", default=None) parser.add_argument("--generation_seed", type=int, default=0) parser.add_argument("--max_new_prompts", type=int, default=None) parser.add_argument("--min_free_gib", type=float, default=500.0) return parser.parse_args() @torch.inference_mode() def generate_chunk0( pipeline, recorder: base.TrajectoryRecorder, prompt: str, generation_seed: int, device: torch.device, ): base.set_seed(generation_seed) base.reset_caches(pipeline, 1, torch.bfloat16, device) recorder.clean_prefeatures = {layer: [] for layer in recorder.layers} conditional = pipeline.text_encoder(text_prompts=[prompt]) # Draw the full 21-latent noise tensor so the RNG state and chunk-0 slice # exactly match the original offline rollout. full_noise = torch.randn( 1, 21, base.LATENT_CHANNELS, base.LATENT_HEIGHT, base.LATENT_WIDTH, dtype=torch.bfloat16, device=device, ) noisy_input = full_noise[:, : pipeline.num_frame_per_block] timesteps = pipeline.denoising_step_list.to(device=device) trajectory = {} torch.cuda.reset_peak_memory_stats() torch.cuda.synchronize() started = time.perf_counter() timestep = None denoised_pred = None for step, current_timestep in enumerate(timesteps): timestep = torch.ones( [1, pipeline.num_frame_per_block], device=device, dtype=torch.int64 ) * current_timestep recorder.start_denoising_step() _, denoised_pred = pipeline.generator( noisy_image_or_video=noisy_input, conditional_dict=conditional, timestep=timestep, kv_cache=pipeline.kv_cache1, crossattn_cache=pipeline.crossattn_cache, current_start=0, ) trajectory[f"chunk_00_step_{step:02d}_final_hidden"] = ( recorder.finish_denoising_step() ) if step < len(timesteps) - 1: next_timestep = timesteps[step + 1] flat = denoised_pred.flatten(0, 1) noisy_input = pipeline.scheduler.add_noise( flat, torch.randn_like(flat), next_timestep * torch.ones( [pipeline.num_frame_per_block], device=device, dtype=torch.long, ), ).unflatten(0, denoised_pred.shape[:2]) if denoised_pred is None or timestep is None: raise RuntimeError("Chunk-0 denoising produced no output") recorder.start_clean_pass() pipeline.generator( noisy_image_or_video=denoised_pred, conditional_dict=conditional, timestep=torch.ones_like(timestep) * pipeline.args.context_noise, kv_cache=pipeline.kv_cache1, crossattn_cache=pipeline.crossattn_cache, current_start=0, ) recorder.finish_clean_pass() torch.cuda.synchronize() return ( trajectory, recorder.clean_prefeatures, time.perf_counter() - started, torch.cuda.max_memory_allocated() / 1024**3, ) def save_context(prompt_dir, trajectory, prefeatures, elapsed_s, peak_gib): destination = prompt_dir / "chunk0_context" partial = prompt_dir / "chunk0_context.partial" if partial.exists(): shutil.rmtree(partial) partial.mkdir(parents=True) metadata = {"dataset_version": "2", "kind": "chunk0_context"} base.atomic_save_safetensors( trajectory, partial / "trajectory.safetensors", metadata ) for layer in range(30): values = prefeatures[layer] if len(values) != 1: raise RuntimeError(f"Layer {layer} captured {len(values)} values") base.atomic_save_safetensors( {"chunk_00": values[0]}, partial / "clean_prefeatures" / f"block_{layer:02d}.safetensors", {**metadata, "block_id": str(layer)}, ) base.atomic_write_json( partial / "metadata.json", { "kind": "context_only", "chunk": 0, "is_training_target": False, "hidden_steps": [0, 1, 2, 3], "layers": list(range(30)), "elapsed_s": elapsed_s, "peak_gpu_gib": peak_gib, }, ) (partial / "_SUCCESS").write_text("ok\n", encoding="utf-8") if destination.exists(): shutil.rmtree(destination) os.replace(partial, destination) return sum(p.stat().st_size for p in destination.rglob("*") if p.is_file()) def main() -> None: args = parse_args() root = args.dataset_root.resolve() config_path = base.resolve_path(args.config_path) checkpoint_path = base.resolve_path(args.checkpoint_path) selected = json.loads((root / "prompt_selection.json").read_text())["prompts"] requested = ( set(range(len(selected))) if args.prompt_ids is None else set(args.prompt_ids) ) pending = [ (index, item) for index, item in enumerate(selected) if index in requested and not (root / f"prompt_{index:04d}" / "chunk0_context" / "_SUCCESS").exists() ] if not pending: print("[context] all requested prompt contexts exist", flush=True) return config = OmegaConf.merge( OmegaConf.load(base.REPO_ROOT / "configs/default_config.yaml"), OmegaConf.load(config_path), ) device = torch.device("cuda") pipeline = base.build_pipeline(config, checkpoint_path, device) recorder = base.TrajectoryRecorder(pipeline.generator.model, list(range(30))) 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(root).free / 1024**3 if free_gib < args.min_free_gib: raise RuntimeError(f"Only {free_gib:.1f} GiB free") trajectory, prefeatures, elapsed_s, peak_gib = generate_chunk0( pipeline, recorder, selection["prompt"], args.generation_seed, device ) size = save_context( root / f"prompt_{index:04d}", trajectory, prefeatures, elapsed_s, peak_gib, ) generated += 1 print( f"[context] saved prompt_{index:04d}: {size / 1024**3:.3f} GiB, " f"{elapsed_s:.1f}s, peak={peak_gib:.1f} GiB", flush=True, ) del trajectory, prefeatures torch.cuda.empty_cache() finally: recorder.close() print(f"[context] generated {generated} prompt contexts", flush=True) if __name__ == "__main__": main()