| """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() |
|
|