Skip to content

Add support for _expected_input_spec in MLX Lambda to match JAX implementation - #17

Open
lancelotblanchard wants to merge 1 commit into
google:mainfrom
lancelotblanchard:mlx-lambda
Open

lancelotblanchard wants to merge 1 commit into
google:mainfrom
lancelotblanchard:mlx-lambda

Conversation

@lancelotblanchard

Copy link
Copy Markdown

This PR implements expected_input_spec support in the MLX Lambda layer to match the JAX implementation. Previously, expected_input_spec was accepted in Lambda.Config for compatibility but was completely ignored by the MLX runtime.

The MLX dynamic shape probing (_probe_output) currently supports a input_dtype argument, but it is ignored by get_output_shape, which assumes mx.float32. The input_dtype argument is now overwritten by the config's expected_input_spec.

@JulianSlzr JulianSlzr self-assigned this Oct 6, 2026
@JulianSlzr
JulianSlzr self-requested a review October 6, 2026 20:18

@JulianSlzr JulianSlzr left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks Lancelot! Will merge this internally with attribution and close once it's reflected in the public version.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants