OpenMotionLab/MotionGPT
118
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 