RIFE_fp32_timestep / patch_rife_timestep.py
walterlow's picture
Upload patch_rife_timestep.py with huggingface_hub
7bca327 verified
Raw
History Blame Contribute Delete
2.79 kB
"""
Expose RIFE's baked t=0.5 as a runtime scalar input.
FuryTMP/RIFE_fp32 (MIT) traced its IFNet with timestep=0.5 folded into a Constant, so the
exported graph takes only 6 channels and can produce the midpoint frame. The timestep plane
is still built internally as ((Slice * 0) + 1) * 0.5 and fed to the head Concat of every
IFBlock -- promoting that 0.5 to a graph input restores arbitrary-phase interpolation.
Verified: output is bit-identical to the original at t=0.5, and ORT folds the patched graph to
the same 336 nodes with zero shape-computation ops under freeDimensionOverrides.
Usage: python patch_rife_timestep.py RIFE_fp32.onnx RIFE_fp32_timestep.onnx
Feed the new input as a 0-d array: np.array(t, dtype=np.float32) (not np.float32(t)).
"""
import sys
import onnx
from onnx import TensorProto, helper, numpy_helper
TIMESTEP_INPUT = "timestep"
def find_baked_timestep(graph):
"""The scalar Constant whose only consumer is the Mul feeding the timestep plane."""
consumers = {}
for node in graph.node:
for name in node.input:
consumers.setdefault(name, []).append(node)
candidates = []
for node in graph.node:
if node.op_type != "Constant":
continue
value = next((a.t for a in node.attribute if a.name == "value"), None)
if value is None:
continue
array = numpy_helper.to_array(value)
if array.shape != () or not (0.0 < float(array) < 1.0):
continue
users = consumers.get(node.output[0], [])
if len(users) == 1 and users[0].op_type == "Mul":
candidates.append((node, float(array), users[0]))
if len(candidates) != 1:
raise SystemExit(
f"expected exactly one baked-timestep Constant, found {len(candidates)}: "
f"{[(n.name, v) for n, v, _ in candidates]}"
)
return candidates[0]
def main(src, dst):
model = onnx.load(src)
graph = model.graph
if any(i.name == TIMESTEP_INPUT for i in graph.input):
raise SystemExit(f"{src} already exposes a '{TIMESTEP_INPUT}' input")
const_node, baked_value, mul_node = find_baked_timestep(graph)
print(f"found {const_node.name} = {baked_value} -> consumed by {mul_node.name}")
graph.node.remove(const_node)
for i, name in enumerate(mul_node.input):
if name == const_node.output[0]:
mul_node.input[i] = TIMESTEP_INPUT
graph.input.append(helper.make_tensor_value_info(TIMESTEP_INPUT, TensorProto.FLOAT, []))
onnx.checker.check_model(model, full_check=True)
onnx.save(model, dst)
print(f"wrote {dst}; inputs = {[i.name for i in graph.input]}")
if __name__ == "__main__":
if len(sys.argv) != 3:
raise SystemExit(__doc__)
main(sys.argv[1], sys.argv[2])