Support fp16, torch.compile

This commit is contained in:
kijai
2025-03-17 17:26:04 +02:00
parent a9735ed9f8
commit 72ecb48014
5 changed files with 104 additions and 78 deletions
+4 -1
View File
@@ -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
):
+5 -4
View File
@@ -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
+1 -1
View File
@@ -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)
+91 -72
View File
@@ -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)