Instructions to use CoderViking/birefnet-lite-onnx with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- BiRefNet
How to use CoderViking/birefnet-lite-onnx with BiRefNet:
# Option 1: use with transformers from transformers import AutoModelForImageSegmentation birefnet = AutoModelForImageSegmentation.from_pretrained("CoderViking/birefnet-lite-onnx", trust_remote_code=True)# Option 2: use with BiRefNet # Install from https://github.com/ZhengPeng7/BiRefNet from models.birefnet import BiRefNet model = BiRefNet.from_pretrained("CoderViking/birefnet-lite-onnx") - Notebooks
- Google Colab
- Kaggle
| #!/usr/bin/env python | |
| """Browser-tuned ONNX export of BiRefNet_lite (ZhengPeng7/BiRefNet_lite, MIT). | |
| Why a custom export exists: the community export (onnx-community/BiRefNet_lite-ONNX) | |
| is a dynamic-shape trace whose deformable convolutions decompose into | |
| GatherND/ScatterND/Clip chains. Those ops fall back to CPU on the ONNX Runtime | |
| Web WebGPU EP and materialize im2col tensors of hundreds of MB, which overflows | |
| the 32-bit wasm heap at 1024x1024 (OrtRun std::bad_alloc on every EP). | |
| This export: | |
| * pins the source checkpoint by commit sha and loads it with | |
| transformers' trust_remote_code path (the model code ships in the repo); | |
| * replaces DeformableConv2d.forward with a numerically identical per-tap | |
| GridSample decomposition (one bilinear GridSample + 1x1 Conv per kernel | |
| tap, accumulated) - peak extra memory is one [1,C,H,W] tensor per tap | |
| instead of one [1,C,K,H,W] im2col blob, and the graph uses only ops the | |
| WebGPU EP implements natively (no GatherND, no ScatterND, no Clip); | |
| * traces with a static input shape (TorchScript exporter, opset 17, fp32, | |
| constant folding) so all Shape/Where dynamism folds away; | |
| * validates the patched module against torchvision.ops.deform_conv2d and the | |
| final ONNX against the unpatched PyTorch reference. | |
| Usage: python export_birefnet_lite.py [--size 1024] [--out birefnet_lite.onnx] | |
| """ | |
| import argparse | |
| import sys | |
| import numpy as np | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| REPO = 'ZhengPeng7/BiRefNet_lite' | |
| REVISION = '7838f1c3472f827cd8ce13ab5ccc2ce48077360f' # pinned 2026-07-18 | |
| def grid_sample_deform_forward(self, x): | |
| """Drop-in for DeformableConv2d.forward (birefnet.py) built on GridSample. | |
| Equivalent to torchvision.ops.deform_conv2d(input, offset, weight, bias, | |
| stride=1, padding=p, mask=modulator): for kernel tap k at (i, j) the op | |
| samples x at (y + i - p + dy_k, x + j - p + dx_k) bilinearly with zero | |
| padding, scales by the modulation mask, and accumulates W[:, :, i, j] as a | |
| 1x1 conv. Offsets are interleaved (dy, dx) per tap, row-major over the | |
| kernel, matching torchvision's layout. | |
| """ | |
| offset = self.offset_conv(x) # [B, 2K, H, W] | |
| modulator = 2.0 * torch.sigmoid(self.modulator_conv(x)) # [B, K, H, W] | |
| B, C, H, W = x.shape | |
| kh, kw = self.regular_conv.kernel_size | |
| pad = self.padding | |
| weight = self.regular_conv.weight # [Cout, Cin, kh, kw] | |
| ys = torch.arange(H, dtype=x.dtype, device=x.device) | |
| xs = torch.arange(W, dtype=x.dtype, device=x.device) | |
| base_y, base_x = torch.meshgrid(ys, xs, indexing='ij') # [H, W] constants | |
| # pre-normalized base grids (grid_sample align_corners=False convention); | |
| # keeping the per-tap shift on the dynamic offset lets constant folding | |
| # share ONE [H, W] grid pair per resolution instead of one per kernel tap | |
| norm_y = (2 * base_y + 1) / H - 1 | |
| norm_x = (2 * base_x + 1) / W - 1 | |
| out = None | |
| for k in range(kh * kw): | |
| i, j = divmod(k, kw) | |
| py = norm_y + (offset[:, 2 * k] + (i - pad)) * (2.0 / H) # [B, H, W] | |
| px = norm_x + (offset[:, 2 * k + 1] + (j - pad)) * (2.0 / W) | |
| grid = torch.stack([px, py], dim=-1) | |
| sampled = F.grid_sample(x, grid, mode='bilinear', padding_mode='zeros', | |
| align_corners=False) | |
| sampled = sampled * modulator[:, k:k + 1] | |
| contrib = F.conv2d(sampled, weight[:, :, i, j].unsqueeze(-1).unsqueeze(-1)) | |
| out = contrib if out is None else out + contrib | |
| if self.regular_conv.bias is not None: | |
| out = out + self.regular_conv.bias.view(1, -1, 1, 1) | |
| return out | |
| def verify_deform_patch(deform_cls): | |
| """Random-weight equivalence check: patched forward vs torchvision op.""" | |
| from torchvision.ops import deform_conv2d | |
| torch.manual_seed(0) | |
| worst = 0.0 | |
| for ksize, pad in [(1, 0), (3, 1), (7, 3)]: | |
| m = deform_cls(8, 16, kernel_size=ksize, padding=pad, bias=False) | |
| for p in m.parameters(): # zero-init offsets would hide layout bugs | |
| nn.init.normal_(p, std=0.3) | |
| x = torch.randn(1, 8, 20, 24) | |
| with torch.no_grad(): | |
| offset = m.offset_conv(x) | |
| modulator = 2.0 * torch.sigmoid(m.modulator_conv(x)) | |
| ref = deform_conv2d(x, offset, m.regular_conv.weight, | |
| m.regular_conv.bias, padding=pad, | |
| mask=modulator, stride=m.stride) | |
| got = grid_sample_deform_forward(m, x) | |
| worst = max(worst, (ref - got).abs().max().item()) | |
| return worst | |
| class LogitsWrapper(nn.Module): | |
| """BiRefNet eval-forward returns a list of multi-scale preds; the last one | |
| is the full-resolution logits map the demo pipeline consumes.""" | |
| def __init__(self, net): | |
| super().__init__() | |
| self.net = net | |
| def forward(self, input_image): | |
| return self.net(input_image)[-1] | |
| def main(): | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument('--size', type=int, default=1024) | |
| ap.add_argument('--out', default='birefnet_lite.onnx') | |
| args = ap.parse_args() | |
| from transformers import AutoModelForImageSegmentation | |
| torch.set_grad_enabled(False) | |
| model = AutoModelForImageSegmentation.from_pretrained( | |
| REPO, revision=REVISION, trust_remote_code=True) | |
| model.eval().float() | |
| birefnet_mod = sys.modules[type(model).__module__] | |
| # reference logits from the unpatched model (torchvision deform_conv2d) | |
| torch.manual_seed(1) | |
| probe = torch.rand(1, 3, args.size, args.size) | |
| ref_logits = LogitsWrapper(model)(probe) | |
| deform_err = verify_deform_patch(birefnet_mod.DeformableConv2d) | |
| print(f'deform patch vs torchvision: max abs dev {deform_err:.3e}') | |
| assert deform_err < 1e-4, 'GridSample decomposition diverges from torchvision' | |
| birefnet_mod.DeformableConv2d.forward = grid_sample_deform_forward | |
| patched_logits = LogitsWrapper(model)(probe) | |
| patch_err = (ref_logits - patched_logits).abs().max().item() | |
| print(f'patched model vs reference: max abs logits dev {patch_err:.3e}') | |
| assert patch_err < 1e-3 | |
| torch.onnx.export( | |
| LogitsWrapper(model), | |
| (probe,), | |
| args.out, | |
| input_names=['input_image'], | |
| output_names=['output_image'], | |
| opset_version=17, | |
| do_constant_folding=True, | |
| dynamo=False, | |
| ) | |
| # The TorchScript trace leaves the (all-constant, static-shape) Swin | |
| # attention-mask construction as live ScatterND/Where/Shape chains; | |
| # onnxslim folds them into initializers. | |
| # FusionGemm is skipped so the shared (deduped) MatMul weights are not | |
| # re-split into per-call-site transposed Gemm copies. | |
| import onnxslim | |
| onnxslim.slim(args.out, output_model=args.out, | |
| skip_fusion_patterns=['FusionGemm', 'FusionGemmMul', 'FusionGemmAdd']) | |
| # The backbone is traced twice (multi-scale 'cat' input), and the | |
| # TorchScript exporter emits one initializer per call site - every Swin | |
| # weight lands in the file twice (+116 MB). Dedupe by content hash. | |
| import hashlib | |
| import onnx | |
| m = onnx.load(args.out) | |
| canonical, rename = {}, {} | |
| keep = [] | |
| for init in m.graph.initializer: | |
| arr = onnx.numpy_helper.to_array(init) | |
| h = (str(arr.dtype), arr.shape, | |
| hashlib.sha256(arr.tobytes()).hexdigest()) | |
| if h in canonical: | |
| rename[init.name] = canonical[h] | |
| else: | |
| canonical[h] = init.name | |
| keep.append(init) | |
| del m.graph.initializer[:] | |
| m.graph.initializer.extend(keep) | |
| for node in m.graph.node: | |
| for idx, name in enumerate(node.input): | |
| if name in rename: | |
| node.input[idx] = rename[name] | |
| print(f'deduped {len(rename)} duplicate initializers') | |
| onnx.save(m, args.out) | |
| import onnx | |
| from collections import Counter | |
| m = onnx.load(args.out, load_external_data=False) | |
| ops = Counter(n.op_type for n in m.graph.node) | |
| print('ops:', sorted(ops.items(), key=lambda kv: -kv[1])) | |
| for banned in ('GatherND', 'ScatterND', 'Clip'): | |
| assert banned not in ops, f'{banned} in graph - not browser-clean' | |
| import onnxruntime as ort_rt | |
| sess = ort_rt.InferenceSession(args.out, providers=['CPUExecutionProvider']) | |
| onnx_logits = sess.run(None, {'input_image': probe.numpy()})[0] | |
| onnx_err = np.abs(ref_logits.numpy() - onnx_logits).max() | |
| alpha_err = np.abs(torch.sigmoid(ref_logits).numpy() | |
| - 1 / (1 + np.exp(-onnx_logits))).max() | |
| print(f'onnx vs torch reference: max abs logits dev {onnx_err:.3e}, ' | |
| f'max sigmoid dev {alpha_err:.3e}') | |
| if __name__ == '__main__': | |
| main() | |