JAX JIT模式下reshape报错:实现entmax-alpha遇类型错误求助
解决JAX中entmax实现的TypeError问题
这个错误的根源是在JAX的自动微分追踪过程中,直接访问被追踪张量的.shape属性并使用整数索引导致的。当value_and_grad触发微分追踪时,input变成了DynamicJaxprTracer对象,它的.shape不是普通的Python元组,无法用整数直接索引,因此抛出TypeError: tuple indices must be integers or slices, not DynamicJaxprTracer。
下面是具体的修复方案和修改后的完整代码:
关键修改点
- 替换直接访问
.shape的方式:用jnp.shape()获取张量形状(返回JAX数组,兼容追踪机制),而非直接访问.shape属性。 - 简化reshape操作:使用
-1让JAX自动推导维度,避免手动计算input.size / reduce_length(该操作在tracer上下文中也可能引发问题)。 - 修复
map_row中的形状访问:将z.shape[0]替换为JAX兼容的形状获取方式,确保静态/动态场景都能正常运行。
修改后的完整代码
import jax import jax.numpy as jnp from jax import custom_jvp from jax import jit from jax import lax from jax import vmap @jax.partial(jit, static_argnums=(2,)) def p_tau(z, tau, alpha=1.5): return jnp.clip((alpha - 1) * z - tau, a_min=0) ** (1 / (alpha - 1)) @jit def get_tau(tau, tau_max, tau_min, z_value): return lax.cond(z_value < 1, lambda _: (tau, tau_min), lambda _: (tau_max, tau), operand=None ) @jit def body(kwargs, x): tau_min = kwargs['tau_min'] tau_max = kwargs['tau_max'] z = kwargs['z'] alpha = kwargs['alpha'] tau = (tau_min + tau_max) / 2 z_value = p_tau(z, tau, alpha).sum() taus = get_tau(tau, tau_max, tau_min, z_value) tau_max, tau_min = taus[0], taus[1] return {'tau_min': tau_min, 'tau_max': tau_max, 'z': z, 'alpha': alpha}, None @jax.partial(jit, static_argnums=(1, 2,)) def map_row(z_input, alpha, T): z = (alpha - 1) * z_input # 用jnp.shape获取形状并转为对应dtype的数组,兼容tracer z_len = jnp.asarray(jnp.shape(z)[0], dtype=z.dtype) tau_min, tau_max = jnp.min(z) - 1, jnp.max(z) - z_len ** (1 - alpha) result, _ = lax.scan(body, {'tau_min': tau_min, 'tau_max': tau_max, 'z': z, 'alpha': alpha}, xs=None, length=T) tau = (result['tau_max'] + result['tau_min']) / 2 result = p_tau(z, tau, alpha) return result / result.sum() @jax.partial(custom_jvp, nondiff_argnums=(1, 2, 3,)) def entmax(input, axis=-1, alpha=1.5, T=10): # 用jnp.shape获取输入形状,避免直接访问tracer的.shape属性 input_shape = jnp.shape(input) reduce_length = input_shape[axis] input = jnp.swapaxes(input, -1, axis) # 用-1自动推导第一个维度,无需手动计算 input = input.reshape(-1, reduce_length) result = vmap(jax.partial(map_row, alpha=alpha, T=T), 0)(input) return jnp.swapaxes(result, -1, axis) @jax.partial(jit, static_argnums=(1, 2,)) def _entmax_jvp_impl(axis, alpha, T, primals, tangents): input = primals[0] Y = entmax(input, axis, alpha, T) gppr = Y ** (2 - alpha) grad_output = tangents[0] dX = grad_output * gppr q = dX.sum(axis=axis) / gppr.sum(axis=axis) q = jnp.expand_dims(q, axis=axis) dX -= q * gppr return Y, dX @entmax.defjvp def entmax_jvp(axis, alpha, T, primals, tangents): return _entmax_jvp_impl(axis, alpha, T, primals, tangents)
验证修复
运行你提供的测试代码:
import numpy as np from jax import value_and_grad input = jnp.array(np.random.randn(64, 10)) weight = jnp.array(np.random.randn(64, 10)) def toy(input, weight): return (weight*entmax(input, axis=-1, alpha=1.5, T=20)).sum() value_and_grad(toy)(input, weight)
现在应该可以正常返回损失值和梯度了。
额外说明
JAX的追踪机制对静态/动态操作的区分很严格:
- 静态操作(比如直接访问Python元组的索引)在tracer上下文中会失效,因为tracer的属性不是常规Python对象。
- 尽量使用JAX提供的函数(如
jnp.shape、jnp.reshape的-1语法)来处理形状,这些函数会自动兼容静态和动态场景。
内容的提问来源于stack exchange,提问作者RanWang
相关产品推荐
相关产品推荐

