Spaces:
Running on Zero
Running on Zero
| import os | |
| import sys | |
| import copy | |
| from pprint import pformat | |
| from functools import partial | |
| sys.path.append(".") | |
| DEVICE_TYPE = os.environ.get("DEVICE_TYPE", "gpu") | |
| import torch | |
| if not torch.cuda.is_available() or DEVICE_TYPE == 'npu': | |
| USE_NPU = True | |
| os.environ['DEVICE_TYPE'] = "npu" | |
| DEVICE_TYPE = "npu" | |
| print("Enable NPU!") | |
| try: | |
| # just before torch_npu, let xformers know there is no gpu | |
| import xformers | |
| import xformers.ops | |
| except Exception as e: | |
| print(f"Got {e} during import xformers!") | |
| import torch_npu | |
| from torch_npu.contrib import transfer_to_npu | |
| else: | |
| USE_NPU = False | |
| import magicdrivedit.utils.module_contrib | |
| import colossalai | |
| import torch.distributed as dist | |
| import torchvision.transforms as TF | |
| from einops import rearrange, repeat | |
| from colossalai.cluster import DistCoordinator, ProcessGroupMesh | |
| from mmengine.runner import set_random_seed | |
| from tqdm import tqdm | |
| from mmcv.parallel import DataContainer | |
| from magicdrivedit.acceleration.communications import gather_tensors, serialize_state, deserialize_state | |
| from magicdrivedit.acceleration.parallel_states import ( | |
| set_sequence_parallel_group, | |
| get_sequence_parallel_group, | |
| set_data_parallel_group, | |
| get_data_parallel_group, | |
| ) | |
| from magicdrivedit.datasets import save_sample | |
| from magicdrivedit.datasets.dataloader import prepare_dataloader | |
| from magicdrivedit.datasets.dataloader import prepare_dataloader | |
| from magicdrivedit.registry import DATASETS, MODELS, SCHEDULERS, build_module | |
| from magicdrivedit.utils.config_utils import parse_configs, define_experiment_workspace, save_training_config, merge_dataset_cfg, mmengine_conf_get, mmengine_conf_set | |
| from magicdrivedit.utils.inference_utils import ( | |
| concat_6_views_pt, | |
| add_null_condition, | |
| enable_offload, | |
| ) | |
| from magicdrivedit.utils.misc import ( | |
| reset_logger, | |
| is_distributed, | |
| to_torch_dtype, | |
| collate_bboxes_to_maxlen, | |
| move_to, | |
| add_box_latent, | |
| ) | |
| from magicdrivedit.utils.train_utils import sp_vae | |
| VIEW_ORDER = [ | |
| "CAM_FRONT_LEFT", | |
| "CAM_FRONT", | |
| "CAM_FRONT_RIGHT", | |
| "CAM_BACK_RIGHT", | |
| "CAM_BACK", | |
| "CAM_BACK_LEFT", | |
| ] | |
| def make_file_dirs(path): | |
| os.makedirs(os.path.dirname(path), exist_ok=True) | |
| def set_omegaconf_key_value(cfg, key, value): | |
| p, m = key.rsplit(".", 1) | |
| node = cfg | |
| for pk in p.split("."): | |
| node = getattr(node, pk) | |
| node[m] = value | |
| def main(): | |
| torch.set_grad_enabled(False) | |
| # ====================================================== | |
| # configs & runtime variables | |
| # ====================================================== | |
| # == parse configs == | |
| cfg = parse_configs(training=False) | |
| if cfg.get("vsdebug", False): | |
| import debugpy | |
| debugpy.listen(5678) | |
| print("Waiting for debugger attach") | |
| debugpy.wait_for_client() | |
| print('Attached, continue...') | |
| # == dataset config == | |
| if cfg.num_frames is None: | |
| num_data_cfgs = len(cfg.data_cfg_names) | |
| datasets = [] | |
| val_datasets = [] | |
| for (res, data_cfg_name), overrides in zip( | |
| cfg.data_cfg_names, cfg.get("dataset_cfg_overrides", [[]] * num_data_cfgs)): | |
| dataset, val_dataset = merge_dataset_cfg(cfg, data_cfg_name, overrides) | |
| datasets.append((res, dataset)) | |
| val_datasets.append((res, val_dataset)) | |
| dataset = {"type": "NuScenesMultiResDataset", "cfg": datasets} | |
| val_dataset = {"type": "NuScenesMultiResDataset", "cfg": val_datasets} | |
| else: | |
| dataset, val_dataset = merge_dataset_cfg( | |
| cfg, cfg.data_cfg_name, cfg.get("dataset_cfg_overrides", []), | |
| cfg.num_frames) | |
| if cfg.get("use_train", False): | |
| cfg.dataset = dataset | |
| tag = cfg.get("tag", "") | |
| cfg.tag = "train" if tag == "" else f"{tag}_train" | |
| else: | |
| cfg.dataset = val_dataset | |
| # set img_collate_param | |
| if hasattr(cfg.dataset, "img_collate_param"): | |
| cfg.dataset.img_collate_param.is_train = False # Important! | |
| else: | |
| for d in cfg.dataset.cfg: | |
| d[1].img_collate_param.is_train = False # Important! | |
| cfg.batch_size = 1 | |
| # for lower cpu memory in dataloading | |
| cfg.ignore_ori_imgs = cfg.get("ignore_ori_imgs", False) | |
| if cfg.ignore_ori_imgs: | |
| cfg.dataset.drop_ori_imgs = True | |
| # post transformation | |
| cfg.use_back_trans = cfg.get("use_back_trans", True) | |
| cfg.save_mode = cfg.get("save_mode", "single-view") | |
| assert cfg.save_mode in ["single-view", "all-in-one", "image_filename"] | |
| cfg.use_map0 = cfg.get("use_map0", False) | |
| # == device and dtype == | |
| device = "cuda" if torch.cuda.is_available() else "cpu" | |
| cfg_dtype = cfg.get("dtype", "bf16") | |
| assert cfg_dtype in ["fp16", "bf16", "fp32"], f"Unknown mixed precision {cfg_dtype}" | |
| dtype = to_torch_dtype(cfg.get("dtype", "bf16")) | |
| torch.backends.cuda.matmul.allow_tf32 = True | |
| torch.backends.cudnn.allow_tf32 = True | |
| if USE_NPU: # disable some kernels | |
| if mmengine_conf_get(cfg, "text_encoder.shardformer", None): | |
| mmengine_conf_set(cfg, "text_encoder.shardformer", False) | |
| if mmengine_conf_get(cfg, "model.bbox_embedder_param.enable_xformers", None): | |
| mmengine_conf_set(cfg, "model.bbox_embedder_param.enable_xformers", False) | |
| if mmengine_conf_get(cfg, "model.frame_emb_param.enable_xformers", None): | |
| mmengine_conf_set(cfg, "model.frame_emb_param.enable_xformers", False) | |
| # == init distributed env == | |
| cfg.sp_size = cfg.get("sp_size", 1) | |
| if is_distributed(): | |
| colossalai.launch_from_torch({}) | |
| else: | |
| dist.init_process_group( | |
| backend="nccl", world_size=1, rank=0, | |
| init_method="tcp://localhost:12355") | |
| coordinator = DistCoordinator() | |
| if cfg.sp_size > 1: | |
| DP_AXIS, SP_AXIS = 0, 1 | |
| dp_size = dist.get_world_size() // cfg.sp_size | |
| pg_mesh = ProcessGroupMesh(dp_size, cfg.sp_size) | |
| dp_group = pg_mesh.get_group_along_axis(DP_AXIS) | |
| sp_group = pg_mesh.get_group_along_axis(SP_AXIS) | |
| set_sequence_parallel_group(sp_group) | |
| else: | |
| # TODO: sequence_parallel_group unset! | |
| dp_group = dist.group.WORLD | |
| set_data_parallel_group(dp_group) | |
| set_random_seed(seed=cfg.get("seed", 1024)) | |
| # == init exp_dir == | |
| cfg.outputs = cfg.get("outputs", "outputs/test") | |
| exp_name, exp_dir = define_experiment_workspace(cfg, use_date=True) | |
| cfg.save_dir = os.path.join(exp_dir, "generation") | |
| coordinator.block_all() | |
| if coordinator.is_master(): | |
| os.makedirs(exp_dir, exist_ok=True) | |
| save_training_config(cfg.to_dict(), exp_dir) | |
| coordinator.block_all() | |
| # == init logger == | |
| logger = reset_logger(exp_dir) | |
| logger.info("Inference configuration:\n %s", pformat(cfg.to_dict())) | |
| verbose = cfg.get("verbose", 1) | |
| # ====================================================== | |
| # 2. build dataset and dataloader | |
| # ====================================================== | |
| if cfg.get("val", None): | |
| validation_index = cfg.val.validation_index | |
| if validation_index == "all": | |
| raise NotImplementedError() | |
| cfg.num_sample = cfg.val.get("num_sample", 1) | |
| cfg.scheduler = cfg.val.get("scheduler", cfg.scheduler) | |
| else: | |
| validation_index = cfg.get("validation_index", "all") | |
| # == build dataset == | |
| logger.info("Building dataset...") | |
| dataset = build_module(cfg.dataset, DATASETS) | |
| if validation_index == "even": | |
| idxs = list(range(0, len(dataset), 2)) | |
| dataset = torch.utils.data.Subset(dataset, idxs) | |
| elif validation_index == "odd": | |
| idxs = list(reversed(list(range(1, len(dataset), 2)))) # reversed! | |
| dataset = torch.utils.data.Subset(dataset, idxs) | |
| elif validation_index != "all": | |
| dataset = torch.utils.data.Subset(dataset, validation_index) | |
| logger.info(f"Your validation index: {validation_index}") | |
| logger.info("Dataset contains %s samples.", len(dataset)) | |
| # == build dataloader == | |
| dataloader_args = dict( | |
| dataset=dataset, | |
| batch_size=cfg.get("batch_size", 1), | |
| num_workers=cfg.get("num_workers", 4), | |
| seed=cfg.get("seed", 1024), | |
| shuffle=isinstance(validation_index, str), # changed | |
| drop_last=False, # changed | |
| pin_memory=True, | |
| process_group=get_data_parallel_group(), | |
| prefetch_factor=cfg.get("prefetch_factor", None), | |
| ) | |
| dataloader, sampler = prepare_dataloader( | |
| bucket_config=cfg.get("bucket_config", None), | |
| num_bucket_build_workers=cfg.get("num_bucket_build_workers", 1), | |
| **dataloader_args, | |
| ) | |
| num_steps_per_epoch = len(dataloader) | |
| def collate_data_container_fn(batch, *, collate_fn_map=None): | |
| return batch | |
| # add datacontainer handler | |
| torch.utils.data._utils.collate.default_collate_fn_map.update({ | |
| DataContainer: collate_data_container_fn | |
| }) | |
| # ====================================================== | |
| # build model & load weights | |
| # ====================================================== | |
| logger.info("Building models...") | |
| # == build text-encoder and vae == | |
| # NOTE: set to true/false, | |
| # https://github.com/huggingface/transformers/issues/5486 | |
| # if the program gets stuck, try set it to false | |
| os.environ['TOKENIZERS_PARALLELISM'] = "true" | |
| text_encoder = build_module(cfg.text_encoder, MODELS, device=device) | |
| vae = build_module(cfg.vae, MODELS).to(device, dtype).eval() | |
| # == prepare video size == | |
| if cfg.use_back_trans: | |
| # FIXME: we should have permuted (0, 1) here, but we did not do it. | |
| back_trans = TF.Compose([ | |
| TF.Resize(cfg.post.resize, interpolation=TF.InterpolationMode.BICUBIC), | |
| TF.Pad(cfg.post.padding), | |
| ]) | |
| cut_length = cfg.post.get("cut_length", None) | |
| else: | |
| def back_trans(x): return x | |
| cut_length = cfg.post.get("cut_length", None) | |
| logger.info(f"Using transform:\n{back_trans}\ncut_length={cut_length}") | |
| # == build diffusion model == | |
| model = ( | |
| build_module( | |
| cfg.model, | |
| MODELS, | |
| input_size=(None, None, None), | |
| in_channels=vae.out_channels, | |
| caption_channels=text_encoder.output_dim, | |
| model_max_length=text_encoder.model_max_length, | |
| enable_sequence_parallelism=cfg.sp_size > 1, | |
| ) | |
| .to(device, dtype) | |
| .eval() | |
| ) | |
| text_encoder.y_embedder = model.y_embedder # HACK: for classifier-free guidance | |
| # == build scheduler == | |
| scheduler = build_module(cfg.scheduler, SCHEDULERS) | |
| # ====================================================== | |
| # inference | |
| # ====================================================== | |
| cfg.cpu_offload = cfg.get("cpu_offload", False) | |
| if cfg.cpu_offload: | |
| text_encoder.t5.model.to("cpu") | |
| model.to("cpu") | |
| vae.to("cpu") | |
| text_encoder.t5.model, model, vae, last_hook = enable_offload( | |
| text_encoder.t5.model, model, vae, device) | |
| # == load prompts == | |
| batch_size = cfg.get("batch_size", 1) | |
| num_sample = cfg.get("num_sample", 1) | |
| save_video_dir = os.path.join(cfg.save_dir, "gen_video") | |
| save_gt_video_dir = os.path.join(cfg.save_dir, "gt_video") | |
| # == Iter over all samples == | |
| start_step = 0 | |
| total_num = 0 | |
| assert batch_size == 1 | |
| sampler.set_epoch(0) | |
| dataloader_iter = iter(dataloader) | |
| generator = torch.Generator("cpu").manual_seed(cfg.seed) | |
| bl_generator = torch.Generator("cpu").manual_seed(cfg.seed) | |
| with tqdm( | |
| enumerate(dataloader_iter, start=start_step), | |
| desc=f"Generating", | |
| disable=not coordinator.is_master() or not verbose, | |
| initial=start_step, | |
| total=num_steps_per_epoch, | |
| ) as pbar: | |
| for i, batch in pbar: | |
| this_token: str = batch['meta_data']['metas'][0][0].data['token'] | |
| if cfg.ignore_ori_imgs: | |
| B, T, NC = 1, *batch["pixel_values_shape"][0].tolist()[:2] | |
| latent_size = vae.get_latent_size( | |
| (T, *batch["pixel_values_shape"][0].tolist()[-2:])) | |
| else: | |
| B, T, NC = batch["pixel_values"].shape[:3] | |
| latent_size = vae.get_latent_size((T, *batch["pixel_values"].shape[-2:])) | |
| # == prepare batch prompts == | |
| y = batch.pop("captions")[0] # B, just take first frame | |
| maps = batch.pop("bev_map_with_aux").to(device, dtype) # B, T, C, H, W | |
| bbox = batch.pop("bboxes_3d_data") | |
| # B len list (T, NC, len, 8, 3) | |
| bbox = [bbox_i.data for bbox_i in bbox] | |
| # B, T, NC, len, 8, 3 | |
| # TODO: `bbox` may have some redundancy on `NC` dim. | |
| # NOTE: we reshape the data later! | |
| bbox = collate_bboxes_to_maxlen(bbox, device, dtype, NC, T) | |
| # B, T, NC, 3, 7 | |
| cams = batch.pop("camera_param").to(device, dtype) | |
| cams = rearrange(cams, "B T NC ... -> (B NC) T 1 ...") # BxNC, T, 1, 3, 7 | |
| rel_pos = batch.pop("frame_emb").to(device, dtype) | |
| rel_pos = repeat(rel_pos, "B T ... -> (B NC) T 1 ...", NC=NC) # BxNC, T, 1, 4, 4 | |
| # == model input format == | |
| model_args = {} | |
| model_args["maps"] = maps | |
| model_args["bbox"] = bbox | |
| model_args["cams"] = cams | |
| model_args["rel_pos"] = rel_pos | |
| model_args["fps"] = batch.pop('fps') | |
| model_args["height"] = batch.pop("height") | |
| model_args["width"] = batch.pop("width") | |
| model_args["num_frames"] = batch.pop("num_frames") | |
| model_args = move_to(model_args, device=device, dtype=dtype) | |
| # no need to move these | |
| model_args["mv_order_map"] = cfg.get("mv_order_map") | |
| model_args["t_order_map"] = cfg.get("t_order_map") | |
| # == Iter over number of sampling for one prompt == | |
| save_fps = int(model_args['fps'][0]) | |
| _fpss = gather_tensors(model_args['fps'], pg=get_data_parallel_group()) | |
| _tokens = [[bytes(_t).decode("utf8") for _t in _tk] for _tk in gather_tensors( | |
| torch.ByteTensor([bytes(this_token, 'utf8')]).to(device=device))] | |
| if cfg.save_mode == "image_filename": | |
| gen_length = cut_length if cut_length is not None else T | |
| # assume bs=1! | |
| _filenames = [ | |
| deserialize_state(_meta) | |
| for _meta in gather_tensors( | |
| serialize_state( | |
| [batch['meta_data']['metas'][i][0].data['filename'] for i in range(gen_length)] | |
| ).cuda(), | |
| pg=get_data_parallel_group(), | |
| ) | |
| ] | |
| for ns in range(num_sample): | |
| z = torch.randn( | |
| len(y), vae.out_channels * NC, *latent_size, generator=generator, | |
| ).to(device=device, dtype=dtype) | |
| # == sample box == | |
| if bbox is not None: | |
| # null set values to all zeros, this should be safe | |
| bbox = add_box_latent(bbox, B, NC, T, | |
| partial(model.sample_box_latent, generator=bl_generator)) | |
| # overwrite! | |
| new_bbox = {} | |
| for k, v in bbox.items(): | |
| new_bbox[k] = rearrange(v, "B T NC ... -> (B NC) T ...") # BxNC, T, len, 3, 7 | |
| model_args["bbox"] = move_to(new_bbox, device=device, dtype=dtype) | |
| # == add null condition == | |
| # y is handled by scheduler.sample | |
| if cfg.scheduler.type == "dpm-solver" and cfg.scheduler.cfg_scale == 1.0 or ( | |
| cfg.scheduler.type in ["rflow-slice",] | |
| ): | |
| _model_args = copy.deepcopy(model_args) | |
| else: | |
| _model_args = add_null_condition( | |
| copy.deepcopy(model_args), | |
| model.camera_embedder.uncond_cam.to(device), | |
| model.frame_embedder.uncond_cam.to(device), | |
| prepend=(cfg.scheduler.type == "dpm-solver"), | |
| use_map0=cfg.get("use_map0", False), | |
| ) | |
| # == inference == | |
| samples = scheduler.sample( | |
| model, | |
| text_encoder, | |
| z=z, | |
| prompts=y, | |
| device=device, | |
| additional_args=_model_args, | |
| progress=verbose >= 1 and coordinator.is_master(), | |
| mask=None, | |
| ) | |
| samples = rearrange(samples, "B (C NC) T ... -> (B NC) C T ...", NC=NC) | |
| if cfg.sp_size > 1: | |
| samples = sp_vae( | |
| samples.to(dtype), | |
| partial(vae.decode, num_frames=_model_args["num_frames"]), | |
| get_sequence_parallel_group(), | |
| ) | |
| else: | |
| samples = vae.decode(samples.to(dtype), num_frames=_model_args["num_frames"]) | |
| samples = rearrange(samples, "(B NC) C T ... -> B NC C T ...", NC=NC) | |
| if cfg.cpu_offload: | |
| last_hook.offload() | |
| # cut to standard length | |
| samples = samples[:, :, :, slice(None, cut_length)] | |
| # gather sample from all processes | |
| coordinator.block_all() | |
| _samples = gather_tensors(samples, pg=get_data_parallel_group()) | |
| # == save samples, one-time-generation only == | |
| if coordinator.is_master(): | |
| video_clips = [] | |
| fpss = [] | |
| tokens = [] | |
| for sample, fps, token in zip(_samples, _fpss, _tokens): # list of B, NC, C, T ... | |
| video_clips += [s.cpu() for s in sample] # list of NC, C, T ... | |
| fpss += [int(_fps) for _fps in fps] | |
| tokens += [_tk for _tk in token] | |
| for idx, videos in enumerate(video_clips): # NC, C, T ... | |
| if cfg.save_mode == "single-view": | |
| for view, video in zip(VIEW_ORDER, videos): | |
| save_path = os.path.join( | |
| save_video_dir, f"{tokens[idx]}_gen{ns}", | |
| f"{tokens[idx]}_{view}") | |
| make_file_dirs(save_path) | |
| save_path = save_sample( | |
| back_trans(video), | |
| fps=save_fps if save_fps else fpss[idx], | |
| save_path=save_path, | |
| high_quality=True, | |
| verbose=verbose >= 2, | |
| with_postfix=False, | |
| ) | |
| elif cfg.save_mode == "all-in-one": | |
| video = concat_6_views_pt(videos, oneline=False) | |
| save_path = os.path.join( | |
| save_video_dir, f"{tokens[idx]}_gen{ns}") | |
| make_file_dirs(save_path) | |
| save_path = save_sample( | |
| back_trans(video), | |
| fps=save_fps if save_fps else fpss[idx], | |
| save_path=save_path, | |
| high_quality=True, | |
| verbose=verbose >= 2, | |
| ) | |
| elif cfg.save_mode == "image_filename": | |
| # save image with their original name | |
| for v_idx, (view, video) in enumerate(zip(VIEW_ORDER, videos)): | |
| # video: C, T, H, W | |
| assert video.shape[1] == len(_filenames[idx]) | |
| for _t in range(video.shape[1]): | |
| _basename = os.path.basename(_filenames[idx][_t][v_idx]) | |
| _basename = os.path.splitext(_basename)[0] | |
| save_path = os.path.join( | |
| save_video_dir, view, | |
| f"{_basename}_gen{ns}.jpg", | |
| ) | |
| make_file_dirs(save_path) | |
| save_path = save_sample( | |
| back_trans(video[:, _t:_t+1]), # take single frame | |
| fps=save_fps if save_fps else fpss[idx], | |
| save_path=save_path, | |
| verbose=verbose >= 2, | |
| with_postfix=False, | |
| ) | |
| coordinator.block_all() | |
| total_num += len(y) | |
| if cfg.ignore_ori_imgs or cfg.get("skip_save_original", False): | |
| coordinator.block_all() | |
| continue | |
| # == save_gt == | |
| x = batch.pop("pixel_values").to(device, dtype) | |
| x = rearrange(x, "B T NC C ... -> B NC C T ...") # B, NC, C, T, H, W | |
| # cut to standard length | |
| x = x[:, :, :, slice(None, cut_length)] | |
| _samples = gather_tensors(x, pg=get_data_parallel_group()) | |
| if coordinator.is_master(): | |
| # gather | |
| samples = [] | |
| fpss = [] | |
| tokens = [] | |
| for sample, fps, token in zip(_samples, _fpss, _tokens): # list of B, NC, C, T ... | |
| samples += [s.cpu() for s in sample] # list of NC, C, T ... | |
| fpss += [int(_fps) for _fps in fps] | |
| tokens += [_tk for _tk in token] | |
| # save | |
| for idx, sample in enumerate(samples): # NC, C, T ... | |
| if cfg.save_mode == "single-view": | |
| for view, video in zip(VIEW_ORDER, sample): | |
| save_path = os.path.join( | |
| save_gt_video_dir, f"{tokens[idx]}", | |
| f"{tokens[idx]}_{view}") | |
| make_file_dirs(save_path) | |
| save_path = save_sample( | |
| back_trans(video), | |
| fps=save_fps if save_fps else fpss[idx], | |
| save_path=save_path, | |
| high_quality=True, | |
| verbose=verbose >= 2, | |
| with_postfix=False, | |
| ) | |
| elif cfg.save_mode == "all-in-one": | |
| vid_sample = concat_6_views_pt(sample, oneline=False) | |
| save_path = os.path.join( | |
| save_gt_video_dir, f"{tokens[idx]}") | |
| make_file_dirs(save_path) | |
| save_path = save_sample( | |
| back_trans(vid_sample), | |
| fps=save_fps if save_fps else fpss[idx], | |
| save_path=save_path, | |
| high_quality=True, | |
| verbose=verbose >= 2, | |
| ) | |
| coordinator.block_all() | |
| logger.info("Inference finished.") | |
| logger.info("Saved %s samples to %s", total_num, cfg.save_dir) | |
| coordinator.destroy() | |
| if __name__ == "__main__": | |
| main() | |