You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.11 08:55:43