Skip to content

MLX: Subtensor and IncSubtensor fail on a symbolic integer index under mx.compile #2422

Description

@jessegrabowski

x[i] and set_subtensor(x[i], y) with a symbolic integer i fail to compile on the MLX backend. mlx_funcify_Subtensor and mlx_funcify_IncSubtensor call int() on every integer index input, and under mx.compile the index is a traced mx.array, so int() raises. mlx accepts an integer mx.array as an index for both reads and writes, and the same graph compiles when the index is a one-element vector, because AdvancedSubtensor passes the array through.

import numpy as np
import pytensor
import pytensor.tensor as pt

x = pt.matrix("x")
i = pt.iscalar("i")
fn = pytensor.function([x, i], x[i], mode="MLX")
print(fn(np.eye(3, dtype="float32"), 1))  # ValueError: [eval] Attempting to eval an array during function transformations
# workaround: index with a one-element vector, x[i[None]][0], which lowers to AdvancedSubtensor

The int() coercion came in with #2240 for slice bounds. Restricting it to slice bounds and passing an integer mx.array index straight to x[index] would cover both ops.

Activity

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

Metadata

Metadata

Assignees

No one assigned

    Labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions