jiabins0303's picture
Add the Split and GatherND patch scripts and a verify script
1ad01ce verified
Raw History Blame Contribute Delete
10.8 kB
"""
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())