File size: 4,677 Bytes
6383c88
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
"""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)