import os import argparse import logging from datetime import datetime import numpy as np import torch import torch.nn as nn import torch.optim as optim import torch.distributed as dist from torch.nn.parallel import DistributedDataParallel as DDP from tqdm import tqdm from models.auto import StackedTrajectorySpatialTextRefiner from models.vae import TripoSGVaeWrapper from models.clip import FrozenCLIPTextEncoder from utils.autodataloader import get_mesh_dataloader from utils.loss import TokenL2Loss class TextProjectionLayer(nn.Module): def __init__(self, in_dim=512, out_dim=64): super().__init__() self.net = nn.Sequential( nn.Linear(in_dim, in_dim), nn.ReLU(), nn.Linear(in_dim, out_dim), ) self.residual = nn.Linear(in_dim, out_dim) if in_dim != out_dim else nn.Identity() def forward(self, x): return self.net(x) + self.residual(x) def setup_logger(run_dir, rank): logger = logging.getLogger("TrainLogger") if logger.hasHandlers(): logger.handlers.clear() if rank == 0: logger.setLevel(logging.INFO) formatter = logging.Formatter("%(asctime)s - %(message)s") log_file = os.path.join(run_dir, "train.log") file_handler = logging.FileHandler(log_file) file_handler.setFormatter(formatter) logger.addHandler(file_handler) console_handler = logging.StreamHandler() console_handler.setFormatter(formatter) logger.addHandler(console_handler) else: logger.setLevel(logging.WARNING) return logger def count_parameters(model, name, logger): trainable = sum(p.numel() for p in model.parameters() if p.requires_grad) logger.info(f"[{name}] trainable parameters: {trainable:,}") return trainable def log_grad_stats(model, name, logger, topk=None): rows = [] for param_name, param in model.named_parameters(): if not param.requires_grad: continue if param.grad is None: rows.append((param_name, "grad=None")) continue grad = param.grad.detach() grad_norm = grad.norm(2).item() grad_mean_abs = grad.abs().mean().item() grad_max_abs = grad.abs().max().item() grad_min_abs = grad.abs().min().item() zero_ratio = (grad == 0).float().mean().item() rows.append( ( param_name, grad_norm, grad_mean_abs, grad_max_abs, grad_min_abs, zero_ratio, ) ) logger.info(f"--- gradient statistics: {name} ---") for row in rows[:topk] if topk is not None else rows: if row[1] == "grad=None": logger.info(f"{name}.{row[0]} | grad=None") else: param_name, grad_norm, grad_mean_abs, grad_max_abs, grad_min_abs, zero_ratio = row logger.info( f"{name}.{param_name} | " f"norm={grad_norm:.3e} | " f"mean_abs={grad_mean_abs:.3e} | " f"max_abs={grad_max_abs:.3e} | " f"min_abs={grad_min_abs:.3e} | " f"zero_ratio={zero_ratio:.3f}" ) def squeeze_extra_text_dim(x): while x.dim() > 3 and x.shape[1] == 1: x = x.squeeze(1) return x def train(args): local_rank = int(os.environ["LOCAL_RANK"]) torch.cuda.set_device(local_rank) device = torch.device("cuda", local_rank) dist.init_process_group(backend="nccl", device_id=device) rank = dist.get_rank() run_name = datetime.now().strftime("%Y%m%d_%H%M%S") run_dir = os.path.join(args.save_dir, run_name) if rank == 0: os.makedirs(run_dir, exist_ok=True) dist.barrier() logger = setup_logger(run_dir, rank) vae = TripoSGVaeWrapper(device=device) vae.eval() for param in vae.parameters(): param.requires_grad = False text_encoder = FrozenCLIPTextEncoder(args.clip_model_path, device=device) text_proj = TextProjectionLayer(in_dim=512, out_dim=args.token_dim).to(device) auto_model = StackedTrajectorySpatialTextRefiner( num_groups=args.num_refine_groups, temporal_kwargs={ "num_tokens": args.num_tokens, "token_dim": args.token_dim, "max_seq_len": 16, "num_blocks": 6, }, spatial_kwargs={ "num_tokens": args.num_tokens, "token_dim": args.token_dim, "num_blocks": 3, }, text_kwargs={ "token_dim": args.token_dim, "text_dim": args.token_dim, }, ).to(device) text_proj = DDP(text_proj, device_ids=[local_rank]) auto_model = DDP(auto_model, device_ids=[local_rank]) optimizer = optim.Adam( list(auto_model.parameters()) + list(text_proj.parameters()), lr=args.lr, ) logger.info("--- parameter statistics ---") count_parameters(auto_model, "AutoRefiner", logger) count_parameters(text_proj, "TextProjection", logger) start_epoch = 0 global_step = 0 if args.auto_resume: if args.last_ckpt and os.path.exists(args.last_ckpt): logger.info(f"Resuming training from checkpoint: {args.last_ckpt}") checkpoint = torch.load(args.last_ckpt, map_location="cpu") auto_model.module.load_state_dict(checkpoint["auto"]) text_proj.module.load_state_dict(checkpoint["proj"]) if "optimizer" in checkpoint: optimizer.load_state_dict(checkpoint["optimizer"]) for state in optimizer.state.values(): for k, v in state.items(): if isinstance(v, torch.Tensor): state[k] = v.to(device) if "epoch" in checkpoint: start_epoch = checkpoint["epoch"] + 1 if "step" in checkpoint: global_step = checkpoint["step"] logger.info(f"Resume finished. Continue from Epoch {start_epoch + 1}, Global Step {global_step}.") else: logger.warning( f"auto_resume=True, but no valid checkpoint was found at: {args.last_ckpt}. " "Training will start from scratch." ) criterion = TokenL2Loss(matching_method=args.matching_method) data_dir = "/home/dataset-assist-0/usr/lh/ysh/dw/RL/AR/data/shards/" dataloader = get_mesh_dataloader( data_dir=data_dir, batch_size=args.batch_size, num_workers=12, ) bar_format = "{l_bar}{bar:30}| {n_fmt}/{total_fmt} [{elapsed}<{remaining}, {rate_fmt}{postfix}]" for epoch in range(start_epoch, args.epochs): auto_model.train() text_proj.train() pbar = tqdm( dataloader, desc=f"Epoch {epoch + 1:03d}", disable=(rank != 0), ncols=1200, bar_format=bar_format, ) for batch_idx, batch in enumerate(pbar): optimizer.zero_grad() vertices, faces = batch["vertices"], batch["faces"] with torch.no_grad(): latents_batch = [] for i in range(len(vertices)): sample_latents = [] for t in range(vertices[i].shape[0]): v_t = vertices[i][t].detach().clone().to(device=device, dtype=torch.float32) f_t = faces[i].detach().clone().to(device=device, dtype=torch.float32) latent = vae.encode_mesh(v_t, f_t) if isinstance(latent, np.ndarray): latent = torch.from_numpy(latent).to(device) sample_latents.append(latent.squeeze(0)) latents_batch.append(torch.stack(sample_latents, dim=0)) latents = torch.stack(latents_batch, dim=0) text_tokens = text_encoder(list(batch["caption"])) text_tokens = squeeze_extra_text_dim(text_tokens) text_embed = text_proj(text_tokens) text_embed = squeeze_extra_text_dim(text_embed) history = latents[:, :-1, :, :] target = latents[:, 1:, :, :] shortest_path_matrix = torch.zeros( (len(batch["vertices"]), args.num_tokens, args.num_tokens), device=device, dtype=torch.long, ) refined_pred = auto_model( history, text_embed=text_embed, shortest_path_matrix=shortest_path_matrix, ) loss = criterion( history=history, refined_pred=refined_pred, target=target, ) total_loss = loss["total"] if torch.isnan(total_loss): logger.error(f"NaN detected at Epoch {epoch + 1}, Batch {batch_idx + 1}") raise ValueError("Training stopped: NaN loss detected.") total_loss.backward() if rank == 0 and global_step % args.grad_log_interval == 0: log_grad_stats(auto_model.module, "AutoRefiner", logger) log_grad_stats(text_proj.module, "TextProjection", logger) for handler in logger.handlers: handler.flush() torch.nn.utils.clip_grad_norm_( list(auto_model.parameters()) + list(text_proj.parameters()), max_norm=1.0, ) optimizer.step() global_step += 1 if rank == 0: if global_step % 500 == 0: ckpt_path = os.path.join(run_dir, f"model_step_{global_step}.pth") torch.save( { "auto": auto_model.module.state_dict(), "proj": text_proj.module.state_dict(), "optimizer": optimizer.state_dict(), "epoch": epoch, "step": global_step, }, ckpt_path, ) logger.info(f"Checkpoint saved at Step {global_step}: {ckpt_path}") if (batch_idx + 1) % 10 == 0: log_msg = ( f"Epoch {epoch + 1:03d} | Batch {batch_idx + 1:04d} | " f"Total: {total_loss.item():.4f} | " f"Original MSE: {loss['original_mse'].item():.4f} | " f"Matched MSE: {loss['matched_mse'].item():.4f}" ) logger.info(log_msg) for handler in logger.handlers: handler.flush() pbar.set_postfix( { "Loss": f"{total_loss.item():.3f}", "Orig": f"{loss['original_mse'].item():.3f}", "Match": f"{loss['matched_mse'].item():.3f}", } ) if rank == 0 and (epoch + 1) % args.save_every_k_epochs == 0: ckpt_path = os.path.join(run_dir, f"model_epoch_{epoch + 1}.pth") torch.save( { "auto": auto_model.module.state_dict(), "proj": text_proj.module.state_dict(), "optimizer": optimizer.state_dict(), "epoch": epoch, "step": global_step, }, ckpt_path, ) logger.info(f"Model saved: {ckpt_path}") dist.destroy_process_group() if __name__ == "__main__": parser = argparse.ArgumentParser() parser.add_argument("--num_tokens", type=int, default=512) parser.add_argument("--token_dim", type=int, default=64) parser.add_argument("--batch_size", type=int, default=4) parser.add_argument("--lr", type=float, default=5e-4) parser.add_argument("--epochs", type=int, default=100) parser.add_argument("--save_every_k_epochs", type=int, default=1) parser.add_argument("--save_dir", type=str, default="./checkpoints") parser.add_argument( "--clip_model_path", type=str, default="/home/dataset-assist-0/usr/lh/ysh/dw/RL/AR/pretrain/clip-vit-base-patch32", ) parser.add_argument("--num_refine_groups", type=int, default=3) parser.add_argument("--matching_method", type=str, default="hungarian", choices=["hungarian", "greedy"]) parser.add_argument("--auto_resume", action="store_true", help="whether to resume training from a checkpoint") parser.add_argument("--last_ckpt", type=str, default="", help="checkpoint path to resume from (.pth)") parser.add_argument("--grad_log_interval", type=int, default=100000) args = parser.parse_args() train(args)