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