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.
x[i]andset_subtensor(x[i], y)with a symbolic integerifail to compile on the MLX backend.mlx_funcify_Subtensorandmlx_funcify_IncSubtensorcallint()on every integer index input, and undermx.compilethe index is a tracedmx.array, soint()raises. mlx accepts an integermx.arrayas an index for both reads and writes, and the same graph compiles when the index is a one-element vector, becauseAdvancedSubtensorpasses the array through.The
int()coercion came in with #2240 for slice bounds. Restricting it to slice bounds and passing an integermx.arrayindex straight tox[index]would cover both ops.