PointRSP / train_vqvae.py
Mo-nan's picture
Upload train_vqvae.py
7b92f7d verified
Raw
History Blame Contribute Delete
5.62 kB
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()