diff --git a/gimmvfi/generalizable_INR/gimmvfi_f.py b/gimmvfi/generalizable_INR/gimmvfi_f.py index 73be634..f0001e5 100644 --- a/gimmvfi/generalizable_INR/gimmvfi_f.py +++ b/gimmvfi/generalizable_INR/gimmvfi_f.py @@ -27,12 +27,14 @@ from .modules.softsplat import softsplat class GIMMVFI_F(nn.Module): Config = GIMMVFIConfig - def __init__(self, config: GIMMVFIConfig): + def __init__(self, dtype, config: GIMMVFIConfig): super().__init__() self.config = config = config.copy() self.hyponet_config = config.hyponet self.raft_iter = config.raft_iter + self.dtype = dtype + ######### Encoder and Decoder Settings ######### #self.flow_estimator = initialize_Flowformer() f_dims = [256, 128] @@ -201,6 +203,7 @@ class GIMMVFI_F(nn.Module): img_warp = mask * img0_warp + (1 - mask) * img1_warp return img_warp + @torch.compiler.disable() def frame_synthesize( self, img_xs, flow_t, features0, features1, corr_fn, cur_t, full_img=None ): diff --git a/gimmvfi/generalizable_INR/gimmvfi_r.py b/gimmvfi/generalizable_INR/gimmvfi_r.py index 62e4c38..3e2b264 100644 --- a/gimmvfi/generalizable_INR/gimmvfi_r.py +++ b/gimmvfi/generalizable_INR/gimmvfi_r.py @@ -34,7 +34,7 @@ from .modules.softsplat import softsplat class GIMMVFI_R(nn.Module): Config = GIMMVFIConfig - def __init__(self, config: GIMMVFIConfig): + def __init__(self, dtype, config: GIMMVFIConfig): super().__init__() self.config = config = config.copy() self.hyponet_config = config.hyponet @@ -44,7 +44,7 @@ class GIMMVFI_R(nn.Module): #self.flow_estimator = initialize_RAFT() cur_f_dims = [128, 96] f_dims = [256, 128] - self.dtype = torch.float32 + self.dtype = dtype skip_channels = f_dims[-1] // 2 self.num_flows = 3 @@ -126,10 +126,10 @@ class GIMMVFI_R(nn.Module): def cal_bidirection_flow(self, im0, im1, iters=20): f01, features0, fnet0 = self.flow_estimator( - im0, im1, return_feat=True, iters=20 + im0.to(self.dtype), im1.to(self.dtype), return_feat=True, iters=20 ) f10, features1, fnet1 = self.flow_estimator( - im1, im0, return_feat=True, iters=20 + im1.to(self.dtype), im0.to(self.dtype), return_feat=True, iters=20 ) corr_fn = BidirCorrBlock(self.amt_fproj(fnet0), self.amt_fproj(fnet1), radius=4) features0 = [ @@ -220,6 +220,7 @@ class GIMMVFI_R(nn.Module): img_warp = mask * img0_warp + (1 - mask) * img1_warp return img_warp + @torch.compiler.disable() def frame_synthesize( self, img_xs, flow_t, features0, features1, corr_fn, cur_t, full_img=None ): diff --git a/gimmvfi/generalizable_INR/modules/softsplat.py b/gimmvfi/generalizable_INR/modules/softsplat.py index 369b99b..415fc51 100644 --- a/gimmvfi/generalizable_INR/modules/softsplat.py +++ b/gimmvfi/generalizable_INR/modules/softsplat.py @@ -261,6 +261,7 @@ def cuda_kernel(strFunction: str, strKernel: str, objVariables: typing.Dict): @cupy.memoize(for_each_device=True) +@torch.compiler.disable() def cuda_launch(strKey: str): try: os.environ.setdefault("CUDA_HOME", cupy.cuda.get_cuda_path()) @@ -284,6 +285,7 @@ def cuda_launch(strKey: str): ########################################################## +@torch.compiler.disable() def softsplat(tenIn, tenFlow, tenMetric, strMode, return_norm=False): assert strMode.split("-")[0] in ["sum", "avg", "linear", "softmax"] @@ -449,6 +451,7 @@ class softsplat_func(torch.autograd.Function): # end @staticmethod + @torch.compiler.disable() @torch.amp.custom_bwd(device_type="cuda") def backward(self, tenOutgrad): tenIn, tenFlow = self.saved_tensors diff --git a/gimmvfi/generalizable_INR/raft/update.py b/gimmvfi/generalizable_INR/raft/update.py index ced6df0..90abe56 100644 --- a/gimmvfi/generalizable_INR/raft/update.py +++ b/gimmvfi/generalizable_INR/raft/update.py @@ -143,7 +143,7 @@ class BasicUpdateBlock(nn.Module): ) def forward(self, net, inp, corr, flow, upsample=True): - motion_features = self.encoder(flow, corr) + motion_features = self.encoder(flow.to(inp), corr.to(inp)) inp = torch.cat([inp, motion_features], dim=1) net = self.gru(net, inp) diff --git a/nodes.py b/nodes.py index 9c4f2d9..afbb69f 100644 --- a/nodes.py +++ b/nodes.py @@ -22,6 +22,8 @@ from .gimmvfi.generalizable_INR.flowformer.configs.submission import get_cfg from .gimmvfi.utils.flow_viz import flow_to_image from .gimmvfi.utils.utils import InputPadder, RaftArgs, easydict_to_dict +from contextlib import nullcontext + import logging logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s') log = logging.getLogger(__name__) @@ -39,6 +41,10 @@ class DownloadAndLoadGIMMVFIModel: "gimmvfi_f_arb_lpips_fp32.safetensors" ],), }, + "optional": { + "precision": (["fp32", "bf16", "fp16"], {"default": "fp32"}), + "torch_compile": ("BOOLEAN", {"default": False, "tooltip": "Compile part of the model with torch.compile, requires Triton"}), + }, } RETURN_TYPES = ("GIMMVIF_MODEL",) @@ -46,11 +52,13 @@ class DownloadAndLoadGIMMVFIModel: FUNCTION = "loadmodel" CATEGORY = "GIMM-VFI" - def loadmodel(self, model): + def loadmodel(self, model, precision="fp32", torch_compile=False): device = mm.get_torch_device() offload_device = mm.unet_offload_device() + dtype = {"fp8_e4m3fn": torch.float8_e4m3fn, "fp8_e4m3fn_fast": torch.float8_e4m3fn, "bf16": torch.bfloat16, "fp16": torch.float16, "fp16_fast": torch.float16, "fp32": torch.float32}[precision] + download_path = os.path.join(folder_paths.models_dir, 'interpolation', 'gimm-vfi') model_path = os.path.join(download_path, model) @@ -93,7 +101,7 @@ class DownloadAndLoadGIMMVFIModel: # load model if "gimmvfi_r" in model: - model = GIMMVFI_R(config) + model = GIMMVFI_R(dtype, config) #load RAFT raft_args = RaftArgs( small=False, @@ -104,31 +112,27 @@ class DownloadAndLoadGIMMVFIModel: raft_model = RAFT(raft_args) raft_sd = load_torch_file(flow_model_path) raft_model.load_state_dict(raft_sd, strict=True) - raft_model.to(device) + raft_model.to(dtype).to(device) flow_estimator = raft_model elif "gimmvfi_f" in model: - model = GIMMVFI_F(config) + model = GIMMVFI_F(dtype, config) cfg = get_cfg() flowformer = FlowFormer(cfg.latentcostformer) flowformer_sd = load_torch_file(flow_model_path) flowformer.load_state_dict(flowformer_sd, strict=True) - flow_estimator = flowformer + flow_estimator = flowformer.to(dtype).to(device) sd = load_torch_file(model_path) model.load_state_dict(sd, strict=False) - model.flow_estimator = flow_estimator - model = model.eval().to(device) + model = model.eval().to(dtype).to(device) + + if torch_compile: + model = torch.compile(model) return (model,) - -def load_image(img_path): - img = Image.open(img_path) - raw_img = np.array(img.convert("RGB")) - img = torch.from_numpy(raw_img.copy()).permute(2, 0, 1) / 255.0 - return img.to(torch.float).unsqueeze(0) #region Interpolate class GIMMVFI_interpolate: @@ -142,6 +146,9 @@ class GIMMVFI_interpolate: "interpolation_factor": ("INT", {"default": 8, "min": 1, "max": 100, "step": 1}), "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), }, + "optional": { + "output_flows": ("BOOLEAN", {"default": False, "tooltip": "Output the flow tensors"}), + }, } RETURN_TYPES = ("IMAGE", "IMAGE",) @@ -149,7 +156,7 @@ class GIMMVFI_interpolate: FUNCTION = "interpolate" CATEGORY = "PyramidFlowWrapper" - def interpolate(self, gimmvfi_model, images, ds_factor, interpolation_factor,seed): + def interpolate(self, gimmvfi_model, images, ds_factor, interpolation_factor,seed, output_flows=False): mm.soft_empty_cache() images = images.permute(0, 3, 1, 2) torch.manual_seed(seed) @@ -158,77 +165,89 @@ class GIMMVFI_interpolate: device = mm.get_torch_device() offload_device = mm.unet_offload_device() - gimmvfi_model.to(device) - + dtype = gimmvfi_model.dtype + + out_images_list = [] flows = [] start = 0 end = images.shape[0] - 1 pbar = ProgressBar(images.shape[0] - 1) - for j in tqdm(range(start, end)): - I0 = images[j].unsqueeze(0) - I2 = images[j+1].unsqueeze(0) - if j == start: - out_images_list.append(I0.squeeze(0).permute(1, 2, 0)) + autocast_device = mm.get_autocast_device(device) + cast_context = torch.autocast(device_type=autocast_device, dtype=dtype) if dtype != torch.float32 else nullcontext() + + with cast_context: + for j in tqdm(range(start, end)): + I0 = images[j].unsqueeze(0) + I2 = images[j+1].unsqueeze(0) + + if j == start: + out_images_list.append(I0.squeeze(0).permute(1, 2, 0)) + + padder = InputPadder(I0.shape, 32) + I0, I2 = padder.pad(I0, I2) + xs = torch.cat((I0.unsqueeze(2), I2.unsqueeze(2)), dim=2).to(device, non_blocking=True) + + batch_size = xs.shape[0] + s_shape = xs.shape[-2:] - padder = InputPadder(I0.shape, 32) - I0, I2 = padder.pad(I0, I2) - xs = torch.cat((I0.unsqueeze(2), I2.unsqueeze(2)), dim=2).to(device, non_blocking=True) - - batch_size = xs.shape[0] - s_shape = xs.shape[-2:] - - coord_inputs = [ - ( - gimmvfi_model.sample_coord_input( - batch_size, - s_shape, - [1 / interpolation_factor * i], - device=xs.device, - upsample_ratio=ds_factor, - ), - None, + coord_inputs = [ + ( + gimmvfi_model.sample_coord_input( + batch_size, + s_shape, + [1 / interpolation_factor * i], + device=xs.device, + upsample_ratio=ds_factor, + ), + None, + ) + for i in range(1, interpolation_factor) + ] + timesteps = [ + i * 1 / interpolation_factor * torch.ones(xs.shape[0]).to(xs.device)#.to(torch.float) + for i in range(1, interpolation_factor) + ] + + all_outputs = gimmvfi_model(xs, coord_inputs, t=timesteps, ds_factor=ds_factor) + out_frames = [padder.unpad(im) for im in all_outputs["imgt_pred"]] + out_flowts = [padder.unpad(f) for f in all_outputs["flowt"]] + + if output_flows: + flowt_imgs = [ + flow_to_image( + flowt.squeeze().detach().cpu().permute(1, 2, 0).numpy(), + convert_to_bgr=True, + ) + for flowt in out_flowts + ] + I1_pred_img = [ + (I1_pred[0].detach().cpu().permute(1, 2, 0)) + for I1_pred in out_frames + ] + + for i in range(interpolation_factor - 1): + out_images_list.append(I1_pred_img[i]) + if output_flows: + flows.append(flowt_imgs[i]) + + out_images_list.append( + ((padder.unpad(I2)).squeeze().detach().cpu().permute(1, 2, 0)) ) - for i in range(1, interpolation_factor) - ] - timesteps = [ - i * 1 / interpolation_factor * torch.ones(xs.shape[0]).to(xs.device).to(torch.float) - for i in range(1, interpolation_factor) - ] - - all_outputs = gimmvfi_model(xs, coord_inputs, t=timesteps, ds_factor=ds_factor) - out_frames = [padder.unpad(im) for im in all_outputs["imgt_pred"]] - out_flowts = [padder.unpad(f) for f in all_outputs["flowt"]] - - flowt_imgs = [ - flow_to_image( - flowt.squeeze().detach().cpu().permute(1, 2, 0).numpy(), - convert_to_bgr=True, - ) - for flowt in out_flowts - ] - I1_pred_img = [ - (I1_pred[0].detach().cpu().permute(1, 2, 0)) - for I1_pred in out_frames - ] - - for i in range(interpolation_factor - 1): - out_images_list.append(I1_pred_img[i]) - flows.append(flowt_imgs[i]) - - out_images_list.append( - ((padder.unpad(I2)).squeeze().detach().cpu().permute(1, 2, 0)) - ) - pbar.update(1) + pbar.update(1) image_tensors = torch.stack(out_images_list) image_tensors = image_tensors.cpu().float() rgb_images = [cv2.cvtColor(flow, cv2.COLOR_BGR2RGB) for flow in flows] - flow_tensors = torch.stack([torch.from_numpy(image) for image in rgb_images]) - flow_tensors = flow_tensors / 255.0 - flow_tensors = flow_tensors.cpu().float() + + if output_flows: + flow_tensors = torch.stack([torch.from_numpy(image) for image in rgb_images]) + flow_tensors = flow_tensors / 255.0 + flow_tensors = flow_tensors.cpu().float() + else: + flow_tensors = torch.zeros(1, 64, 64, 3) return (image_tensors, flow_tensors)