MCPcopy Create free account
hub / github.com/OpenImagingLab/4DSloMo / training_report

Function training_report

train.py:277–346  ·  view source on GitHub ↗
(tb_writer, iteration, Ll1, loss, l1_loss, elapsed, testing_iterations, scene : Scene, renderFunc, renderArgs, loss_dict=None)

Source from the content-addressed store, hash-verified

275 return tb_writer
276
277def training_report(tb_writer, iteration, Ll1, loss, l1_loss, elapsed, testing_iterations, scene : Scene, renderFunc, renderArgs, loss_dict=None):
278 if tb_writer:
279 tb_writer.add_scalar('train_loss_patches/l1_loss', Ll1.item(), iteration)
280 tb_writer.add_scalar('train_loss_patches/ssim_loss', Ll1.item(), iteration)
281 tb_writer.add_scalar('train_loss_patches/total_loss', loss.item(), iteration)
282 tb_writer.add_scalar('iter_time', elapsed, iteration)
283 tb_writer.add_scalar('total_points', scene.gaussians.get_xyz.shape[0], iteration)
284 tb_writer.add_histogram("scene/opacity_histogram", scene.gaussians.get_opacity, iteration)
285 if loss_dict is not None:
286 if "Lrigid" in loss_dict:
287 tb_writer.add_scalar('train_loss_patches/rigid_loss', loss_dict['Lrigid'].item(), iteration)
288 if "Ldepth" in loss_dict:
289 tb_writer.add_scalar('train_loss_patches/depth_loss', loss_dict['Ldepth'].item(), iteration)
290 if "Ltv" in loss_dict:
291 tb_writer.add_scalar('train_loss_patches/tv_loss', loss_dict['Ltv'].item(), iteration)
292 if "Lopa" in loss_dict:
293 tb_writer.add_scalar('train_loss_patches/opa_loss', loss_dict['Lopa'].item(), iteration)
294 if "Lptsopa" in loss_dict:
295 tb_writer.add_scalar('train_loss_patches/pts_opa_loss', loss_dict['Lptsopa'].item(), iteration)
296 if "Lsmooth" in loss_dict:
297 tb_writer.add_scalar('train_loss_patches/smooth_loss', loss_dict['Lsmooth'].item(), iteration)
298 if "Llaplacian" in loss_dict:
299 tb_writer.add_scalar('train_loss_patches/laplacian_loss', loss_dict['Llaplacian'].item(), iteration)
300
301 psnr_test_iter = 0.0
302 # Report test and samples of training set
303 if iteration in testing_iterations:
304 validation_configs = ({'name': 'train', 'cameras' : [scene.getTrainCameras()[idx % len(scene.getTrainCameras())] for idx in range(5, 30, 5)]},
305 {'name': 'test', 'cameras' : [scene.getTestCameras()[idx] for idx in range(len(scene.getTestCameras()))]})
306
307 for config in validation_configs:
308 if config['cameras'] and len(config['cameras']) > 0:
309 l1_test = 0.0
310 psnr_test = 0.0
311 ssim_test = 0.0
312 msssim_test = 0.0
313 for idx, batch_data in enumerate(tqdm(config['cameras'])):
314 gt_image, viewpoint = batch_data
315 gt_image = gt_image.cuda()
316 viewpoint = viewpoint.cuda()
317
318 render_pkg = renderFunc(viewpoint, scene.gaussians, *renderArgs)
319 image = torch.clamp(render_pkg["render"], 0.0, 1.0)
320
321 depth = easy_cmap(render_pkg['depth'][0])
322 alpha = torch.clamp(render_pkg['alpha'], 0.0, 1.0).repeat(3,1,1)
323 if tb_writer and (idx < 5):
324 grid = [gt_image, image, alpha, depth]
325 grid = make_grid(grid, nrow=2)
326 tb_writer.add_images(config['name'] + "_view_{}/gt_vs_render".format(viewpoint.image_name), grid[None], global_step=iteration)
327
328 l1_test += l1_loss(image, gt_image).mean().double()
329 psnr_test += psnr(image, gt_image).mean().double()
330 ssim_test += ssim(image, gt_image).mean().double()
331 msssim_test += msssim(image[None].cpu(), gt_image[None].cpu())
332 psnr_test /= len(config['cameras'])
333 l1_test /= len(config['cameras'])
334 ssim_test /= len(config['cameras'])

Callers 1

trainingFunction · 0.85

Calls 8

easy_cmapFunction · 0.90
l1_lossFunction · 0.90
psnrFunction · 0.90
ssimFunction · 0.90
msssimFunction · 0.90
getTrainCamerasMethod · 0.80
getTestCamerasMethod · 0.80
cudaMethod · 0.80

Tested by

no test coverage detected