indic-transcribe-flex-coreml / export_coreml.py
phequals's picture
Publish validated experimental Flex CoreML port with mixed-script Swift runtime
18039eb verified
Raw
History Blame Contribute Delete
4.16 kB
"""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()