""" 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())