MCPcopy Create free account
hub / github.com/InternScience/InternAgent / main

Function main

tasks/AutoCls3D/code/experiment.py:266–410  ·  view source on GitHub ↗
(config)

Source from the content-addressed store, hash-verified

264 return point_set, label[0]
265
266def 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

Callers 1

experiment.pyFile · 0.70

Calls 15

updateMethod · 0.95
all_accMethod · 0.95
cal_accMethod · 0.95
ConfusionMatrixClass · 0.90
parametersMethod · 0.80
state_dictMethod · 0.80
ModelNetDatasetClass · 0.70
ModelClass · 0.70
getenvMethod · 0.45
trainMethod · 0.45
backwardMethod · 0.45
stepMethod · 0.45

Tested by

no test coverage detected