Cccccz/HY
0
1#!/usr/bin/env python32"""Single-GPU training for AR-context HY-WorldPlay Predictor-v2."""3 4from __future__ import annotations5 6import argparse7import json8import math9import os10import random11import time12from pathlib import Path13 14import swanlab15import torch16import torch.nn.functional as F17from torch.optim import AdamW18from torch.optim.lr_scheduler import LambdaLR19from torch.utils.data import DataLoader20 21from hyvideo.commons.infer_state import initialize_infer_state22from models import HYWorldPlayPredictorV2, HYWorldPlayPredictorV3, PredictorV2Config23from predictor_training.checkpoint import (24 save_predictor_weights,25 save_training_state,26 trainable_state_dict,27)28from predictor_training.dataset_v2 import (29 PredictorV2PairDataset,30 move_v2_pair_to_device,31 predictor_v2_collate,32)33from predictor_training.dataset_v3 import (34 PredictorV3PairDataset,35 move_v3_pair_to_device,36 predictor_v3_collate,37)38 39 40def parse_source_blocks(value: str) -> tuple[int, int]:41 result = tuple(int(part.strip()) for part in value.split(",") if part.strip())42 if result not in {(0, 53), (1, 52)}:43 raise argparse.ArgumentTypeError("source_blocks must be 0,53 or 1,52")44 return result45 46 47def parse_args() -> argparse.Namespace:48 parser = argparse.ArgumentParser(description=__doc__)49 parser.add_argument("--manifest", default="datasets/predictor_v2/train_eval_manifest.jsonl")50 parser.add_argument("--predictor_version", choices=("v2", "v3"), default="v2")51 parser.add_argument("--source_blocks", type=parse_source_blocks, required=True)52 parser.add_argument(53 "--teacher_checkpoint",54 default="models-ms/HY-WorldPlay/ar_distilled_action_model/diffusion_pytorch_model.safetensors",55 )56 parser.add_argument("--output_dir", required=True)57 parser.add_argument("--max_steps", type=int, default=2000)58 parser.add_argument("--batch_size", type=int, default=1)59 parser.add_argument("--gradient_accumulation_steps", type=int, default=4)60 parser.add_argument("--warmup_blocks_steps", type=int, default=100)61 parser.add_argument("--scheduler_warmup_steps", type=int, default=100)62 parser.add_argument("--fusion_lr", type=float, default=1e-4)63 parser.add_argument("--blocks_lr", type=float, default=1e-5)64 parser.add_argument("--weight_decay", type=float, default=0.01)65 parser.add_argument("--hidden_weight", type=float, default=0.1)66 parser.add_argument("--velocity_weight", type=float, default=1.0)67 parser.add_argument("--grad_clip", type=float, default=1.0)68 parser.add_argument("--log_every", type=int, default=10)69 parser.add_argument("--save_every", type=int, default=500)70 parser.add_argument("--keep_snapshots", type=int, default=2)71 parser.add_argument("--seed", type=int, default=0)72 parser.add_argument("--num_workers", type=int, default=0)73 parser.add_argument("--max_records", type=int, default=None)74 parser.add_argument("--gradient_checkpointing", action=argparse.BooleanOptionalAction, default=True)75 parser.add_argument("--save_optimizer_state", action=argparse.BooleanOptionalAction, default=True)76 parser.add_argument("--resume", default=None)77 parser.add_argument("--swanlab_project", default="HY-WorldPlay-Predictor-v2")78 parser.add_argument("--swanlab_workspace", default=None)79 parser.add_argument("--swanlab_experiment", default=None)80 parser.add_argument("--swanlab_mode", choices=("cloud", "offline", "local", "disabled"), default="cloud")81 return parser.parse_args()82 83 84def initialize_runtime_state() -> None:85 class Args:86 sage_blocks_range = "0-53"87 use_sageattn = False88 enable_torch_compile = False89 use_fp8_gemm = False90 quant_type = "fp8-per-block"91 include_patterns = "double_blocks"92 use_vae_parallel = False93 94 initialize_infer_state(Args())95 96 97def make_scheduler(optimizer: torch.optim.Optimizer, warmup: int, total: int) -> LambdaLR:98 def factor(step: int) -> float:99 if step < warmup:100 return max(1e-8, float(step + 1) / max(1, warmup))101 progress = min(1.0, float(step - warmup) / max(1, total - warmup))102 return 0.5 * (1.0 + math.cos(math.pi * progress))103 104 return LambdaLR(optimizer, factor)105 106 107def set_ar_blocks_trainable(model: HYWorldPlayPredictorV2, enabled: bool) -> None:108 for block in model.double_blocks:109 for name, parameter in block.named_parameters():110 parameter.requires_grad_(enabled and not name.startswith("txt_"))111 112 113def prune_snapshots(output_dir: Path, keep: int) -> None:114 snapshots = sorted(output_dir.glob("predictor_step_*.safetensors"))115 for path in snapshots[:-keep] if keep > 0 else snapshots:116 path.unlink()117 118 119def init_swanlab(args: argparse.Namespace, config: dict):120 if args.swanlab_mode == "cloud":121 api_key = os.environ.get("SWANLAB_API_KEY")122 swanlab.login(api_key=api_key, save=False) if api_key else swanlab.login()123 fusion_lr = f"{args.fusion_lr:.0e}".replace("e-0", "e-").replace("e+0", "e+")124 blocks_lr = f"{args.blocks_lr:.0e}".replace("e-0", "e-").replace("e+0", "e+")125 experiment = args.swanlab_experiment or (126 f"predictor-{args.predictor_version}-blocks"127 + "-".join(str(value) for value in args.source_blocks)128 + f"-flr{fusion_lr}-blr{blocks_lr}-schedsteps{args.max_steps}"129 + f"-swarmup{args.scheduler_warmup_steps}"130 + f"-bwarmup{args.warmup_blocks_steps}"131 )132 return swanlab.init(133 project=args.swanlab_project,134 workspace=args.swanlab_workspace,135 experiment_name=experiment,136 description=(137 f"HY-WorldPlay {args.predictor_version} AR-context Predictor, "138 "fresh exact 10-case targets"139 ),140 tags=[141 f"predictor-{args.predictor_version}",142 "ar-context-kv",143 f"blocks-{args.source_blocks[0]}-{args.source_blocks[1]}",144 ],145 config=config,146 logdir=str(Path(args.output_dir).resolve() / "swanlab"),147 mode=args.swanlab_mode,148 )149 150 151def main() -> None:152 args = parse_args()153 if not torch.cuda.is_available():154 raise RuntimeError("CUDA is required")155 random.seed(args.seed)156 torch.manual_seed(args.seed)157 torch.cuda.manual_seed_all(args.seed)158 initialize_runtime_state()159 device = torch.device("cuda:0")160 dtype = torch.bfloat16161 output_dir = Path(args.output_dir).resolve()162 output_dir.mkdir(parents=True, exist_ok=True)163 log_path = output_dir / "train_log.jsonl"164 165 dataset_class = (166 PredictorV3PairDataset if args.predictor_version == "v3" else PredictorV2PairDataset167 )168 collate_fn = (169 predictor_v3_collate if args.predictor_version == "v3" else predictor_v2_collate170 )171 move_pair = (172 move_v3_pair_to_device173 if args.predictor_version == "v3"174 else move_v2_pair_to_device175 )176 dataset = dataset_class(177 args.manifest,178 source_block_ids=args.source_blocks,179 max_records=args.max_records,180 )181 generator = torch.Generator().manual_seed(args.seed)182 loader = DataLoader(183 dataset,184 batch_size=args.batch_size,185 shuffle=True,186 num_workers=args.num_workers,187 pin_memory=False,188 collate_fn=collate_fn,189 generator=generator,190 )191 192 model_class = (193 HYWorldPlayPredictorV3 if args.predictor_version == "v3" else HYWorldPlayPredictorV2194 )195 model = model_class(PredictorV2Config(source_block_ids=args.source_blocks))196 model.load_teacher_initialization(args.teacher_checkpoint)197 model.to(device=device, dtype=dtype)198 model.enable_gradient_checkpointing(args.gradient_checkpointing)199 model.train()200 201 fusion_parameters = list(model.img_fusion.parameters()) + list(model.residual_out.parameters())202 block_parameters = [203 parameter for parameter in model.double_blocks.parameters() if parameter.requires_grad204 ]205 intended_trainable_parameters = sum(206 parameter.numel() for parameter in fusion_parameters + block_parameters207 )208 intended_parameter_breakdown = model.trainable_parameter_breakdown()209 optimizer = AdamW(210 [211 {"params": fusion_parameters, "lr": args.fusion_lr},212 {"params": block_parameters, "lr": args.blocks_lr},213 ],214 betas=(0.9, 0.95),215 weight_decay=args.weight_decay,216 )217 scheduler = make_scheduler(optimizer, args.scheduler_warmup_steps, args.max_steps)218 global_step = 0219 micro_step = 0220 set_ar_blocks_trainable(model, args.warmup_blocks_steps <= 0)221 222 if args.resume:223 state = torch.load(args.resume, map_location="cpu", weights_only=False)224 if tuple(state["predictor_config"]["source_block_ids"]) != args.source_blocks:225 raise ValueError("Resume checkpoint source blocks do not match")226 model.load_state_dict(state["model"], strict=False)227 optimizer.load_state_dict(state["optimizer"])228 scheduler.load_state_dict(state["scheduler"])229 global_step = int(state["global_step"])230 micro_step = int(state.get("micro_step", global_step * args.gradient_accumulation_steps))231 if "torch_rng_state" in state:232 torch.set_rng_state(state["torch_rng_state"])233 if "cuda_rng_state" in state:234 torch.cuda.set_rng_state_all(state["cuda_rng_state"])235 set_ar_blocks_trainable(model, global_step >= args.warmup_blocks_steps)236 237 run_config = {238 **vars(args),239 "source_blocks": list(args.source_blocks),240 "manifest": str(Path(args.manifest).resolve()),241 "teacher_checkpoint": str(Path(args.teacher_checkpoint).resolve()),242 "dataset_pairs": len(dataset),243 "effective_batch_size": args.batch_size * args.gradient_accumulation_steps,244 "trainable_parameters": intended_trainable_parameters,245 "parameter_breakdown": intended_parameter_breakdown,246 "dtype": "bfloat16",247 "predictor_version": args.predictor_version,248 "gpu": torch.cuda.get_device_name(device),249 }250 run = init_swanlab(args, run_config)251 print(252 json.dumps(253 {254 "dataset_pairs": len(dataset),255 "source_blocks": args.source_blocks,256 "trainable_parameters": intended_trainable_parameters,257 "breakdown": intended_parameter_breakdown,258 "start_step": global_step,259 "swanlab_run_id": getattr(run, "id", None),260 },261 indent=2,262 ),263 flush=True,264 )265 266 optimizer.zero_grad(set_to_none=True)267 running: list[dict[str, float | str]] = []268 epoch = 0269 try:270 while global_step < args.max_steps:271 epoch += 1272 for item in loader:273 if global_step >= args.max_steps:274 break275 if global_step == args.warmup_blocks_steps:276 set_ar_blocks_trainable(model, True)277 278 inputs, target_hidden, target_velocity = move_pair(279 item, device=device, dtype=dtype280 )281 started = time.perf_counter()282 with torch.autocast(device_type="cuda", dtype=dtype):283 output = model(**inputs)284 hidden_loss = F.mse_loss(285 output["pred_hidden"].float(), target_hidden.float()286 )287 velocity_loss = F.mse_loss(288 output["pred_velocity"].float(), target_velocity.float()289 )290 loss = args.hidden_weight * hidden_loss + args.velocity_weight * velocity_loss291 scaled_loss = loss / args.gradient_accumulation_steps292 scaled_loss.backward()293 micro_step += 1294 running.append(295 {296 "pair_id": ",".join(item["pair_id"]),297 "loss": float(loss.detach()),298 "hidden_mse": float(hidden_loss.detach()),299 "velocity_mse": float(velocity_loss.detach()),300 "micro_time_s": time.perf_counter() - started,301 "context_frames": sum(item["context_frames"])302 / len(item["context_frames"]),303 }304 )305 del (306 output,307 loss,308 scaled_loss,309 hidden_loss,310 velocity_loss,311 inputs,312 target_hidden,313 target_velocity,314 item,315 )316 317 if micro_step % args.gradient_accumulation_steps:318 continue319 grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), args.grad_clip)320 optimizer.step()321 scheduler.step()322 optimizer.zero_grad(set_to_none=True)323 global_step += 1324 325 if global_step % args.log_every == 0 or global_step == 1:326 torch.cuda.synchronize()327 count = len(running)328 record = {329 "global_step": global_step,330 "epoch": epoch,331 "loss": sum(float(x["loss"]) for x in running) / count,332 "hidden_mse": sum(float(x["hidden_mse"]) for x in running) / count,333 "velocity_mse": sum(float(x["velocity_mse"]) for x in running) / count,334 "grad_norm": float(grad_norm),335 "lr_fusion": optimizer.param_groups[0]["lr"],336 "lr_blocks": optimizer.param_groups[1]["lr"],337 "peak_memory_gib": torch.cuda.max_memory_allocated() / 2**30,338 "avg_micro_time_s": sum(float(x["micro_time_s"]) for x in running) / count,339 "avg_context_frames": sum(float(x["context_frames"]) for x in running) / count,340 }341 with log_path.open("a", encoding="utf-8") as handle:342 handle.write(json.dumps(record, sort_keys=True) + "\n")343 swanlab.log(record, step=global_step)344 print(json.dumps(record, sort_keys=True), flush=True)345 running.clear()346 347 should_save = global_step % args.save_every == 0 or global_step == args.max_steps348 if should_save:349 weights_path = output_dir / f"predictor_step_{global_step:05d}.safetensors"350 save_predictor_weights(model, weights_path)351 if args.save_optimizer_state:352 latest = {353 "model": trainable_state_dict(model),354 "optimizer": optimizer.state_dict(),355 "scheduler": scheduler.state_dict(),356 "global_step": global_step,357 "micro_step": micro_step,358 "torch_rng_state": torch.get_rng_state(),359 "cuda_rng_state": torch.cuda.get_rng_state_all(),360 "args": vars(args),361 "predictor_config": model.config_dict,362 "weights_path": str(weights_path),363 }364 save_training_state(output_dir / "training_latest.pt", latest)365 prune_snapshots(output_dir, args.keep_snapshots)366 except BaseException as error:367 swanlab.finish(error=error)368 raise369 else:370 swanlab.finish()371 372 373if __name__ == "__main__":374 main()375 