| 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(): |
| |
| input_dim = 3 |
| 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 |
| |
| |
| 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: |
| opt, curr_epoch = model.init_from_ckpt(ckpt_path, None, False) |
| else: |
| _, _ = model.init_from_ckpt(ckpt_path, None, True) |
|
|
| if opt is not None: |
| optimizer.load_state_dict(opt) |
|
|
| |
| root_dir = "data/ShapeNetCore.v2.PC15k" |
| categories = ['airplane'] |
| 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) |
| |
| |
| 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: |
| |
| 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() |
|
|
| |
| 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() |
|
|