"""Export on Linux; predictions and performance validation require macOS.""" import argparse import json import sys from pathlib import Path import numpy as np import torch import coremltools as ct from coreml_port import Encoder, CrossKV, Decoder, causal_mask def main(): p = argparse.ArgumentParser() p.add_argument('--root', default='/home/indicASR/bodhan-core') p.add_argument('--component', choices=['encoder', 'cross', 'decoder'], required=True) a = p.parse_args() root = Path(a.root) sys.path.insert(0, str(root / 'upstream')) from modeling_indic_canary import IndicCanaryForConditionalGeneration torch.set_num_threads(8) model = IndicCanaryForConditionalGeneration.from_pretrained(str(root / 'upstream'), dtype=torch.float32).eval() fixture = np.load(root / 'baseline/fixture_0.npz') enc = torch.from_numpy(fixture['encoder']) out = root / 'coreml' out.mkdir(exist_ok=True) t = ct.RangeDim(1, 376, default=enc.shape[1]) if a.component == 'encoder': wrapper = Encoder(model).eval() examples = (torch.from_numpy(fixture['features']), torch.from_numpy(fixture['mask']).int()) mel_t = ct.RangeDim(101, 3001, default=examples[0].shape[-1]) inputs = [ct.TensorType(name='features', shape=(1,128,mel_t)), ct.TensorType(name='mask', shape=(1,mel_t), dtype=np.int32)] outputs = [ct.TensorType(name='encoder'), ct.TensorType(name='lengths', dtype=np.int32)] states = [] elif a.component == 'cross': wrapper = CrossKV(model).eval() examples = (enc,) inputs = [ct.TensorType(name='encoder', shape=(1,t,1024))] outputs = [ct.TensorType(name='cross_kv')] states = [] else: wrapper = Decoder(model).eval() with torch.no_grad(): cross = CrossKV(model)(enc) ids = torch.from_numpy(fixture['prompt']).int() examples = (ids, torch.tensor([0],dtype=torch.int32), causal_mask(0,ids.shape[1],512), cross, torch.zeros(1,1,1,enc.shape[1])) q = ct.RangeDim(1, 16, default=ids.shape[1]) inputs = [ct.TensorType(name='ids',shape=(1,q),dtype=np.int32), ct.TensorType(name='position',shape=(1,),dtype=np.int32), ct.TensorType(name='self_mask',shape=(1,1,q,512)), ct.TensorType(name='cross_kv',shape=(48,1,8,t,128)), ct.TensorType(name='cross_mask',shape=(1,1,1,t))] outputs = [ct.TensorType(name='logits')] states = [ct.StateType(wrapped_type=ct.TensorType(shape=b.shape,dtype=np.float16),name=n) for n,b in wrapper.named_buffers() if n.startswith(('k_','v_'))] with torch.no_grad(): logits = wrapper(*examples).numpy() delta = np.abs(logits-fixture['logits']) parity = {'max_abs':float(delta.max()),'mean_abs':float(delta.mean()), 'last_token_equal':bool(logits[0,-1].argmax()==fixture['logits'][0,-1].argmax())} (out/'decoder_eager_parity.json').write_text(json.dumps(parity,indent=2)) print(parity,flush=True) if not parity['last_token_equal'] or parity['max_abs'] > 0.1: raise RuntimeError('Decoder adapter does not match source; refusing export') wrapper.reset() with torch.no_grad(): traced = torch.jit.trace(wrapper, examples, check_trace=False) if isinstance(wrapper, Decoder): wrapper.reset() print('Converting', a.component, flush=True) ml = ct.convert(traced, inputs=inputs, outputs=outputs, states=states, minimum_deployment_target=ct.target.macOS15, compute_precision=ct.precision.FLOAT16, skip_model_load=True) ml.user_defined_metadata['upstream_revision'] = (root/'upstream/REVISION').read_text().strip() if (root/'upstream/REVISION').exists() else '4d29eeb7a0990de4a8febf9a4d5a9c5c61134a0d' ml.user_defined_metadata['upstream_model'] = 'bodhan-ai/indic-transcribe-flex' ml.save(str(out/f'{a.component}.mlpackage')) print('Saved',out/f'{a.component}.mlpackage',flush=True) if __name__ == '__main__': main()