如何在Google JAX中对自定义卷积包装器正确使用grad梯度计算功能
问题原因
- 传参逻辑错误:JAX的
grad默认对函数的第一个输入参数求梯度,你调用grad(conv)(y, k)时,将输出特征维度为64的特征图y作为输入传入conv函数,而卷积核k的输入特征维度仅为3,二者维度不匹配触发了卷积算子的参数校验错误,就是你看到的报错信息。 - 梯度计算逻辑不合法:
grad要求目标函数的输出必须是标量才能计算梯度,你直接对输出为4维张量的conv函数求梯度,即使维度匹配也会报错。
解决方案
首先构造一个输出为标量的损失函数,再根据你需要求导的参数指定argnums参数调用grad,示例代码如下:
init = normal() rng = PRNGKey(42) x = init(rng, [128, 3, 224, 224]) k = init(rng, [64, 3, 3, 3]) # 定义标量损失,比如对卷积输出的所有元素求和 def loss(x, weight): return conv(x, weight).sum() # 对卷积核weight(也就是函数的第二个输入参数,索引为1)求梯度 k_grad = grad(loss, argnums=1)(x, k) # 如果需要同时对输入x和卷积核k求梯度,写法如下 x_grad, k_grad = grad(loss, argnums=(0, 1))(x, k)
可选优化(避免后续使用出错)
你当前的conv包装器存在两个小问题,可以提前修正:
- dilation参数赋值错误:普通卷积的空洞率是作用在卷积核上的,对应
lax.conv_general_dilated的rhs_dilation参数,lhs_dilation仅用于转置卷积场景,赋值错误会导致带空洞的卷积计算结果异常。 - 不支持分组卷积:你写死了
feature_group_count=1,没有用到函数传入的groups参数,无法实现分组卷积功能。
修正后的conv函数返回部分代码如下:
return lax.conv_general_dilated( lhs=input, rhs=weight, window_strides=stride, padding=padding, rhs_dilation=dilation, dimension_numbers=torch_dims[n], feature_group_count=groups, batch_group_count=1, precision=None, preferred_element_type=None )
内容的提问来源于stack exchange,提问作者TIM
相关产品推荐
相关产品推荐

