memef4rmer/edit_anything
0
1from pathlib import Path2import numpy as np3import matplotlib.pyplot as plt4from mpl_toolkits.axes_grid1 import ImageGrid5from matplotlib.colors import Normalize, LogNorm6import torch7from torchvision.utils import make_grid8from einops import rearrange9from .data import blend_rgba10 11import imageio12 13from my.utils.plot import mpl_fig_to_buffer14from my.utils.event import read_stats15 16 17def vis(ref_img, pred_img, pred_depth, *, msg="", return_buffer=False):18 # plt the 2 images side by side and compare19 fig = plt.figure(figsize=(15, 6))20 grid = ImageGrid(21 fig, 111, nrows_ncols=(1, 3),22 cbar_location="right", cbar_mode="single",23 )24 25 grid[0].imshow(ref_img)26 grid[0].set_title("gt")27 28 grid[1].imshow(pred_img)29 grid[1].set_title(f"rendering {msg}")30 31 h = grid[2].imshow(pred_depth, norm=LogNorm(vmin=2, vmax=10), cmap="Spectral")32 grid[2].set_title("expected depth")33 plt.colorbar(h, cax=grid.cbar_axes[0])34 plt.tight_layout()35 36 if return_buffer:37 plot = mpl_fig_to_buffer(fig)38 return plot39 else:40 plt.show()41 42 43def _bad_vis(pred_img, pred_depth, *, return_buffer=False):44 """emergency function for one-off use"""45 fig, grid = plt.subplots(1, 2, squeeze=True, figsize=(10, 6))46 47 grid[0].imshow(pred_img)48 grid[0].set_title("rendering")49 50 h = grid[1].imshow(pred_depth, norm=LogNorm(vmin=0.5, vmax=10), cmap="Spectral")51 grid[1].set_title("expected depth")52 # plt.colorbar(h, cax=grid.cbar_axes[0])53 plt.tight_layout()54 55 if return_buffer:56 plot = mpl_fig_to_buffer(fig)57 return plot58 else:59 plt.show()60 61 62colormap = plt.get_cmap('Spectral')63 64 65def bad_vis(pred_img, pred_depth, final_H=512):66 # pred_img = pred_img.cpu()67 depth = pred_depth.cpu().numpy()68 del pred_depth69 70 depth = np.log(1. + depth + 1e-12)71 depth = depth / np.log(1+10.)72 # depth = 1 - depth73 depth = colormap(depth)74 depth = blend_rgba(depth)75 depth = rearrange(depth, "h w c -> 1 c h w", c=3)76 depth = torch.from_numpy(depth)77 78 depth = torch.nn.functional.interpolate(79 depth, (final_H, final_H), mode='bilinear', antialias=True80 )81 pred_img = torch.nn.functional.interpolate(82 pred_img, (final_H, final_H), mode='bilinear', antialias=True83 )84 pred_img = (pred_img + 1) / 285 pred_img = pred_img.clamp(0, 1).cpu()86 stacked = torch.cat([pred_img, depth], dim=0)87 pane = make_grid(stacked, nrow=2)88 pane = rearrange(pane, "c h w -> h w c")89 pane = (pane * 255.).clamp(0, 255)90 pane = pane.to(torch.uint8)91 pane = pane.numpy()92 # plt.imshow(pane)93 # plt.show()94 return pane95 96 97def export_movie(seqs, fname, fps=30):98 fname = Path(fname)99 if fname.suffix == "":100 fname = fname.with_suffix(".mp4")101 writer = imageio.get_writer(fname, fps=fps)102 for img in seqs:103 writer.append_data(img)104 writer.close()105 106 107def stitch_vis(save_fn, img_fnames, fps=10):108 figs = [imageio.imread(fn) for fn in img_fnames]109 export_movie(figs, save_fn, fps)110 