(config)
| 264 | return point_set, label[0] |
| 265 | |
| 266 | def main(config): |
| 267 | |
| 268 | pathlib.Path(config.out_dir).mkdir(parents=True, exist_ok=True) |
| 269 | |
| 270 | # 创建TensorBoard的SummaryWriter |
| 271 | log_dir = os.getenv("TENSORBOARD_LOG_PATH", "/tensorboard_logs/") |
| 272 | |
| 273 | pathlib.Path(log_dir).mkdir(parents=True, exist_ok=True) |
| 274 | writer = SummaryWriter(log_dir) |
| 275 | |
| 276 | datasets, dataloaders = {}, {} |
| 277 | for split in ['train', 'test']: |
| 278 | datasets[split] = ModelNetDataset(config.data_root, config.num_category, config.num_points, split) |
| 279 | dataloaders[split] = DataLoader(datasets[split], batch_size=config.batch_size, shuffle=(split == 'train'), |
| 280 | drop_last=(split == 'train'), num_workers=8) |
| 281 | |
| 282 | model = Model(in_channels=config.in_channels).cuda() |
| 283 | optimizer = torch.optim.Adam( |
| 284 | model.parameters(), lr=config.learning_rate, |
| 285 | betas=(0.9, 0.999), eps=1e-8, |
| 286 | weight_decay=1e-4 |
| 287 | ) |
| 288 | scheduler = torch.optim.lr_scheduler.StepLR( |
| 289 | optimizer, step_size=20, gamma=0.7 |
| 290 | ) |
| 291 | train_losses = [] |
| 292 | print("Training model...") |
| 293 | model.train() |
| 294 | global_step = 0 |
| 295 | cur_epoch = 0 |
| 296 | best_oa = 0 |
| 297 | best_acc = 0 |
| 298 | |
| 299 | start_time = time.time() |
| 300 | for epoch in tqdm(range(config.max_epoch), desc='training'): |
| 301 | model.train() |
| 302 | cm = ConfusionMatrix(num_classes=len(datasets['train'].classes)) |
| 303 | epoch_loss = 0.0 |
| 304 | batch_count = 0 |
| 305 | for points, target in tqdm(dataloaders['train'], desc=f'epoch {cur_epoch}/{config.max_epoch}'): |
| 306 | # data transforms |
| 307 | points = points.data.numpy() |
| 308 | points = data_transforms.random_point_dropout(points) |
| 309 | points[:, :, 0:3] = data_transforms.random_scale_point_cloud(points[:, :, 0:3]) |
| 310 | points[:, :, 0:3] = data_transforms.shift_point_cloud(points[:, :, 0:3]) |
| 311 | points = torch.from_numpy(points).transpose(2, 1).contiguous() |
| 312 | |
| 313 | points, target = points.cuda(), target.long().cuda() |
| 314 | |
| 315 | loss, logits = model(points, target) |
| 316 | loss.backward() |
| 317 | |
| 318 | torch.nn.utils.clip_grad_norm_(model.parameters(), 1, norm_type=2) |
| 319 | optimizer.step() |
| 320 | model.zero_grad() |
| 321 | loss_value = loss.detach().item() |
| 322 | epoch_loss += loss_value |
| 323 | batch_count += 1 |
no test coverage detected