Download scripts/export.py from swiftail/BiRefNet_lite-onnx-coreml: direct link, hf CLI and curl.
- Browser
- Download file 4.68 kB
-
https://huggingface.co/swiftail/BiRefNet_lite-onnx-coreml/resolve/main/scripts/export.py
- Command line
-
hf download hf://swiftail/BiRefNet_lite-onnx-coreml/scripts/export.py
-
curl -L -o export.py https://huggingface.co/swiftail/BiRefNet_lite-onnx-coreml/resolve/main/scripts/export.py
4.68 kB
| """Re-export BiRefNet (HF remote code) to ONNX with CoreML/memory-friendly rewrites: | |
| 1. DeformableConv2d: torchvision deform_conv2d -> per-tap GridSample + 1x1 conv accumulation (no 784MB im2col tensors) | |
| 2. Swin window_partition/window_reverse: 6-D view/permute -> 5-D (batch=1), CoreML rank limit is 5 | |
| usage: export.py <hf_repo> <out.onnx> [size] [--check]""" | |
| import sys, os, time | |
| os.environ.setdefault("HF_HOME", "/tmp/mixer-bg-spike/hf") | |
| import torch, torch.nn.functional as F | |
| from transformers import AutoModelForImageSegmentation | |
| repo, out = sys.argv[1], sys.argv[2] | |
| size = int(sys.argv[3]) if len(sys.argv) > 3 and sys.argv[3].isdigit() else 1024 | |
| check = "--check" in sys.argv | |
| torch.set_grad_enabled(False) | |
| model = AutoModelForImageSegmentation.from_pretrained(repo, trust_remote_code=True).eval().float() | |
| mod = sys.modules[type(model).__module__] | |
| def deform_forward_gs(self, x): | |
| offset = self.offset_conv(x) | |
| mask = 2. * torch.sigmoid(self.modulator_conv(x)) | |
| w = self.regular_conv.weight; O, C, kh, kw = w.shape | |
| sh, sw = self.stride | |
| ph, pw = (self.padding, self.padding) if isinstance(self.padding, int) else self.padding | |
| B, _, H, W = x.shape; Ho, Wo = offset.shape[2], offset.shape[3] | |
| ys = (torch.arange(Ho, dtype=x.dtype) * sh - ph).view(1, Ho, 1) | |
| xs = (torch.arange(Wo, dtype=x.dtype) * sw - pw).view(1, 1, Wo) | |
| out = None | |
| for i in range(kh): | |
| for j in range(kw): | |
| k = i * kw + j | |
| py = ys + i + offset[:, 2 * k] # [B,Ho,Wo] | |
| px = xs + j + offset[:, 2 * k + 1] | |
| grid = torch.stack((2 * px / (W - 1) - 1, 2 * py / (H - 1) - 1), dim=-1) | |
| s = F.grid_sample(x, grid, mode="bilinear", padding_mode="zeros", align_corners=True) | |
| s = s * mask[:, k:k + 1] | |
| t = F.conv2d(s, w[:, :, i:i + 1, j:j + 1]) | |
| out = t if out is None else out + t | |
| if self.regular_conv.bias is not None: | |
| out = out + self.regular_conv.bias.view(1, -1, 1, 1) | |
| return out | |
| def window_partition5(x, ws): | |
| B, H, W, C = x.shape # B == 1 for export | |
| x = x.view(H // ws, ws, W // ws, ws, C).permute(0, 2, 1, 3, 4).contiguous() | |
| return x.view(-1, ws, ws, C) | |
| def window_reverse5(windows, ws, H, W): | |
| C = windows.shape[-1] | |
| x = windows.view(H // ws, W // ws, ws, ws, C).permute(0, 2, 1, 3, 4).contiguous() | |
| return x.view(1, H, W, C) | |
| class Wrap(torch.nn.Module): | |
| def __init__(s, m): super().__init__(); s.m = m | |
| def forward(s, x): return s.m(x)[-1] | |
| x = torch.randn(1, 3, size, size) | |
| if check: | |
| xs = torch.randn(1, 3, 512, 512) | |
| t = time.time(); ref = Wrap(model)(xs); print("ref", time.time() - t) | |
| mod.DeformableConv2d.forward = deform_forward_gs | |
| mod.window_partition = window_partition5 | |
| mod.window_reverse = window_reverse5 | |
| # 3. qkv[0..2] (Gather, rejected by CoreML EP) -> unbind (Split+Squeeze) | |
| import inspect, textwrap | |
| src = textwrap.dedent(inspect.getsource(mod.WindowAttention.forward)).replace("qkv[0], qkv[1], qkv[2]", "qkv.unbind(0)") | |
| assert "unbind" in src; ns = {}; exec(src, mod.__dict__, ns); mod.WindowAttention.forward = ns["forward"] | |
| # 4. image2patches 'b c (hg h) (wg w) -> b (c hg wg) h w' as 5-D (b == 1) | |
| _orig_i2p = mod.image2patches | |
| def image2patches5(image, grid_h=2, grid_w=2, patch_ref=None, transformation=None): | |
| if transformation != 'b c (hg h) (wg w) -> b (c hg wg) h w': return _orig_i2p(image, grid_h, grid_w, patch_ref, transformation) | |
| if patch_ref is not None: | |
| grid_h, grid_w = image.shape[-2] // patch_ref.shape[-2], image.shape[-1] // patch_ref.shape[-1] | |
| _, c, H, W = image.shape; h, w = H // grid_h, W // grid_w | |
| return image.view(c, grid_h, h, grid_w, w).permute(0, 1, 3, 2, 4).reshape(1, c * grid_h * grid_w, h, w) | |
| mod.image2patches = image2patches5 | |
| if check: | |
| t = time.time(); new = Wrap(model)(xs); print("new", time.time() - t) | |
| print("max abs logit diff", (ref - new).abs().max().item(), "max sigmoid diff", (ref.sigmoid() - new.sigmoid()).abs().max().item()) | |
| sys.exit(0) | |
| t = time.time() | |
| if os.environ.get("DYNAMO"): | |
| prog = torch.onnx.export(Wrap(model), (x,), None, input_names=["input_image"], output_names=["output_image"], opset_version=18, dynamo=True, optimize=os.environ.get("OPT", "1") == "1") | |
| prog.save(out, external_data=bool(os.environ.get("EXTDATA"))); print("exported(dynamo)", out, time.time() - t); sys.exit(0) | |
| torch.onnx.export(Wrap(model), x, out, input_names=["input_image"], output_names=["output_image"], | |
| opset_version=17, do_constant_folding=False, dynamo=False, | |
| dynamic_axes={"input_image": {2: "h", 3: "w"}, "output_image": {2: "h", 3: "w"}}) | |
| print("exported", out, time.time() - t) | |