| """Repro of prior spike's crasher: 2 cross-dependent states + attention. |
| State A (mem_bank) is READ -> attention -> result WRITTEN to state B (ptr_bank). |
| Prior result on torch 2.13 + coremltools 9.0: KeyError in optimize_state.py:93. |
| Now testing under torch 2.7.0 (coremltools' tested ceiling). |
| """ |
|
|
| import numpy as np |
| import torch |
| import torch.nn as nn |
| import coremltools as ct |
|
|
| C, T, S = 256, 7, 64 |
|
|
|
|
| class CrossStateModel(nn.Module): |
| def __init__(self): |
| super().__init__() |
| self.register_buffer("mem_bank", torch.zeros(T, S, C)) |
| self.register_buffer("ptr_bank", torch.zeros(T, C)) |
| self.attn = nn.MultiheadAttention(C, 4, batch_first=True) |
| self.enc = nn.Linear(C, C) |
|
|
| def forward(self, x): |
| |
| mem = self.mem_bank.reshape(1, T * S, C) |
| ptrs = self.ptr_bank.reshape(1, T, C) |
| kv = torch.cat([mem, ptrs], dim=1) |
| y, _ = self.attn(x, kv, kv) |
| |
| new_ptr = y.mean(dim=1) |
| ptr_updated = torch.cat([self.ptr_bank[1:], new_ptr], dim=0) |
| self.ptr_bank.copy_(ptr_updated) |
| |
| new_mem = self.enc(y) |
| mem_updated = torch.cat([self.mem_bank[1:], new_mem], dim=0) |
| self.mem_bank.copy_(mem_updated) |
| return y |
|
|
|
|
| def main(): |
| torch.manual_seed(0) |
| m = CrossStateModel().eval() |
| m.requires_grad_(False) |
| x = torch.randn(1, S, C) |
|
|
| with torch.no_grad(): |
| ep = torch.export.export(m, (x,)) |
| ep = ep.run_decompositions({}) |
| print("torch.export OK") |
|
|
| mlmodel = ct.convert( |
| ep, |
| minimum_deployment_target=ct.target.iOS18, |
| compute_units=ct.ComputeUnit.CPU_ONLY, |
| ) |
| print("coremltools convert OK") |
|
|
| state = mlmodel.make_state() |
| out1 = mlmodel.predict({"x": x.numpy()}, state=state) |
| out2 = mlmodel.predict({"x": x.numpy()}, state=state) |
| k = list(out1)[0] |
| print("predict OK; outputs differ across calls (state evolving):", |
| not np.allclose(out1[k], out2[k])) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|