Cccccz/HY
0
1#!/usr/bin/env python32"""Single-GPU training for the dense image/text HY-WorldPlay Predictor."""3 4from __future__ import annotations5 6import argparse7import json8import math9import os10import random11import time12from pathlib import Path13 14import torch15import torch.nn.functional as F16from torch.optim import AdamW17from torch.optim.lr_scheduler import LambdaLR18from torch.utils.data import DataLoader19 20from hyvideo.commons.infer_state import initialize_infer_state21from models import HYWorldPlayPredictor22from predictor_training.checkpoint import (23 load_predictor_weights,24 save_predictor_weights,25 save_training_state,26 trainable_state_dict,27)28from predictor_training.dataset import (29 PredictorPairDataset,30 move_pair_to_device,31 predictor_collate,32)33 34 35def parse_args() -> argparse.Namespace:36 parser = argparse.ArgumentParser(description=__doc__)37 parser.add_argument("--manifest", default="datasets/predictor_v1/train_eval_manifest.jsonl")38 parser.add_argument(39 "--teacher_checkpoint",40 default="models-ms/HY-WorldPlay/ar_distilled_action_model/diffusion_pytorch_model.safetensors",41 )42 parser.add_argument(43 "--output_dir",44 default=(45 "checkpoints/predictor_v1_blocks1-52_bf16_1gpu_"46 "flr1e-4_blr1e-5_schedsteps2000_swarmup100_bwarmup100"47 ),48 )49 parser.add_argument("--max_steps", type=int, default=2000)50 parser.add_argument("--gradient_accumulation_steps", type=int, default=4)51 parser.add_argument("--warmup_blocks_steps", type=int, default=100)52 parser.add_argument("--scheduler_warmup_steps", type=int, default=100)53 parser.add_argument("--fusion_lr", type=float, default=1e-4)54 parser.add_argument("--blocks_lr", type=float, default=1e-5)55 parser.add_argument("--weight_decay", type=float, default=0.01)56 parser.add_argument("--hidden_weight", type=float, default=0.1)57 parser.add_argument("--velocity_weight", type=float, default=1.0)58 parser.add_argument("--grad_clip", type=float, default=1.0)59 parser.add_argument("--log_every", type=int, default=10)60 parser.add_argument("--save_every", type=int, default=500)61 parser.add_argument("--keep_snapshots", type=int, default=2)62 parser.add_argument("--seed", type=int, default=0)63 parser.add_argument("--num_workers", type=int, default=0)64 parser.add_argument("--max_records", type=int, default=None)65 parser.add_argument("--gradient_checkpointing", action=argparse.BooleanOptionalAction, default=True)66 parser.add_argument("--save_optimizer_state", action=argparse.BooleanOptionalAction, default=True)67 parser.add_argument("--resume", default=None)68 return parser.parse_args()69 70 71def initialize_runtime_state() -> None:72 class Args:73 sage_blocks_range = "0-53"74 use_sageattn = False75 enable_torch_compile = False76 use_fp8_gemm = False77 quant_type = "fp8-per-block"78 include_patterns = "double_blocks"79 use_vae_parallel = False80 81 initialize_infer_state(Args())82 83 84def make_scheduler(optimizer: torch.optim.Optimizer, warmup: int, total: int) -> LambdaLR:85 def factor(step: int) -> float:86 if step < warmup:87 return max(1e-8, float(step + 1) / max(1, warmup))88 progress = min(1.0, float(step - warmup) / max(1, total - warmup))89 return 0.5 * (1.0 + math.cos(math.pi * progress))90 91 return LambdaLR(optimizer, factor)92 93 94def set_blocks_trainable(model: HYWorldPlayPredictor, enabled: bool) -> None:95 model.double_blocks.requires_grad_(enabled)96 97 98def prune_snapshots(output_dir: Path, keep: int) -> None:99 snapshots = sorted(output_dir.glob("predictor_step_*.safetensors"))100 for path in snapshots[:-keep] if keep > 0 else snapshots:101 path.unlink()102 103 104def main() -> None:105 args = parse_args()106 if not torch.cuda.is_available():107 raise RuntimeError("CUDA is required")108 random.seed(args.seed)109 torch.manual_seed(args.seed)110 torch.cuda.manual_seed_all(args.seed)111 initialize_runtime_state()112 device = torch.device("cuda:0")113 dtype = torch.bfloat16114 output_dir = Path(args.output_dir).resolve()115 output_dir.mkdir(parents=True, exist_ok=True)116 log_path = output_dir / "train_log.jsonl"117 118 dataset = PredictorPairDataset(args.manifest, max_records=args.max_records)119 generator = torch.Generator().manual_seed(args.seed)120 loader = DataLoader(121 dataset,122 batch_size=1,123 shuffle=True,124 num_workers=args.num_workers,125 pin_memory=True,126 collate_fn=predictor_collate,127 generator=generator,128 )129 130 model = HYWorldPlayPredictor()131 model.load_teacher_initialization(args.teacher_checkpoint)132 model.to(device=device, dtype=dtype)133 model.enable_gradient_checkpointing(args.gradient_checkpointing)134 model.train()135 136 optimizer = AdamW(137 [138 {139 "params": list(model.img_fusion.parameters())140 + list(model.txt_fusion.parameters())141 + list(model.residual_out.parameters()),142 "lr": args.fusion_lr,143 },144 {"params": model.double_blocks.parameters(), "lr": args.blocks_lr},145 ],146 betas=(0.9, 0.95),147 weight_decay=args.weight_decay,148 )149 scheduler = make_scheduler(optimizer, args.scheduler_warmup_steps, args.max_steps)150 global_step = 0151 micro_step = 0152 set_blocks_trainable(model, args.warmup_blocks_steps <= 0)153 154 if args.resume:155 state = torch.load(args.resume, map_location="cpu", weights_only=False)156 model.load_state_dict(state["model"], strict=False)157 optimizer.load_state_dict(state["optimizer"])158 scheduler.load_state_dict(state["scheduler"])159 global_step = int(state["global_step"])160 micro_step = int(state.get("micro_step", global_step * args.gradient_accumulation_steps))161 if "torch_rng_state" in state:162 torch.set_rng_state(state["torch_rng_state"])163 if "cuda_rng_state" in state:164 torch.cuda.set_rng_state_all(state["cuda_rng_state"])165 set_blocks_trainable(model, global_step >= args.warmup_blocks_steps)166 167 print(168 json.dumps(169 {170 "dataset_pairs": len(dataset),171 "trainable_parameters": model.trainable_parameter_count(),172 "breakdown": model.trainable_parameter_breakdown(),173 "start_step": global_step,174 },175 indent=2,176 ),177 flush=True,178 )179 180 optimizer.zero_grad(set_to_none=True)181 running: list[dict[str, float | str]] = []182 epoch = 0183 while global_step < args.max_steps:184 epoch += 1185 for item in loader:186 if global_step >= args.max_steps:187 break188 if global_step == args.warmup_blocks_steps:189 set_blocks_trainable(model, True)190 191 inputs, target_hidden, target_velocity = move_pair_to_device(192 item, device=device, dtype=dtype193 )194 started = time.perf_counter()195 with torch.autocast(device_type="cuda", dtype=dtype):196 output = model(**inputs)197 hidden_loss = F.mse_loss(198 output["pred_hidden"].float(), target_hidden.float()199 )200 velocity_loss = F.mse_loss(201 output["pred_velocity"].float(), target_velocity.float()202 )203 loss = (204 args.hidden_weight * hidden_loss205 + args.velocity_weight * velocity_loss206 )207 scaled_loss = loss / args.gradient_accumulation_steps208 scaled_loss.backward()209 micro_step += 1210 running.append(211 {212 "pair_id": item["pair_id"],213 "loss": float(loss.detach()),214 "hidden_mse": float(hidden_loss.detach()),215 "velocity_mse": float(velocity_loss.detach()),216 "micro_time_s": time.perf_counter() - started,217 }218 )219 del output, loss, scaled_loss, hidden_loss, velocity_loss220 221 if micro_step % args.gradient_accumulation_steps:222 continue223 grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), args.grad_clip)224 optimizer.step()225 scheduler.step()226 optimizer.zero_grad(set_to_none=True)227 global_step += 1228 229 if global_step % args.log_every == 0 or global_step == 1:230 torch.cuda.synchronize()231 count = len(running)232 record = {233 "global_step": global_step,234 "epoch": epoch,235 "loss": sum(float(x["loss"]) for x in running) / count,236 "hidden_mse": sum(float(x["hidden_mse"]) for x in running) / count,237 "velocity_mse": sum(float(x["velocity_mse"]) for x in running) / count,238 "grad_norm": float(grad_norm),239 "lr_fusion": optimizer.param_groups[0]["lr"],240 "lr_blocks": optimizer.param_groups[1]["lr"],241 "peak_memory_gib": torch.cuda.max_memory_allocated() / 2**30,242 "avg_micro_time_s": sum(float(x["micro_time_s"]) for x in running) / count,243 }244 with log_path.open("a", encoding="utf-8") as handle:245 handle.write(json.dumps(record, sort_keys=True) + "\n")246 print(json.dumps(record, sort_keys=True), flush=True)247 running.clear()248 249 should_save = global_step % args.save_every == 0 or global_step == args.max_steps250 if should_save:251 weights_path = output_dir / f"predictor_step_{global_step:05d}.safetensors"252 save_predictor_weights(model, weights_path)253 if args.save_optimizer_state:254 latest = {255 "model": trainable_state_dict(model),256 "optimizer": optimizer.state_dict(),257 "scheduler": scheduler.state_dict(),258 "global_step": global_step,259 "micro_step": micro_step,260 "torch_rng_state": torch.get_rng_state(),261 "cuda_rng_state": torch.cuda.get_rng_state_all(),262 "args": vars(args),263 "predictor_config": model.config_dict,264 "weights_path": str(weights_path),265 }266 save_training_state(output_dir / "training_latest.pt", latest)267 prune_snapshots(output_dir, args.keep_snapshots)268 269 270if __name__ == "__main__":271 main()272 