Image Segmentation
Transformers.js
ONNX
swin
background-removal
matting
alpha-matting
image-matting
webgpu
client-side
in-browser
fp16
Instructions to use jiabins0303/birefnet-lite-1024-webgpu with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers.js
How to use jiabins0303/birefnet-lite-1024-webgpu with Transformers.js:
// npm i @huggingface/transformers import { pipeline } from '@huggingface/transformers'; // Allocate pipeline const pipe = await pipeline('image-segmentation', 'jiabins0303/birefnet-lite-1024-webgpu');
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())
|