MCPcopy Create free account
hub / github.com/computational-imaging/bacon / write_sdf_slice

Function write_sdf_slice

utils.py:583–638  ·  view source on GitHub ↗

write slices of sdf in each plane

(model, writer, total_steps, prefix='train_', is_multi=False)

Source from the content-addressed store, hash-verified

581
582
583def write_sdf_slice(model, writer, total_steps, prefix='train_', is_multi=False):
584 ''' write slices of sdf in each plane '''
585
586 slice_coords_2d = dataio.get_mgrid(512)
587
588 with torch.no_grad():
589
590 yz_slice_coords = torch.cat((torch.zeros_like(slice_coords_2d[:, :1]),
591 slice_coords_2d), dim=-1)
592 yz_slice_model_input = {'coords': yz_slice_coords.cuda()[None, ...]}
593 yz_model_out = model(yz_slice_model_input)
594
595 sdf_values = yz_model_out['model_out']
596 all_sdf_values = sdf_values if is_multi else [sdf_values, ]
597
598 fig, axs = plt.subplots(1, len(all_sdf_values), figsize=(2.75*len(all_sdf_values), 2.75), dpi=100)
599 for idx, sdf_values in enumerate(all_sdf_values):
600 sdf_values = dataio.lin2img(sdf_values).squeeze().cpu().numpy()
601 ax = axs if not isinstance(axs, np.ndarray) else axs[idx]
602 make_contour_plot(sdf_values, ax=ax)
603
604 writer.add_figure(prefix + 'yz_sdf_slice', fig, global_step=total_steps)
605
606 xz_slice_coords = torch.cat((slice_coords_2d[:, :1],
607 torch.zeros_like(slice_coords_2d[:, :1]),
608 slice_coords_2d[:, -1:]), dim=-1)
609 xz_slice_model_input = {'coords': xz_slice_coords.cuda()[None, ...]}
610
611 xz_model_out = model(xz_slice_model_input)
612 sdf_values = xz_model_out['model_out']
613 all_sdf_values = sdf_values if is_multi else [sdf_values, ]
614
615 fig, axs = plt.subplots(1, len(all_sdf_values), figsize=(2.75*len(all_sdf_values), 2.75), dpi=100)
616 for idx, sdf_values in enumerate(all_sdf_values):
617 sdf_values = dataio.lin2img(sdf_values).squeeze().cpu().numpy()
618 ax = axs if not isinstance(axs, np.ndarray) else axs[idx]
619 make_contour_plot(sdf_values, ax=ax)
620
621 writer.add_figure(prefix + 'xz_sdf_slice', fig, global_step=total_steps)
622
623 xy_slice_coords = torch.cat((slice_coords_2d[:, :2],
624 -0.75*torch.ones_like(slice_coords_2d[:, :1])), dim=-1)
625 xy_slice_model_input = {'coords': xy_slice_coords.cuda()[None, ...]}
626
627 xy_model_out = model(xy_slice_model_input)
628 sdf_values = xy_model_out['model_out']
629 all_sdf_values = sdf_values if is_multi else [sdf_values, ]
630
631 fig, axs = plt.subplots(1, len(all_sdf_values), figsize=(2.75*len(all_sdf_values), 2.75), dpi=100)
632 for idx, sdf_values in enumerate(all_sdf_values):
633 sdf_values = dataio.lin2img(sdf_values).squeeze().cpu().numpy()
634 ax = axs if not isinstance(axs, np.ndarray) else axs[idx]
635 make_contour_plot(sdf_values, ax=ax)
636
637 writer.add_figure(prefix + 'xy_sdf_slice', fig, global_step=total_steps)
638 plt.close('all')
639
640

Callers 2

write_sdf_summaryFunction · 0.85

Calls 1

make_contour_plotFunction · 0.85

Tested by

no test coverage detected