File size: 10,792 Bytes
1ad01ce
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
"""
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())