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: 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())
|