CoolFace
Modelpublic

Cccccz/HY

sourceHugging Faceupdated 11d agoView on Hugging Face
0likes
train_predictor.py272 linesDownload Raw Back to root
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