CoolFace
Apppublic

memef4rmer/edit_anything

sourceHugging Faceccupdated 3y agoView on Hugging Face
0likes
vis.py110 linesDownload Raw Back to voxnerf
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