Diffusers
Safetensors
AR / train.py
xfcghj's picture
Upload folder using huggingface_hub
394918c verified
Raw
History Blame Contribute Delete
12.9 kB
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)