web: hand the upscaler's activations to the chip that is running them
The WebGPU execution provider has no PReLU kernel. The model is 34 convolutions with a PReLU after every one of them, so an export that took the GPU path was split 33 times: each activation came off the chip to be activated on the processor and went straight back, a 64-channel map in both directions, per tile. A machine with a good graphics chip was not exporting any faster for having it. PReLU(x) is exactly Relu(x) - slope * Relu(-x), and Relu, Neg, Mul and Sub the provider does implement, so scripts/realesr-gpu.py writes the 33 activations out as those four and drops the slopes nobody reads any more. The model file is the output of that script, not the file as published. One 256x256 tile through the model before and after, on a WebGPU session: the runtime no longer reports nodes left off the preferred provider (it did, once, before) and the processor path answers bit for bit what it answered before. The warning itself cannot be switched off from here - env.logLevel is read when the runtime module initialises, before any of this runs - so the graph was fixed rather than the lines hidden.
This commit is contained in:
@@ -0,0 +1,73 @@
|
||||
# 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 <downloaded>.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])
|
||||
Reference in New Issue
Block a user