""" Move BiRefNet's deformable-convolution nodes off the CPU execution provider. WHY. `patch_split.py` fixed the FIRST ceiling (WebGPU's storage-buffer limit). The graph then still died, and ORT's own verbose log named the second one exactly: transformer_memcpy.cc:340 AddCopyNode] Add MemcpyFromHost after /decoder/decoder_block1/dec_att/aspp_deforms.2/atrous_conv/GatherND_output_0 for WebGpuExecutionProvider 100 such copies at 1024, and NOT ONE of them outside `atrous_conv`: 80 GatherND plus 20 Sum, across the 20 deformable-conv blocks (decoder_block1-4 and squeeze_module, four ASPP branches each). `deform_conv2d` has no ONNX op, so the exporter decomposes it into bilinear sampling built from GatherND - and in the onnxruntime-web build we tested (August 2026) the WebGPU EP did not take those nodes, so they ran on the CPU EP. CPU tensors live on the 32-bit wasm heap, and at 1024 each GatherND output is [1, 1, 64, 49, 256, 256] fp16 = 392MB with five live at once. That is the `std::bad_alloc`. At 512 the same tensors are 98MB and it fits, which is the entire reason 512 ships and 1024 does not. WHAT THIS DOES. Rewrites each GatherND into ops the WebGPU EP demonstrably has - plain `Gather` is used 72 times elsewhere in this graph and never appears in the memcpy list, so it is known-supported here rather than assumed: GatherND(data[1,1,H,W,C], idx[1,1,P,2], batch_dims=2) -> [1,1,P,C] becomes flat = Cast(Reshape(idx,[P,2]), int32) # WebGPU has no int64 maths y,x = Slice(flat,0:1,axis1), Slice(flat,1:2,axis1) lin = Reshape(Add(Mul(y, W), x), [P]) out = Reshape(Gather(Reshape(data,[H*W,C]), lin, axis=0), [1,1,P,C]) and each variadic `Sum` into a left-to-right chain of binary `Add`, for the same reason (see --sum-mode for why a chain and not a tree). `Mul(y, W) + x` is only equal to the original gather if the indices are already in range, so this refuses to rewrite unless it can SEE a clamp upstream. A silent wrap on a negative index would corrupt one pixel in a way no aggregate metric would ever show. Verify with `verify_patch.py` - this is a structural rewrite and must be bit-identical, not merely close. python patch_deform.py in.onnx out.onnx """ import argparse import collections import sys import onnx from onnx import TensorProto, helper, numpy_helper # Ops that establish an index is already in range. If none of these appears # between the index arithmetic and the gather, the rewrite is refused. CLAMP_OPS = {"Clip", "Min", "Max"} def shapes_of(model): inferred = onnx.shape_inference.infer_shapes(model, strict_mode=False, data_prop=True) out = {} for vi in list(inferred.graph.value_info) + list(inferred.graph.output) + list(inferred.graph.input): dims = [] ok = True for d in vi.type.tensor_type.shape.dim: if d.HasField("dim_value"): dims.append(d.dim_value) else: ok = False break if ok: out[vi.name] = dims return out def clamped_upstream(name, producer, depth=12): """Walk back from `name` looking for a clamp, stopping at the first fan-in.""" seen = set() stack = [(name, 0)] while stack: n, d = stack.pop() if d > depth or n in seen: continue seen.add(n) node = producer.get(n) if node is None: continue if node.op_type in CLAMP_OPS: return True for inp in node.input: if inp: stack.append((inp, d + 1)) return False def main(): ap = argparse.ArgumentParser() ap.add_argument("src") ap.add_argument("dst") ap.add_argument("--allow-unclamped", action="store_true", help="rewrite even when no clamp is visible upstream (unsafe; prints which)") ap.add_argument("--sum-mode", choices=["chain", "tree", "none"], default="chain", help="how to decompose variadic Sum. 'chain' is left-to-right and is the " "DEFAULT because floating-point addition is not associative: a balanced " "'tree' computes (a+b)+(c+d) where ORT's Sum accumulates ((a+b)+c)+d, " "and in fp16 that alone moved the output by 4.9e-03 - measured, and the " "sole reason a structural rewrite failed its bit-identity gate. " "'none' leaves Sum alone, to isolate the GatherND rewrite.") args = ap.parse_args() model = onnx.load(args.src) graph = model.graph shapes = shapes_of(model) producer = {o: n for n in graph.node for o in n.output} new_nodes = [] consts = [] stats = collections.Counter() skipped = [] uid = 0 def const_i64(vals, tag): nonlocal uid uid += 1 name = f"deformpatch_{tag}_{uid}" consts.append(numpy_helper.from_array( __import__("numpy").array(vals, dtype="int64"), name)) return name def const_i32(vals, tag): nonlocal uid uid += 1 name = f"deformpatch_{tag}_{uid}" consts.append(numpy_helper.from_array( __import__("numpy").array(vals, dtype="int32"), name)) return name for node in graph.node: if node.op_type == "Sum" and len(node.input) > 2 and args.sum_mode != "none": # Variadic Sum has no WebGPU kernel; Add does. if args.sum_mode == "chain": # Left-to-right, matching ORT's own accumulation order so the # rewrite stays bit-identical in fp16. See --sum-mode. acc = node.input[0] for i in range(1, len(node.input)): last = i == len(node.input) - 1 out = node.output[0] if last else f"{node.name}_acc{i}" new_nodes.append(helper.make_node( "Add", [acc, node.input[i]], [out], name=f"{node.name}_acc{i}")) acc = out stats["Sum->Add chain"] += 1 else: level = list(node.input) tier = 0 while len(level) > 1: nxt = [] for i in range(0, len(level), 2): if i + 1 == len(level): nxt.append(level[i]) continue last = len(level) <= 2 out = node.output[0] if last else f"{node.name}_add{tier}_{i}" new_nodes.append(helper.make_node( "Add", [level[i], level[i + 1]], [out], name=f"{node.name}_add{tier}_{i}")) nxt.append(out) level = nxt tier += 1 stats["Sum->Add tree"] += 1 continue if node.op_type != "GatherND": new_nodes.append(node) continue attrs = {a.name: helper.get_attribute_value(a) for a in node.attribute} batch_dims = attrs.get("batch_dims", 0) data_s = shapes.get(node.input[0]) idx_s = shapes.get(node.input[1]) # Only the exact deform-conv pattern is rewritten. Anything else keeps # its GatherND rather than being guessed at. if not (batch_dims == 2 and data_s and idx_s and len(data_s) == 5 and len(idx_s) == 4 and data_s[0] == 1 and data_s[1] == 1 and idx_s[0] == 1 and idx_s[1] == 1 and idx_s[3] == 2): skipped.append((node.name, f"batch_dims={batch_dims} data={data_s} idx={idx_s}")) new_nodes.append(node) stats["GatherND kept (pattern)"] += 1 continue if not clamped_upstream(node.input[1], producer) and not args.allow_unclamped: skipped.append((node.name, "no clamp visible upstream")) new_nodes.append(node) stats["GatherND kept (unclamped)"] += 1 continue _, _, H, W, C = data_s P = idx_s[2] b = node.name.replace("/", "_") n_idx2 = f"{b}__idx2" n_i32 = f"{b}__i32" n_y = f"{b}__y" n_x = f"{b}__x" n_yw = f"{b}__yw" n_lin2 = f"{b}__lin2" n_lin = f"{b}__lin" n_data2 = f"{b}__data2" n_gat = f"{b}__gat" # ORDER MATTERS, and it is measured. The index chain arrives as int64 and # WebGPU has no int64, so everything up to the Cast necessarily runs on # the CPU EP. Casting FIRST (the obvious reading) then doing the # arithmetic on the GPU means both [P,1] halves cross the boundary: # 160 copies of 12.8MB each at 1024. Doing the arithmetic in int64 while # it is already on the CPU and casting ONCE means a single [P] tensor # crosses per node - half the transfers and half the bytes. new_nodes += [ helper.make_node("Reshape", [node.input[1], const_i64([P, 2], "s")], [n_idx2], name=n_idx2), helper.make_node("Slice", [n_idx2, const_i64([0], "st"), const_i64([1], "en"), const_i64([1], "ax")], [n_y], name=n_y), helper.make_node("Slice", [n_idx2, const_i64([1], "st"), const_i64([2], "en"), const_i64([1], "ax")], [n_x], name=n_x), helper.make_node("Mul", [n_y, const_i64([W], "w")], [n_yw], name=n_yw), helper.make_node("Add", [n_yw, n_x], [n_lin2], name=n_lin2), helper.make_node("Reshape", [n_lin2, const_i64([P], "s")], [n_i32], name=n_i32), helper.make_node("Cast", [n_i32], [n_lin], to=TensorProto.INT32, name=n_lin), helper.make_node("Reshape", [node.input[0], const_i64([H * W, C], "s")], [n_data2], name=n_data2), helper.make_node("Gather", [n_data2, n_lin], [n_gat], axis=0, name=n_gat), helper.make_node("Reshape", [n_gat, const_i64([1, 1, P, C], "s")], [node.output[0]], name=f"{b}__out"), ] stats["GatherND->Gather"] += 1 del graph.node[:] graph.node.extend(new_nodes) graph.initializer.extend(consts) onnx.checker.check_model(model, full_check=False) onnx.save(model, args.dst) print(f"{args.src} -> {args.dst}") for k, v in stats.most_common(): print(f" {k}: {v}") if skipped: print(f"\n {len(skipped)} node(s) left as-is:") for name, why in skipped[:10]: print(f" {name}: {why}") print(f"\n nodes: {len(new_nodes)} (+{len(consts)} new initializers)") print("\nNow run verify_patch.py - a structural rewrite must be bit-identical.") return 0 if __name__ == "__main__": sys.exit(main())