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');
Download scripts/patch_deform.py from jiabins0303/birefnet-lite-1024-webgpu: direct link, hf CLI and curl.
- Browser
- Download file 10.8 kB
-
https://huggingface.co/jiabins0303/birefnet-lite-1024-webgpu/resolve/main/scripts/patch_deform.py
- Command line
-
hf download hf://jiabins0303/birefnet-lite-1024-webgpu/scripts/patch_deform.py
-
curl -L -o patch_deform.py https://huggingface.co/jiabins0303/birefnet-lite-1024-webgpu/resolve/main/scripts/patch_deform.py
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()) | |