File size: 6,195 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
"""
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())