# The WebGPU execution provider has no PReLU kernel. The upscaler is 34 convolutions # with a PReLU after every one of them, so the runtime splits the graph 33 times: # each activation travels off the GPU to the processor and the result travels back, # a 64-channel map in both directions, per tile. A machine with a good graphics chip # was therefore not exporting any faster for having it. # # Relu, Neg, Mul and Sub it does have, and # PRelu(x) = Relu(x) - slope * Relu(-x) # holds exactly, negative slopes included, so this rewrites the activations into # those four and leaves the weights untouched. # # python3 realesr-gpu.py .onnx public/models/realesr-general-x4v3.onnx # # The source is the x4v3 model of https://github.com/xinntao/Real-ESRGAN # (realesr-general-x4v3.pth, converted to ONNX at opset 17). import sys import numpy as np import onnx from onnx import TensorProto, helper, numpy_helper def rewrite(src, dst): model = onnx.load(src) graph = model.graph rank = len(graph.input[0].type.tensor_type.shape.dim) by_name = {i.name: i for i in graph.initializer} nodes, made = [], 0 for node in graph.node: if node.op_type != "PRelu": nodes.append(node) continue x, slope = node.input # The slope comes out of PyTorch as (C,1,1) and has to line up with a # (1,C,H,W) activation, so give it the batch axis it is missing. s = numpy_helper.to_array(by_name[slope]) if s.shape != (1,) + tuple(s.shape): s = s.reshape((1,) + tuple(s.shape)) name = f"{slope}__broadcast" if name not in by_name: graph.initializer.append(numpy_helper.from_array(s, name)) by_name[name] = graph.initializer[-1] slope = name p, n, q, scaled = f"{node.output[0]}.pos", f"{node.output[0]}.neg", f"{node.output[0]}.relu", f"{node.output[0]}.scaled" nodes += [ helper.make_node("Relu", [x], [p], name=p), helper.make_node("Neg", [x], [n], name=n), helper.make_node("Relu", [n], [q], name=q), helper.make_node("Mul", [q, slope], [scaled], name=scaled), helper.make_node("Sub", [p, scaled], [node.output[0]], name=node.name), ] made += 1 del graph.node[:] graph.node.extend(nodes) # The slopes the PReLU nodes used to read are now unreachable, and the runtime # warns about every one of them while it loads unless they go. used = {name for node in graph.node for name in node.input} kept = [i for i in graph.initializer if i.name in used] print(f"dropping {len(graph.initializer) - len(kept)} unused initializers") del graph.initializer[:] graph.initializer.extend(kept) onnx.checker.check_model(model) onnx.save(model, dst) print(f"rewrote {made} PReLU nodes -> {dst}") assert made, "no PReLU node found: wrong model?" assert rank >= 2 if __name__ == "__main__": rewrite(sys.argv[1], sys.argv[2])