import os import torch import torch.optim as optim import torch.nn as nn from model.vqvae import VQVAE from datasets.data_processing import get_dataset from datetime import datetime from copy import deepcopy import time def main(): # Hyperparameters input_dim = 3 # x, y, z coordinates hidden_dim = 1024 codebook_size = 4096 embedding_dim = 64 batch_size = 1 num_epochs = store_epoch + 1 default_lr = 4e-5 encoder_lr = 3e-5 num_points = 15000 voxel_size = 10 spconv_channels = 32 blocks_num = 6 smooth_end_epoch = 10 beta = 0.35 lambda_chamfer = 40.0 lambda_vq_start = 0.01 lambda_vq_final = 0.05 lambda_construction = 100.0 lambda_usage = 0.02 warmup_steps = 80 ema_decay = 0.9995 token_N1_stage = 9 token_N2_stage = 14 token_N3_stage = 14 decoder_layer = 5 buffer_stage_num = 3 ckpts_save_path = None ckpt_save_step = 20 if_load_ckpt = True if_only_init_param = True ckpt_name = None device = torch.device("cuda" if torch.cuda.is_available() else "cpu") time_commit = 0 # Initialize model model = VQVAE(input_dim, hidden_dim, codebook_size, embedding_dim, num_points, voxel_size, spconv_channels, blocks_num, smooth_end_epoch, beta, lambda_chamfer, lambda_vq_start, lambda_vq_final, lambda_construction, lambda_usage, warmup_steps, ema_decay, token_N1_stage, token_N2_stage, token_N3_stage, decoder_layer, buffer_stage_num).to(device) optimizer = torch.optim.Adam([ {"params": model.encoder.parameters(), "lr": encoder_lr}, {"params": model.pos_linear.parameters(), "lr": encoder_lr}, {"params": model.act.parameters(), "lr": encoder_lr}, {"params": model.after_spconv.parameters(), "lr": encoder_lr}, {"params": model.phi.parameters(), "lr": default_lr}, {"params": model.decoder.parameters(), "lr": default_lr}, ], weight_decay=1e-5) curr_epoch = 0 opt = None if if_load_ckpt: ckpt_path = os.path.join(ckpts_save_path, ckpt_name) if not if_only_init_param: # start from curr_epoch opt, curr_epoch = model.init_from_ckpt(ckpt_path, None, False) else: # start from 0 epoch _, _ = model.init_from_ckpt(ckpt_path, None, True) if opt is not None: optimizer.load_state_dict(opt) # Load ShapeNetV2 data root_dir = "data/ShapeNetCore.v2.PC15k" # PC15k categories = ['airplane'] # Example: cars and chairs train_dataset, test_dataset = get_dataset(root_dir, num_points, category=categories) dataloader = torch.utils.data.DataLoader(train_dataset, batch_size=batch_size, shuffle=True, drop_last=True) # Train the model for epoch in range(curr_epoch, num_epochs): total_loss = 0 total_construction_loss = 0 total_vq_loss = 0 total_chamfer_loss = 0 total_usage_loss = 0 global_step = 0 log_data_total = torch.zeros(8, dtype=torch.float32) for batch in dataloader: # time1 = time.perf_counter() x = batch['train_points'].to(device) optimizer.zero_grad() reconstructed, loss, construction_loss, vq_loss, chamfer_loss, usage_loss = model(x, epoch, global_step) loss.backward() # Gradient Clipping to avoid exploding updates torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) optimizer.step() total_loss += loss.item() total_construction_loss += construction_loss.item() total_vq_loss += vq_loss.item() total_chamfer_loss += chamfer_loss.item() total_usage_loss += usage_loss.item() print(f"VQ_Loss: {vq_loss:.4f}, construction_Loss: {construction_loss:.4f} , chamfer_Loss: {chamfer_loss:.4f} , usage_Loss: {usage_loss:.4f}" ) log_data_total += model.log_data global_step += 1 if epoch % ckpt_save_step == 0: save_checkpoint(model, optimizer, None, epoch, ckpts_save_path) log_data_avg = log_data_total / global_step avg_loss = total_loss / len(dataloader) avg_construction_loss = total_construction_loss / len(dataloader) avg_vq_loss = total_vq_loss / len(dataloader) print(f"Epoch [{epoch+1}/{num_epochs}], " f"Loss: {avg_loss:.4f}, " f"Construction Loss: {avg_construction_loss:.4f}, " f"VQ Loss: {avg_vq_loss:.4f}") def save_checkpoint(model, optimizer=None, scheduler=None, epoch=0, save_root="checkpoints"): date_str = datetime.now().strftime("%m%d") save_dir = os.path.join(save_root, date_str) os.makedirs(save_dir, exist_ok=True) ckpt_path = os.path.join(save_dir, f"epoch{epoch}.pth") ckpt = { "epoch": epoch, "model": model.state_dict(), } if optimizer is not None: ckpt["optimizer"] = optimizer.state_dict() if scheduler is not None: ckpt["scheduler"] = scheduler.state_dict() torch.save(ckpt, ckpt_path) print(f"✔ Checkpoint saved to: {ckpt_path}") if __name__ == "__main__": main()