Team Ai
Apppublic

OpenMotionLab/MotionGPT

sourceHugging Facemitupdated 1y agoView on Hugging Face
118likes
visualize.py748 linesDownload Raw Back to render
1from operator import mod2import os3# from cv2 import CAP_PROP_INTELPERC_DEPTH_LOW_CONFIDENCE_VALUE4import imageio5import shutil6import numpy as np7import torch8from tqdm import tqdm9 10from scipy.spatial.transform import Rotation as R11from mGPT.render.renderer import get_renderer12from mGPT.render.rendermotion import render_video13# from mld.utils.img_utils import convert_img14# from mld.utils.uicap_utils import output_pkl15 16 17def parsename(path):18    basebane = os.path.basename(path)19    base = os.path.splitext(basebane)[0]20    strs = base.split('_')21    key = strs[-2]22    action = strs[-1]23    return key, action24 25 26def load_anim(path, timesize=None):27    data = np.array(imageio.mimread(path, memtest=False))  #[..., :3]28    if timesize is None:29        return data30 31    # take the last frame and put shadow repeat the last frame but with a little shadow32    # lastframe = add_shadow(data[-1])33    # alldata = np.tile(lastframe, (timesize, 1, 1, 1))34    alldata = data35 36    # debug fix mat dim37    if len(data.shape) == 3 and len(alldata.shape) == 4:38        data = data[:, None, :, :]39 40    # copy the first frames41    lenanim = data.shape[0]42    alldata[:lenanim] = data[:lenanim]43    return alldata44 45 46def plot_3d_motion_dico(x):47    motion, length, save_path, params, kargs = x48    plot_3d_motion(motion, length, save_path, params, **kargs)49 50 51def plot_3d_motion(motion,52                   length,53                   save_path,54                   params,55                   title="",56                   interval=50,57                   pred_cam=None,58                   imgs=None,59                   bbox=None,60                   side=None):61    # render smpl62    # [nframes, nVs, 3]63    if motion.shape[1] == 6890:64        # width = 25065        # height = 25066        width = 60067        height = 60068        if pred_cam is None:69            # cam=(0.75, 0.75, 0, 0.1)70            cam = (0.8, 0.8, 0, 0.1)71            # cam=(0.9, 0.9, 0, 0.1)72        else:73            assert bbox is not None74            assert imgs is not None75 76            # Tmp visulize77            # weak perspective camera parameters in cropped image space (s,tx,ty)78            # to79            # weak perspective camera parameters in original image space (sx,sy,tx,ty)80            cam = np.concatenate(81                (pred_cam[:, [0]], pred_cam[:, [0]], pred_cam[:, 1:3]), axis=1)82 83            # ToDo convert to original cam84            # load original img?85            # calculate cam after padding???86            #87            # cam = convert_crop_cam_to_orig_img(88            #     cam=pred_cam,89            #     bbox=bbox,90            #     img_width=width,91            #     img_height=height92            # )93        cam_pose = np.eye(4)94        cam_pose[0:3, 0:3] = R.from_euler('x', -90, degrees=True).as_matrix()95        cam_pose[0:3, 3] = [0, 0, 0]96        if side:97            rz = np.eye(4)98            rz[0:3, 0:3] = R.from_euler('z', -90, degrees=True).as_matrix()99            cam_pose = np.matmul(rz, cam_pose)100 101        # # reshape input imgs102        # if imgs is not None:103        #     imgs = convert_img(imgs.unsqueeze(0), height)[:,0]104        backgrounds = imgs if imgs is not None else np.ones(105            (height, width, 3)) * 255106        renderer = get_renderer(width, height, cam_pose)107 108        # [nframes, nVs, 3]109        meshes = motion110        key, action = parsename(save_path)111        render_video(meshes,112                     key,113                     action,114                     renderer,115                     save_path,116                     backgrounds,117                     cam_pose,118                     cams=cam)119        return120 121 122def stack_images(real, real_gens, gen, real_imgs=None):123    # change to 3 channel124    # print(real.shape)125    # print(real_gens.shape)126    # print(real_gens.shape)127    # real = real[:3]128    # real_gens = real_gens[:3]129    # gen = gen[:3]130 131    nleft_cols = len(real_gens) + 1132    print("Stacking frames..")133    allframes = np.concatenate(134        (real[:, None, ...], *[x[:, None, ...] for x in real_gens], gen), 1)135    nframes, nspa, nats, h, w, pix = allframes.shape136 137    blackborder = np.zeros((w // 30, h * nats, pix), dtype=allframes.dtype)138    # blackborder = np.ones((w//30, h*nats, pix), dtype=allframes.dtype)*255139    frames = []140    for frame_idx in tqdm(range(nframes)):141        columns = np.vstack(allframes[frame_idx].transpose(1, 2, 3, 4,142                                                           0)).transpose(143                                                               3, 1, 0, 2)144        frame = np.concatenate(145            (*columns[0:nleft_cols], blackborder, *columns[nleft_cols:]),146            0).transpose(1, 0, 2)147 148        frames.append(frame)149 150    if real_imgs is not None:151        resize_imgs = convert_img(real_imgs, h)[:nframes, ...]152 153        for i in range(len(frames)):154            imgs = np.vstack(resize_imgs[i, ...])155            imgs4 = np.ones(156                (imgs.shape[0], imgs.shape[1], 4), dtype=np.uint8) * 255157            imgs4[:, :, :3] = imgs158            #imgs = torch2numpy(imgs)159            frames[i] = np.concatenate((imgs4, frames[i]), 1)160    return np.stack(frames)161 162 163def stack_images_gen(gen, real_imgs=None):164    print("Stacking frames..")165    allframes = gen166    nframes, nspa, nats, h, w, pix = allframes.shape167    blackborder = np.zeros((w * nspa, h // 30, pix), dtype=allframes.dtype)168    blackborder = blackborder[None, ...].repeat(nats,169                                                axis=0).transpose(0, 2, 1, 3)170 171    frames = []172    for frame_idx in tqdm(range(nframes)):173        rows = np.vstack(allframes[frame_idx].transpose(0, 3, 2, 4,174                                                        1)).transpose(175                                                            3, 1, 0, 2)176        rows = np.concatenate((rows, blackborder), 1)177        frame = np.concatenate(rows, 0)178        frames.append(frame)179 180    if real_imgs is not None:181        # ToDo Add images182        resize_imgs = convert_img(real_imgs, h)[:nframes, ...]183        for i in range(len(frames)):184            imgs = np.vstack(resize_imgs[i, ...])185            #imgs = torch2numpy(imgs)186            frames[i] = np.concatenate((imgs, frames[i]), 1)187    return np.stack(frames)188 189 190def generate_by_video(visualization, reconstructions, generation,191                      label_to_action_name, params, nats, nspa, tmp_path):192    # shape : (17, 3, 4, 480, 640, 3)193    # (nframes, row, column, h, w, 3)194    fps = params["fps"]195 196    params = params.copy()197 198    gen_only = False199    if visualization is None:200        gen_only = True201        outputkey = "output_vertices"202        params["pose_rep"] = "vertices"203    elif "output_vertices" in visualization:204        outputkey = "output_vertices"205        params["pose_rep"] = "vertices"206    elif "output_xyz" in visualization:207        outputkey = "output_xyz"208        params["pose_rep"] = "xyz"209    else:210        outputkey = "poses"211 212    keep = [outputkey, 'lengths', "y"]213    gener = {key: generation[key].data.cpu().numpy() for key in keep}214    if not gen_only:215        visu = {key: visualization[key].data.cpu().numpy() for key in keep}216        recons = {}217        # visualize regressor results218        if 'vertices_hat' in reconstructions['ntf']:219            recons['regressor'] = {220                'output_vertices':221                reconstructions['ntf']['vertices_hat'].data.cpu().numpy(),222                'lengths':223                reconstructions['ntf']['lengths'].data.cpu().numpy(),224                'y':225                reconstructions['ntf']['y'].data.cpu().numpy()226            }227 228            recons['regressor_side'] = {229                'output_vertices':230                reconstructions['ntf']['vertices_hat'].data.cpu().numpy(),231                'lengths':232                reconstructions['ntf']['lengths'].data.cpu().numpy(),233                'y':234                reconstructions['ntf']['y'].data.cpu().numpy(),235                'side':236                True237            }238            # ToDo rendering overlap results239            # recons['overlap'] = {'output_vertices':reconstructions['ntf']['vertices_hat'].data.cpu().numpy(),240            #                        'lengths':reconstructions['ntf']['lengths'].data.cpu().numpy(),241            #                        'y':reconstructions['ntf']['y'].data.cpu().numpy(),242            #                        'imgs':reconstructions['ntf']['imgs'],243            #                        'bbox':reconstructions['ntf']['bbox'].data.cpu().numpy(),244            #                        'cam':reconstructions['ntf']['preds'][0]['cam'].data.cpu().numpy()}245        for mode, reconstruction in reconstructions.items():246            recons[mode] = {247                key: reconstruction[key].data.cpu().numpy()248                for key in keep249            }250            recons[mode + '_side'] = {251                key: reconstruction[key].data.cpu().numpy()252                for key in keep253            }254            recons[mode + '_side']['side'] = True255 256    # lenmax = max(gener['lengths'].max(), visu['lengths'].max())257    # timesize = lenmax + 5 longer visulization258    lenmax = gener['lengths'].max()259    timesize = lenmax260 261    import multiprocessing262 263    def pool_job_with_desc(pool, iterator, desc, max_, save_path_format, isij):264        with tqdm(total=max_, desc=desc.format("Render")) as pbar:265            for data in iterator:266                plot_3d_motion_dico(data)267            # for _ in pool.imap_unordered(plot_3d_motion_dico, iterator):268            #     pbar.update()269        if isij:270            array = np.stack([[271                load_anim(save_path_format.format(i, j), timesize)272                for j in range(nats)273            ] for i in tqdm(range(nspa), desc=desc.format("Load"))])274            return array.transpose(2, 0, 1, 3, 4, 5)275        else:276            array = np.stack([277                load_anim(save_path_format.format(i), timesize)278                for i in tqdm(range(nats), desc=desc.format("Load"))279            ])280            return array.transpose(1, 0, 2, 3, 4)281 282    pool = None283    # if True:284    with multiprocessing.Pool() as pool:285        # Generated samples286        save_path_format = os.path.join(tmp_path, "gen_{}_{}.gif")287        iterator = ((gener[outputkey][i, j], gener['lengths'][i, j],288                     save_path_format.format(i, j), params, {289                         "title":290                         f"gen: {label_to_action_name(gener['y'][i, j])}",291                         "interval": 1000 / fps292                     }) for j in range(nats) for i in range(nspa))293        gener["frames"] = pool_job_with_desc(pool, iterator,294                                             "{} the generated samples",295                                             nats * nspa, save_path_format,296                                             True)297        if not gen_only:298            # Real samples299            save_path_format = os.path.join(tmp_path, "real_{}.gif")300            iterator = ((visu[outputkey][i], visu['lengths'][i],301                         save_path_format.format(i), params, {302                             "title":303                             f"real: {label_to_action_name(visu['y'][i])}",304                             "interval": 1000 / fps305                         }) for i in range(nats))306            visu["frames"] = pool_job_with_desc(pool, iterator,307                                                "{} the real samples", nats,308                                                save_path_format, False)309            for mode, recon in recons.items():310                # Reconstructed samples311                save_path_format = os.path.join(312                    tmp_path, f"reconstructed_{mode}_" + "{}.gif")313                if mode == 'overlap':314                    iterator = ((315                        recon[outputkey][i], recon['lengths'][i],316                        save_path_format.format(i), params, {317                            "title":318                            f"recons: {label_to_action_name(recon['y'][i])}",319                            "interval": 1000 / fps,320                            "pred_cam": recon['cam'][i],321                            "imgs": recon['imgs'][i],322                            "bbox": recon['bbox'][i]323                        }) for i in range(nats))324                else:325                    side = True if 'side' in recon.keys() else False326                    iterator = ((327                        recon[outputkey][i], recon['lengths'][i],328                        save_path_format.format(i), params, {329                            "title":330                            f"recons: {label_to_action_name(recon['y'][i])}",331                            "interval": 1000 / fps,332                            "side": side333                        }) for i in range(nats))334                recon["frames"] = pool_job_with_desc(335                    pool, iterator, "{} the reconstructed samples", nats,336                    save_path_format, False)337    # vis img in visu338    if not gen_only:339        input_imgs = visualization["imgs"] if visualization[340            "imgs"] is not None else None341        vis = visu["frames"] if not gen_only else None342        rec = [recon["frames"]343               for recon in recons.values()] if not gen_only else None344        gen = gener["frames"]345        frames = stack_images(vis, rec, gen, input_imgs)346    else:347        gen = gener["frames"]348        frames = stack_images_gen(gen)349    return frames350 351 352def viz_epoch(model,353              dataset,354              epoch,355              params,356              folder,357              module=None,358              writer=None,359              exps=''):360    """ Generate & viz samples """361    module = model if module is None else module362 363    # visualize with joints3D364    model.outputxyz = True365 366    print(f"Visualization of the epoch {epoch}")367 368    noise_same_action = params["noise_same_action"]369    noise_diff_action = params["noise_diff_action"]370    duration_mode = params["duration_mode"]371    reconstruction_mode = params["reconstruction_mode"]372    decoder_test = params["decoder_test"]373 374    fact = params["fact_latent"]375    figname = params["figname"].format(epoch)376 377    nspa = params["num_samples_per_action"]378    nats = params["num_actions_to_sample"]379 380    num_classes = params["num_classes"]381    # nats = min(num_classes, nats)382 383    # define some classes384    classes = torch.randperm(num_classes)[:nats]385    # duplicate same classes when sampling too much386    if nats > num_classes:387        classes = classes.expand(nats)388 389    meandurations = torch.from_numpy(390        np.array([391            round(dataset.get_mean_length_label(cl.item())) for cl in classes392        ]))393 394    if duration_mode == "interpolate" or decoder_test == "diffduration":395        points, step = np.linspace(-nspa, nspa, nspa, retstep=True)396        # points = np.round(10*points/step).astype(int)397        points = np.array([5, 10, 16, 30, 60, 80]).astype(int)398        # gendurations = meandurations.repeat((nspa, 1)) + points[:, None]399        gendurations = torch.from_numpy(points[:, None]).expand(400            (nspa, 1)).repeat((1, nats))401    else:402        gendurations = meandurations.repeat((nspa, 1))403    print("Duration time: ")404    print(gendurations[:, 0])405 406    # extract the real samples407    # real_samples, real_theta, mask_real, real_lengths, imgs, paths408    batch = dataset.get_label_sample_batch(classes.numpy())409 410    # ToDo411    # clean these data412    # Visualizaion of real samples413    visualization = {414        "x": batch['x'].to(model.device),415        "y": classes.to(model.device),416        "mask": batch['mask'].to(model.device),417        'lengths': batch['lengths'].to(model.device),418        "output": batch['x'].to(model.device),419        "theta":420        batch['theta'].to(model.device) if 'theta' in batch.keys() else None,421        "imgs":422        batch['imgs'].to(model.device) if 'imgs' in batch.keys() else None,423        "paths": batch['paths'] if 'paths' in batch.keys() else None,424    }425 426    # Visualizaion of real samples427    if reconstruction_mode == "both":428        reconstructions = {429            "tf": {430                "x":431                batch['x'].to(model.device),432                "y":433                classes.to(model.device),434                'lengths':435                batch['lengths'].to(model.device),436                "mask":437                batch['mask'].to(model.device),438                "teacher_force":439                True,440                "theta":441                batch['theta'].to(model.device)442                if 'theta' in batch.keys() else None443            },444            "ntf": {445                "x":446                batch['x'].to(model.device),447                "y":448                classes.to(model.device),449                'lengths':450                batch['lengths'].to(model.device),451                "mask":452                batch['mask'].to(model.device),453                "theta":454                batch['theta'].to(model.device)455                if 'theta' in batch.keys() else None456            }457        }458    else:459        reconstructions = {460            reconstruction_mode: {461                "x":462                batch['x'].to(model.device),463                "y":464                classes.to(model.device),465                'lengths':466                batch['lengths'].to(model.device),467                "mask":468                batch['mask'].to(model.device),469                "teacher_force":470                reconstruction_mode == "tf",471                "imgs":472                batch['imgs'].to(model.device)473                if 'imgs' in batch.keys() else None,474                "theta":475                batch['theta'].to(model.device)476                if 'theta' in batch.keys() else None,477                "bbox":478                batch['bbox'] if 'bbox' in batch.keys() else None479            }480        }481    print("Computing the samples poses..")482 483    # generate the repr (joints3D/pose etc)484    model.eval()485    with torch.no_grad():486        # Reconstruction of the real data487        for mode in reconstructions:488            # update reconstruction dicts489            reconstructions[mode] = model(reconstructions[mode])490        reconstruction = reconstructions[list(reconstructions.keys())[0]]491 492        if decoder_test == "gt":493            # Generate the new data494            gt_input = {495                "x": batch['x'].repeat(nspa, 1, 1, 1).to(model.device),496                "y": classes.repeat(nspa).to(model.device),497                "mask": batch['mask'].repeat(nspa, 1).to(model.device),498                'lengths': batch['lengths'].repeat(nspa).to(model.device)499            }500            generation = model(gt_input)501        if decoder_test == "new":502            # Generate the new data503            generation = module.generate(gendurations,504                                         classes=classes,505                                         nspa=nspa,506                                         noise_same_action=noise_same_action,507                                         noise_diff_action=noise_diff_action,508                                         fact=fact)509        elif decoder_test == "diffaction":510            assert nats == nspa511            # keep the same noise for each "sample"512            z = reconstruction["z"].repeat((nspa, 1))513            mask = reconstruction["mask"].repeat((nspa, 1))514            lengths = reconstruction['lengths'].repeat(nspa)515            # but use other labels516            y = classes.repeat_interleave(nspa).to(model.device)517            generation = {"z": z, "y": y, "mask": mask, 'lengths': lengths}518            model.decoder(generation)519 520        elif decoder_test == "diffduration":521            z = reconstruction["z"].repeat((nspa, 1))522            lengths = gendurations.reshape(-1).to(model.device)523            mask = model.lengths_to_mask(lengths)524            y = classes.repeat(nspa).to(model.device)525            generation = {"z": z, "y": y, "mask": mask, 'lengths': lengths}526            model.decoder(generation)527 528        elif decoder_test == "interpolate_action":529            assert nats == nspa530            # same noise for each sample531            z_diff_action = torch.randn(1,532                                        model.latent_dim,533                                        device=model.device).repeat(nats, 1)534            z = z_diff_action.repeat((nspa, 1))535 536            # but use combination of labels and labels below537            y = F.one_hot(classes.to(model.device),538                          model.num_classes).to(model.device)539            y_below = F.one_hot(torch.cat((classes[1:], classes[0:1])),540                                model.num_classes).to(model.device)541            convex_factors = torch.linspace(0, 1, nspa, device=model.device)542            y_mixed = torch.einsum("nk,m->mnk", y, 1-convex_factors) + \543                torch.einsum("nk,m->mnk", y_below, convex_factors)544            y_mixed = y_mixed.reshape(nspa * nats, y_mixed.shape[-1])545 546            durations = gendurations[0].to(model.device)547            durations_below = torch.cat((durations[1:], durations[0:1]))548 549            gendurations = torch.einsum("l,k->kl", durations, 1-convex_factors) + \550                torch.einsum("l,k->kl", durations_below, convex_factors)551            gendurations = gendurations.to(dtype=durations.dtype)552 553            lengths = gendurations.to(model.device).reshape(z.shape[0])554            mask = model.lengths_to_mask(lengths)555 556            generation = {557                "z": z,558                "y": y_mixed,559                "mask": mask,560                'lengths': lengths561            }562            generation = model.decoder(generation)563 564        visualization = module.prepare(visualization)565        visualization["output_xyz"] = visualization["x_xyz"]566        visualization["output_vertices"] = visualization["x_vertices"]567        # Get xyz for the real ones568        # visualization["output_xyz"] = module.rot2xyz(visualization["output"], visualization["mask"], jointstype="smpl")569        # # Get smpl vertices for the real ones570        # if module.cvae.pose_rep != "xyz":571        #     visualization["output_vertices"] = module.rot2xyz(visualization["output"], visualization["mask"], jointstype="vertices")572 573    for key, val in generation.items():574        if len(generation[key].shape) == 1:575            generation[key] = val.reshape(nspa, nats)576        else:577            generation[key] = val.reshape(nspa, nats, *val.shape[1:])578 579    finalpath = os.path.join(folder, figname + exps + ".gif")580    tmp_path = os.path.join(folder, f"subfigures_{figname}")581    os.makedirs(tmp_path, exist_ok=True)582 583    print("Generate the videos..")584    frames = generate_by_video(visualization, reconstructions, generation,585                               dataset.label_to_action_name, params, nats,586                               nspa, tmp_path)587 588    print(f"Writing video {finalpath}")589    imageio.mimsave(finalpath.replace('gif', 'mp4'), frames, fps=params["fps"])590    shutil.rmtree(tmp_path)591 592    # output npy593    output = {594        "data_id": batch['id'],595        "paths": batch['paths'],596        "x": batch['x'].cpu().numpy(),597        "x_vertices": visualization["x_vertices"].cpu().numpy(),598        "output_vertices":599        reconstructions['ntf']["output_vertices"].cpu().numpy(),600        "gen_vertices": generation["output_vertices"].cpu().numpy()601    }602 603    outputpath = finalpath.replace('gif', 'npy')604    np.save(outputpath, output)605 606    # output pkl607    batch_recon = reconstructions["ntf"]608    outputpath = finalpath.replace('gif', 'pkl')609    # output_pkl([batch_recon], outputpath)610 611    if writer is not None:612        writer.add_video(f"Video/Epoch {epoch}",613                         frames.transpose(0, 3, 1, 2)[None],614                         epoch,615                         fps=params["fps"])616    return finalpath617 618 619def viz_dataset(dataset, params, folder):620    """ Generate & viz samples """621    print("Visualization of the dataset")622 623    nspa = params["num_samples_per_action"]624    nats = params["num_actions_to_sample"]625 626    num_classes = params["num_classes"]627 628    figname = "{}_{}_numframes_{}_sampling_{}_step_{}".format(629        params["dataset"], params["pose_rep"], params["num_frames"],630        params["sampling"], params["sampling_step"])631 632    # define some classes633    classes = torch.randperm(num_classes)[:nats]634 635    allclasses = classes.repeat(nspa, 1).reshape(nspa * nats)636    # extract the real samples637    real_samples, mask_real, real_lengths = dataset.get_label_sample_batch(638        allclasses.numpy())639    # to visualize directly640 641    # Visualizaion of real samples642    visualization = {643        "x": real_samples,644        "y": allclasses,645        "mask": mask_real,646        'lengths': real_lengths,647        "output": real_samples648    }649 650    from mGPT.models.rotation2xyz import Rotation2xyz651 652    device = params["device"]653    rot2xyz = Rotation2xyz(device=device)654 655    rot2xyz_params = {656        "pose_rep": params["pose_rep"],657        "glob_rot": params["glob_rot"],658        "glob": params["glob"],659        "jointstype": params["jointstype"],660        "translation": params["translation"]661    }662 663    output = visualization["output"]664    visualization["output_xyz"] = rot2xyz(output.to(device),665                                          visualization["mask"].to(device),666                                          **rot2xyz_params)667 668    for key, val in visualization.items():669        if len(visualization[key].shape) == 1:670            visualization[key] = val.reshape(nspa, nats)671        else:672            visualization[key] = val.reshape(nspa, nats, *val.shape[1:])673 674    finalpath = os.path.join(folder, figname + ".gif")675    tmp_path = os.path.join(folder, f"subfigures_{figname}")676    os.makedirs(tmp_path, exist_ok=True)677 678    print("Generate the videos..")679    frames = generate_by_video_sequences(visualization,680                                         dataset.label_to_action_name, params,681                                         nats, nspa, tmp_path)682 683    print(f"Writing video {finalpath}..")684    imageio.mimsave(finalpath, frames, fps=params["fps"])685 686 687def generate_by_video_sequences(visualization, label_to_action_name, params,688                                nats, nspa, tmp_path):689    # shape : (17, 3, 4, 480, 640, 3)690    # (nframes, row, column, h, w, 3)691    fps = params["fps"]692    if "output_vetices" in visualization:693        outputkey = "output_vetices"694        params["pose_rep"] = "vertices"695    elif "output_xyz" in visualization:696        outputkey = "output_xyz"697        params["pose_rep"] = "xyz"698    else:699        outputkey = "poses"700 701    keep = [outputkey, 'lengths', "y"]702    visu = {key: visualization[key].data.cpu().numpy() for key in keep}703    lenmax = visu['lengths'].max()704 705    timesize = lenmax + 5706 707    # import multiprocessing708 709    def pool_job_with_desc(pool, iterator, desc, max_, save_path_format):710        for data in iterator:711            plot_3d_motion_dico(data)712        # with tqdm(total=max_, desc=desc.format("Render")) as pbar:713        #     for _ in pool.imap_unordered(plot_3d_motion_dico, iterator):714        #         pbar.update()715        array = np.stack([[716            load_anim(save_path_format.format(i, j), timesize)717            for j in range(nats)718        ] for i in tqdm(range(nspa), desc=desc.format("Load"))])719        return array.transpose(2, 0, 1, 3, 4, 5)720 721    pool = None722    # with multiprocessing.Pool() as pool:723    # Real samples724    save_path_format = os.path.join(tmp_path, "real_{}_{}.gif")725    iterator = ((visu[outputkey][i, j], visu['lengths'][i, j],726                 save_path_format.format(i, j), params, {727                     "title": f"real: {label_to_action_name(visu['y'][i, j])}",728                     "interval": 1000 / fps729                 }) for j in range(nats) for i in range(nspa))730    visu["frames"] = pool_job_with_desc(pool, iterator, "{} the real samples",731                                        nats, save_path_format)732    frames = stack_images_sequence(visu["frames"])733    return frames734 735 736def stack_images_sequence(visu):737    print("Stacking frames..")738    allframes = visu739    nframes, nspa, nats, h, w, pix = allframes.shape740    frames = []741    for frame_idx in tqdm(range(nframes)):742        columns = np.vstack(allframes[frame_idx].transpose(1, 2, 3, 4,743                                                           0)).transpose(744                                                               3, 1, 0, 2)745        frame = np.concatenate(columns).transpose(1, 0, 2)746        frames.append(frame)747    return np.stack(frames)748