write slices of sdf in each plane
(model, writer, total_steps, prefix='train_', is_multi=False)
| 581 | |
| 582 | |
| 583 | def 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 |
no test coverage detected