Files
automatic/modules/rife/model_rife.py
T
CalamitousFelicitousness dc4c58d0cb feat(rife): upgrade vendored RIFE to Practical-RIFE v4.25
- Vendor IFNet_HDv3 v4.25 (5 IFBlocks, Head encoder, feat channel) and
  v4 warplayer with explicit (tenFlow_div, backwarp_tenGrid) signature
- Rewrite RifeModel.inference for the new forward signature with
  per-(H,W,device,dtype) caching of tenFlow_div and backwarp_tenGrid
- Force fp32 inference: bf16 produced visible checkerboard at the new
  IFNet's depth (was hidden by v3.9's shallower architecture)
- Crop padded frames in interpolate_nchw before output (was missing,
  produced gray bar on non-128-aligned inputs)
- Drop training scaffolding (AdamW, EPE/SOBEL, update method)
- Log obsolete legacy v3.9 weights file on first v4.25 load instead of
  silently deleting user data
- Default download URL is HolyWu vs-rife mirror (MIT, byte-identical
  upstream weights); swap to project-hosted URL before merge
2026-04-25 21:14:18 +01:00

70 lines
2.8 KiB
Python

import torch
from modules.rife.model_ifnet import IFNet
from modules import devices
class RifeModel:
def __init__(self, local_rank=-1):
self.flownet = IFNet()
self.device()
self.version = 4.25
self.tenFlow_div_cache = {}
self.backwarp_tenGrid_cache = {}
if local_rank != -1:
from torch.nn.parallel import DistributedDataParallel as DDP
self.flownet = DDP(self.flownet, device_ids=[local_rank], output_device=local_rank)
def train(self):
self.flownet.train()
def eval(self):
self.flownet.eval()
def device(self):
self.flownet.to(devices.device)
self.flownet.to(torch.float32) # bfloat16 produces visible checkerboard artifacts at the new IFNet's depth
def load_model(self, model_file, rank=0):
def convert(param):
if rank == -1:
return { k.replace("module.", ""): v for k, v in param.items() if "module." in k }
else:
return param
if rank <= 0:
if torch.cuda.is_available():
self.flownet.load_state_dict(convert(torch.load(model_file)), False)
else:
self.flownet.load_state_dict(convert(torch.load(model_file, map_location='cpu')), False)
def save_model(self, model_file, rank=0):
if rank == 0:
torch.save(self.flownet.state_dict(), model_file)
def grid_for(self, h, w, device, dtype):
key = (h, w, str(device), str(dtype))
grid = self.backwarp_tenGrid_cache.get(key)
if grid is None:
tenHorizontal = torch.linspace(-1.0, 1.0, w, device=device, dtype=dtype).view(1, 1, 1, w).expand(1, -1, h, -1)
tenVertical = torch.linspace(-1.0, 1.0, h, device=device, dtype=dtype).view(1, 1, h, 1).expand(1, -1, -1, w)
grid = torch.cat([tenHorizontal, tenVertical], 1)
self.backwarp_tenGrid_cache[key] = grid
div = self.tenFlow_div_cache.get(key)
if div is None:
div = torch.tensor([(w - 1.0) / 2.0, (h - 1.0) / 2.0], device=device, dtype=dtype)
self.tenFlow_div_cache[key] = div
return grid, div
def inference(self, img0, img1, timestep=0.5, scale=1.0):
in_dtype = img0.dtype
img0 = img0.float()
img1 = img1.float()
n, _c, h, w = img0.shape
device = img0.device
backwarp_tenGrid, tenFlow_div = self.grid_for(h, w, device, torch.float32)
self.flownet.scale_list = [16 / scale, 8 / scale, 4 / scale, 2 / scale, 1 / scale]
f0 = self.flownet.encode(img0)
f1 = self.flownet.encode(img1)
timestep_t = torch.full((n, 1, h, w), timestep, device=device, dtype=torch.float32)
out = self.flownet(img0, img1, timestep_t, tenFlow_div, backwarp_tenGrid, f0, f1)
return out.to(in_dtype)