fix: support MPS in RIFE/FILM frame interpolation warp
The MPS backend implements grid_sample but rejects padding_mode="border" (RuntimeError: MPS: Unsupported Border padding mode), so the Frame Interpolation node fails on Apple Silicon for every RIFE and FILM model. PYTORCH_ENABLE_MPS_FALLBACK=1 does not help, as the op is implemented and only the argument is rejected. The upstream PyTorch issue (#125098) has been open since April 2024. Clamping the sampling grid to the valid range makes padding_mode="zeros" exactly equivalent to "border": with align_corners=True that range is [-1, 1], and with align_corners=False it is +-(1 - 1/size) per axis. Gated on device.type == "mps", so CUDA and CPU paths are unchanged.
This commit is contained in:
parent
15989f87ca
commit
012d0559cd
|
|
@ -33,7 +33,17 @@ def _warp_core(image, flow, grid_x, grid_y):
|
||||||
dx = flow[:, 0].float() / (W * 0.5)
|
dx = flow[:, 0].float() / (W * 0.5)
|
||||||
dy = flow[:, 1].float() / (H * 0.5)
|
dy = flow[:, 1].float() / (H * 0.5)
|
||||||
grid = torch.stack([grid_x[None, None, :] + dx, grid_y[None, :, None] + dy], dim=3)
|
grid = torch.stack([grid_x[None, None, :] + dx, grid_y[None, :, None] + dy], dim=3)
|
||||||
return F.grid_sample(image.float(), grid, mode="bilinear", padding_mode="border", align_corners=False).to(dtype)
|
padding_mode = "border"
|
||||||
|
if image.device.type == "mps":
|
||||||
|
# MPS does not implement "border" padding. With align_corners=False the valid
|
||||||
|
# sampling range is +-(1 - 1/size) per axis (grid[..., 0] is x -> width,
|
||||||
|
# grid[..., 1] is y -> height), so clamping there makes "zeros" produce
|
||||||
|
# results identical to "border".
|
||||||
|
bx = 1.0 - 1.0 / image.shape[-1]
|
||||||
|
by = 1.0 - 1.0 / image.shape[-2]
|
||||||
|
grid = torch.stack([grid[..., 0].clamp(-bx, bx), grid[..., 1].clamp(-by, by)], dim=-1)
|
||||||
|
padding_mode = "zeros"
|
||||||
|
return F.grid_sample(image.float(), grid, mode="bilinear", padding_mode=padding_mode, align_corners=False).to(dtype)
|
||||||
|
|
||||||
|
|
||||||
def build_image_pyramid(image, pyramid_levels):
|
def build_image_pyramid(image, pyramid_levels):
|
||||||
|
|
|
||||||
|
|
@ -12,7 +12,14 @@ def _warp(img, flow, warp_grids):
|
||||||
base_grid, flow_div = warp_grids[(H, W)]
|
base_grid, flow_div = warp_grids[(H, W)]
|
||||||
flow_norm = torch.cat([flow[:, 0:1] / flow_div[0], flow[:, 1:2] / flow_div[1]], 1).float()
|
flow_norm = torch.cat([flow[:, 0:1] / flow_div[0], flow[:, 1:2] / flow_div[1]], 1).float()
|
||||||
grid = (base_grid.expand(B, -1, -1, -1) + flow_norm).permute(0, 2, 3, 1)
|
grid = (base_grid.expand(B, -1, -1, -1) + flow_norm).permute(0, 2, 3, 1)
|
||||||
return F.grid_sample(img.float(), grid, mode="bilinear", padding_mode="border", align_corners=True).to(img.dtype)
|
padding_mode = "border"
|
||||||
|
if img.device.type == "mps":
|
||||||
|
# MPS does not implement "border" padding. With align_corners=True the valid
|
||||||
|
# sampling range is exactly [-1, 1], so clamping the grid to it makes "zeros"
|
||||||
|
# produce results identical to "border".
|
||||||
|
grid = grid.clamp(-1.0, 1.0)
|
||||||
|
padding_mode = "zeros"
|
||||||
|
return F.grid_sample(img.float(), grid, mode="bilinear", padding_mode=padding_mode, align_corners=True).to(img.dtype)
|
||||||
|
|
||||||
|
|
||||||
class Head(nn.Module):
|
class Head(nn.Module):
|
||||||
|
|
|
||||||
Loading…
Reference in New Issue