import numpy as np, torch, torch.nn as nn import coremltools as ct torch.manual_seed(0) class MinState(nn.Module): def __init__(self): super().__init__() self.register_buffer("bank", torch.zeros(1, 64, 256)) # all dims 32-friendly self.enc = nn.Linear(256, 256) def forward(self, x): # (1,64,256) y = self.enc(x + self.bank) self.bank.copy_(y) # whole-buffer update, no slice/cat return y m = MinState().eval(); m.requires_grad_(False) x = torch.randn(1, 64, 256) with torch.no_grad(): ep = torch.export.export(m, (x,)); ep = ep.run_decompositions({}) ml = ct.convert(ep, minimum_deployment_target=ct.target.iOS18, compute_units=ct.ComputeUnit.CPU_AND_NE) st = ml.make_state() o = ml.predict({"x": x.numpy()}, state=st) print("minimal aligned whole-buffer state on CPU_AND_NE: predict OK")