jiabins0303's picture
Add the Split and GatherND patch scripts and a verify script
1ad01ce verified
Raw History Blame Contribute Delete
6.2 kB
"""
Rewrite wide `Split` nodes into trees of narrow ones, so BiRefNet runs on WebGPU.
THE PROBLEM. ORT's WebGPU backend binds one storage buffer per shader variable
and throws above `maxStorageBuffersPerShaderStage` (8 by WebGPU spec default;
some adapters report 10). It chunks `Concat` inputs to stay under that limit but does NOT
chunk `Split` outputs. BiRefNet_lite@1024's decoder contains 59 Splits with 32
outputs each - 33 buffers apiece - so the graph cannot execute at all.
Full diagnosis in the model card (README.md).
WHY THIS IS SAFE. Splitting [7,7,6,6,6] and then splitting each group is exactly
equivalent to splitting all 32 at once: same axis, same sizes, same order, same
output names. It is a structural rewrite with no numerical content, which is why
the acceptance gate is bit-level correlation against the original rather than a
quality metric.
WHY NOT PATCH `Concat` TOO. The same graph has a 1024-input Concat, but ORT
already chunks Concat, and the 512 export proves that path works. Leave it.
python patch_split.py in.onnx out.onnx [--max-outputs 6]
"""
import argparse
import sys
import numpy as np
import onnx
from onnx import helper, numpy_helper
# 6, not 7, and the extra margin is deliberate.
#
# A Split shader almost certainly binds one buffer for the data input plus one
# per output - the `split` sizes tensor has to be read on the host to compute
# output shapes, so it should not consume a binding. That would make 7 outputs
# (8 buffers) fit the WebGPU spec default of 8.
#
# But "almost certainly" is the wrong confidence level for the thing that decides
# whether the model runs on a stranger's GPU, and being wrong costs one extra
# node per group. 6 outputs is 8 buffers even under the pessimistic accounting
# where the sizes input IS bound.
DEFAULT_MAX_OUTPUTS = 6
def group_sizes(sizes, max_outputs):
"""Partition `sizes` into as few near-equal runs as possible, each <= max_outputs."""
n = len(sizes)
chunks = -(-n // max_outputs) # ceil
base, rem = divmod(n, chunks)
groups, i = [], 0
for g in range(chunks):
take = base + (1 if g < rem else 0)
groups.append(sizes[i : i + take])
i += take
return groups
def main() -> int:
ap = argparse.ArgumentParser()
ap.add_argument("src")
ap.add_argument("dst")
ap.add_argument("--max-outputs", type=int, default=DEFAULT_MAX_OUTPUTS)
args = ap.parse_args()
model = onnx.load(args.src)
graph = model.graph
# Split sizes can arrive as an initializer or as a Constant node's output.
consts = {init.name: numpy_helper.to_array(init) for init in graph.initializer}
for node in graph.node:
if node.op_type == "Constant" and node.output:
for attr in node.attribute:
if attr.name == "value":
consts[node.output[0]] = numpy_helper.to_array(attr.t)
new_nodes, added_inits, patched, skipped = [], [], 0, []
for node in graph.node:
if node.op_type != "Split" or len(node.output) <= args.max_outputs:
new_nodes.append(node)
continue
axis = next((a.i for a in node.attribute if a.name == "axis"), 0)
sizes = None
if len(node.input) > 1 and node.input[1] in consts:
sizes = [int(v) for v in consts[node.input[1]]]
else:
attr = next((a for a in node.attribute if a.name == "split"), None)
if attr is not None:
sizes = [int(v) for v in attr.ints]
# Without explicit sizes the split is equal-division by output count, and
# rewriting it would require the input shape - which is not reliably
# known here. Leave it and report it as unfixed rather than emit a
# graph that is subtly wrong.
if sizes is None or len(sizes) != len(node.output):
skipped.append(node.name or "(unnamed)")
new_nodes.append(node)
continue
groups = group_sizes(sizes, args.max_outputs)
stem = (node.name or f"split_{patched}").replace("/", "_")
# Level 1: one output per group, each carrying that group's total width.
l1_sizes = np.array([sum(g) for g in groups], dtype=np.int64)
l1_name = f"{stem}__l1_sizes"
added_inits.append(numpy_helper.from_array(l1_sizes, l1_name))
l1_outputs = [f"{stem}__g{i}" for i in range(len(groups))]
new_nodes.append(
helper.make_node(
"Split",
inputs=[node.input[0], l1_name],
outputs=l1_outputs,
name=f"{stem}__l1",
axis=axis,
)
)
# Level 2: each group back to the ORIGINAL output names, in order, so
# every downstream consumer is untouched.
cursor = 0
for gi, group in enumerate(groups):
outs = list(node.output[cursor : cursor + len(group)])
cursor += len(group)
if len(group) == 1:
# A one-way Split is illegal; alias instead.
new_nodes.append(
helper.make_node("Identity", [l1_outputs[gi]], outs, name=f"{stem}__g{gi}_id")
)
continue
g_name = f"{stem}__g{gi}_sizes"
added_inits.append(numpy_helper.from_array(np.array(group, dtype=np.int64), g_name))
new_nodes.append(
helper.make_node(
"Split",
inputs=[l1_outputs[gi], g_name],
outputs=outs,
name=f"{stem}__g{gi}",
axis=axis,
)
)
patched += 1
del graph.node[:]
graph.node.extend(new_nodes)
graph.initializer.extend(added_inits)
onnx.checker.check_model(model, full_check=False)
onnx.save(model, args.dst, save_as_external_data=False)
print(f"patched {patched} Split nodes -> max {args.max_outputs} outputs each")
if skipped:
print(f"SKIPPED {len(skipped)} (no explicit sizes): {skipped[:5]}")
print(f"nodes {len(model.graph.node)}, wrote {args.dst}")
return 1 if skipped else 0
if __name__ == "__main__":
sys.exit(main())