如何用flax.linen.checkpoint的static_argnums处理__call__布尔参数
Flax模块含布尔参数时的梯度Checkpointing正确实现
问题背景
子类化flax.linen.Module时,通过__call__方法的deterministic布尔参数区分训练/推理模式(训练启用Dropout等随机操作,推理使用确定性行为)。尝试用flax.linen.checkpoint实现梯度Checkpointing以降低GPU内存占用时,遇到静态参数处理错误:
- 直接装饰模块类触发
jax.errors.ConcretizationTypeError; - 指定
static_argnums=2触发ValueError提示索引过大; - 其他
static_argnums取值无法解决问题。
原测试代码:
import jax import jax.numpy as jnp import flax.linen as nn class MLPWithDropout(nn.Module): @nn.compact def __call__(self, x, deterministic=False): x = nn.Dense(128)(x) x = nn.Dropout(rate=0.5, deterministic=deterministic)(x) x = nn.relu(x) x = nn.Dense(1)(x) return x # 以下方式均失败 # CheckpointedMLPWithDropout = nn.checkpoint(MLPWithDropout) # CheckpointedMLPWithDropout = nn.checkpoint(MLPWithDropout, static_argnums=2) # CheckpointedMLPWithDropout = nn.checkpoint(MLPWithDropout, static_argnums=-1) # CheckpointedMLPWithDropout = nn.checkpoint(MLPWithDropout, static_argnums=1) model = CheckpointedMLPWithDropout() x = jnp.ones((1, 16)) rng = jax.random.PRNGKey(0) vars = model.init(rng, x, deterministic=True) print(model.apply(vars, x, deterministic=False, rngs={'dropout': rng}))
解决方案
将nn.checkpoint直接装饰在模块的__call__方法上,并通过static_argnums指定deterministic参数的索引(相对于self之外的参数)。
修改后可正常运行的代码:
import jax import jax.numpy as jnp import flax.linen as nn class MLPWithDropout(nn.Module): @nn.compact @nn.checkpoint(static_argnums=1) # deterministic是__call__的第2个参数(x为第1个,索引0) def __call__(self, x, deterministic=False): x = nn.Dense(128)(x) x = nn.Dropout(rate=0.5, deterministic=deterministic)(x) x = nn.relu(x) x = nn.Dense(1)(x) return x model = MLPWithDropout() x = jnp.ones((1, 16)) rng = jax.random.PRNGKey(0) vars = model.init(rng, x, deterministic=True) print(model.apply(vars, x, deterministic=False, rngs={'dropout': rng}))
原理说明
- 直接用
nn.checkpoint装饰模块类时,实际包装的是模块的apply函数,其参数顺序为(variables, *args, **kwargs),此时deterministic的索引与__call__方法中的索引不一致,容易导致参数索引错误; - 装饰
__call__方法时,static_argnums对应的是排除self后的方法参数索引:x对应索引0,deterministic对应索引1,标记该参数为静态后,JAX会跳过对它的追踪,避免布尔值的 concretization 错误。
内容的提问来源于stack exchange,提问作者Ian Holmes
相关产品推荐
相关产品推荐

