Support fp16, torch.compile
This commit is contained in:
@@ -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
|
||||
):
|
||||
|
||||
@@ -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
|
||||
):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user