Skip to content

Repository files navigation

stablehlo-coreml

Convert StableHLO models into Apple Core ML format.

StableHLO is the portability layer used by ML frameworks like JAX and PyTorch. This library converts StableHLO programs into Apple's Core ML format via coremltools, enabling deployment on Apple hardware (iOS, macOS, etc.).

Installation

pip install stablehlo-coreml

Requires Python 3.10–3.13 and targets iOS/macOS 18+.

Supported Frameworks

Models can be exported from any framework that produces StableHLO:

  • JAX / Flax / Equinox — via jax.export
  • PyTorch — via torchax to trace the model into JAX, then jax.export to StableHLO

The test suite validates against a broad set of models, including full HuggingFace Transformers such as TinyLlama, T5, DistilBERT, GPT-2, BERT, and Whisper, as well as vision models like ResNet, EfficientNet, ViT, ConvNeXt, and more.

For a real-world example, see gemma-coreml-chat, which exports Google's Gemma 4 model to Core ML using this library.

Converting a Model

To convert a StableHLO module:

import coremltools as ct
from stablehlo_coreml import StateSpec, build_pass_pipeline, convert

mil_program = convert(hlo_module, minimum_deployment_target=ct.target.iOS18)
cml_model = ct.convert(
    mil_program,
    source="milinternal",
    minimum_deployment_target=ct.target.iOS18,
    pass_pipeline=build_pass_pipeline(),
)

build_pass_pipeline() returns a fresh ct.PassPipeline built from ct.PassPipeline.DEFAULT with the stablehlo-coreml graph passes inserted at the right places. Pass your own base pipeline to customise it, e.g. build_pass_pipeline(my_pipeline). The pre-built stablehlo_coreml.DEFAULT_HLO_PIPELINE is the same thing, constructed once at import time — copy it before mutating it.

Obtaining a StableHLO Module from JAX

import jax
from jax._src.lib.mlir import ir
from jax._src.interpreters import mlir as jax_mlir
from jax.export import export

import jax.numpy as jnp

def jax_function(a, b):
    return jnp.einsum("ij,jk -> ik", a, b)

context = jax_mlir.make_ir_context()
input_shapes = (jnp.zeros((2, 4)), jnp.zeros((4, 3)))
jax_exported = export(jax.jit(jax_function))(*input_shapes)
hlo_module = ir.Module.parse(jax_exported.mlir_module(), context=context)

For the JAX example to work, you will additionally need to install absl-py and flatbuffers as dependencies.

Stateful models

Core ML can keep tensors across model invocations as state instead of passing them in and out every time. Mark those tensors when converting by mapping each state input to the output that holds its updated value:

def step(cache, x):
    new_cache = cache + x
    return new_cache * x, new_cache

mil_program = convert(
    hlo_module,
    minimum_deployment_target=ct.target.iOS18,
    states={
        "main": {
            "cache": StateSpec(output=1),
        },
    },
)
cml_model = ct.convert(
    mil_program,
    source="milinternal",
    minimum_deployment_target=ct.target.iOS18,
    pass_pipeline=build_pass_pipeline(),
)

state = cml_model.make_state()
y = cml_model.predict({"x": x}, state=state)
# `cache` is updated in place; inspect or reset it with
# state.read_state(...) / state.write_state(...)

Inner keys may be argument indices or names, and StateSpec.output an output index or JAX result name. Use output=None for read-only state and name=... to override the Core ML state name. A flat {input: output} mapping works for single-function modules.

State tensors must have a static shape and a floating-point dtype (stored as fp16). They are removed from the model's inputs, and the outputs that update them are dropped.

See tests/test_stateful.py for multi-step examples.

Dynamic / symbolic shapes

JAX models exported with symbolic dimensions are supported. Symbolic dims flow through GetDimensionSizeOp, DynamicBroadcastInDimOp, DynamicIotaOp, and shape-assertion CustomCallOps automatically, producing CoreML models with flexible inputs.

import jax
import jax.numpy as jnp
from jax.export import export, symbolic_shape

jax_exported = export(jax.jit(jax_function))(
    jax.ShapeDtypeStruct(symbolic_shape("batch, 4"), jnp.float32),
    jax.ShapeDtypeStruct((4, 3), jnp.float32),
)

When converting to a CoreML model, specify RangeDim for each symbolic dimension so the model accepts a range of sizes at inference time:

cml_model = ct.convert(
    mil_program,
    source="milinternal",
    minimum_deployment_target=ct.target.iOS18,
    pass_pipeline=build_pass_pipeline(),
    inputs=[
        ct.TensorType(name="arg0", shape=(ct.RangeDim(1, 2048, 1), 4)),
        ct.TensorType(name="arg1", shape=(4, 3)),
    ],
)

See tests/test_symbolic_shapes.py for symbolic matmul, batched einsum, and multi-axis patterns (for example transformer-style projections).

Examples in the test suite

The tests/ directory has end-to-end export and conversion examples:

Development

  • coremltools supports up to Python 3.13. Do not run hatch with a newer version. Can be controlled using e.g. export HATCH_PYTHON=python3.13
  • Run tests using hatch run test:pytest tests

About

Convert StableHLO models into Apple Core ML format

Resources

Stars

22 stars

Watchers

1 watching

Forks

Releases

Packages

Used by

Contributors

Languages