Instructions to use xfcghj/AR with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Diffusers
How to use xfcghj/AR with Diffusers:
pip install -U diffusers transformers accelerate
import torch from diffusers import DiffusionPipeline # switch to "mps" for apple devices pipe = DiffusionPipeline.from_pretrained("xfcghj/AR", dtype=torch.bfloat16, device_map="cuda") prompt = "Astronaut in a jungle, cold color palette, muted colors, detailed, 8k" image = pipe(prompt).images[0] - Notebooks
- Google Colab
- Kaggle
| 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) | |