"""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 [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)