swiftail's picture
CoreML-friendly ONNX re-export of BiRefNet_lite (ZhengPeng7)
6383c88 verified
Raw History Blame Contribute Delete
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)