Files

142 lines
5.6 KiB
Python

# https://github.com/Fannovel16/ComfyUI-Frame-Interpolation/tree/main/vfi_utils.py
# https://github.com/yester31/TensorRT_Examples/blob/main/timm_to_trt_python1/onnx_export.py
import torch
import pathlib
import traceback
import os
from urllib.parse import urlparse
from torch.hub import download_url_to_file, get_dir
from rife_arch import IFNet
import onnx
from onnxsim import simplify
MODEL_TYPE = pathlib.Path(__file__).parent.name
CKPT_NAME_VER_DICT = {
"rife40.pth": "4.0",
"rife41.pth": "4.0",
"rife42.pth": "4.2",
"rife43.pth": "4.3",
"rife44.pth": "4.3",
"rife45.pth": "4.5",
"rife46.pth": "4.6",
"rife47.pth": "4.7",
"rife48.pth": "4.7",
"rife49.pth": "4.7",
"sudo_rife4_269.662_testV1_scale1.pth": "4.0"
}
BASE_MODEL_DOWNLOAD_URLS = [
"https://github.com/styler00dollar/VSGAN-tensorrt-docker/releases/download/models/",
"https://github.com/Fannovel16/ComfyUI-Frame-Interpolation/releases/download/models/",
"https://github.com/dajes/frame-interpolation-pytorch/releases/download/v1.0.0/"
]
TORCH_DEVICE = "cuda:0"
def load_file_from_url(url, model_dir=None, progress=True, file_name=None):
"""Load file form http url, will download models if necessary.
Ref:https://github.com/1adrianb/face-alignment/blob/master/face_alignment/utils.py
Args:
url (str): URL to be downloaded.
model_dir (str): The path to save the downloaded model. Should be a full path. If None, use pytorch hub_dir.
Default: None.
progress (bool): Whether to show the download progress. Default: True.
file_name (str): The downloaded file name. If None, use the file name in the url. Default: None.
Returns:
str: The path to the downloaded file.
"""
if model_dir is None: # use the pytorch hub_dir
hub_dir = get_dir()
model_dir = os.path.join(hub_dir, 'checkpoints')
os.makedirs(model_dir, exist_ok=True)
parts = urlparse(url)
file_name = os.path.basename(parts.path)
if file_name is not None:
file_name = file_name
cached_file = os.path.abspath(os.path.join(model_dir, file_name))
if not os.path.exists(cached_file):
print(f'Downloading: "{url}" to {cached_file}\n')
download_url_to_file(url, cached_file, hash_prefix=None, progress=progress)
return cached_file
def load_file_from_github_release(model_type, ckpt_name):
os.makedirs("models", exist_ok=True)
error_strs = []
for i, base_model_download_url in enumerate(BASE_MODEL_DOWNLOAD_URLS):
try:
return load_file_from_url(base_model_download_url + ckpt_name, "models")
except Exception:
traceback_str = traceback.format_exc()
if i < len(BASE_MODEL_DOWNLOAD_URLS) - 1:
print("Failed! Trying another endpoint.")
error_strs.append(f"Error when downloading from: {base_model_download_url + ckpt_name}\n\n{traceback_str}")
error_str = '\n\n'.join(error_strs)
raise Exception(f"Tried all GitHub base urls to download {ckpt_name} but no suceess. Below is the error log:\n\n{error_str}")
def export_onnx(ckpt_name, ensemble, scale_factor):
print(f"PyTorch version: {torch.__version__}")
print(f"ONNX version: {onnx.__version__}")
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
print(f"Using device: {device}")
model_path = load_file_from_github_release(MODEL_TYPE, ckpt_name)
arch_ver = CKPT_NAME_VER_DICT[ckpt_name]
interpolation_model = IFNet(arch_ver=arch_ver)
interpolation_model.load_state_dict(torch.load(model_path, map_location=TORCH_DEVICE))
interpolation_model.eval().to(TORCH_DEVICE)
# # dummy data
img0 = torch.randn(1, 3, 512, 512).to(TORCH_DEVICE)
img1 = torch.randn(1, 3, 512, 512).to(TORCH_DEVICE)
timestep = torch.tensor([0.5], dtype=torch.float32).to(TORCH_DEVICE)
# result = (interpolation_model(img0, img1, timestep))
# print(result)
onnx_save_name = f"{ckpt_name.split('.')[0]}_ensemble_{ensemble}_scale_{scale_factor}.onnx"
onnx_save_path = os.path.join("models",onnx_save_name)
dynamic_axes = {
'img0': {2: 'height', 3: 'width'},
'img1': {2: 'height', 3: 'width'},
'output': {2: 'height', 3: 'width'},
}
with torch.no_grad(): # Disable gradients for efficiency
torch.onnx.export(interpolation_model,
(img0, img1, timestep), # ensemble + scale_factor set in forward fn
onnx_save_path,
export_params=True,
opset_version=19,
do_constant_folding=True,
input_names=['img0', 'img1', 'timestep'],
output_names=['output'],
dynamic_axes=dynamic_axes
)
print(f"ONNX model exported to: {onnx_save_path}")
# Verify the exported ONNX model
onnx_model = onnx.load(onnx_save_path)
onnx.checker.check_model(onnx_model) # Perform a validity check
print("ONNX model validation successful!")
# print(onnx.helper.printable_graph(onnx_model.graph))
sim_model_path = os.path.join("models", f"{ckpt_name.split('.')[0]}_ensemble_{ensemble}_scale_{scale_factor}_sim.onnx")
print("=> ONNX simplify start!")
sim_onnx_model, check = simplify(onnx_model) # convert(simplify)
onnx.save(sim_onnx_model, sim_model_path)
print("=> ONNX simplify done!")
sim_model = onnx.load(sim_model_path)
onnx.checker.check_model(sim_model)
print("=> ONNX Model exported at ", sim_model_path)
print("=> sim ONNX Model check done!")
export_onnx(ckpt_name="rife47.pth", ensemble=True, scale_factor=1)