PyTorch GCN层间输出比特精度截断遇inplace操作报错解决方案
PyTorch GCN比特截断量化原位操作梯度报错修复
问题背景
- 任务:基于PyTorch在Cora数据集上训练2层GCN模型完成节点分类
- 自定义精度截断逻辑:每层输出float32张量后,先转换为32位二进制表示,仅保留指定数量比特位,再将截断后的二进制值转回float32,作为下一层输入
- 运行时抛出梯度计算错误,报错信息如下:
RuntimeError: one of the variables needed for gradient computation has been modified by an inplace operation: [torch.FloatTensor [2708, 7, 1]], which is output 0 of PowBackward1, is at version 2; expected version 0 instead. Hint: the backtrace further above shows the operation that failed to compute its gradient. The variable in question was changed in there or anywhere later. Good luck!
原始问题代码
def forward(self, data): torch.autograd.set_detect_anomaly(True) x, edge_index = data.x, data.edge_index x = self.conv1(x, edge_index) bit_x = float2bit(x) float_x = bit2float(bit_x) x = torch.sigmoid(float_x) bit_x = float2bit(x) bit_x[super_node,:,2:] = 0 float_x = bit2float(bit_x) x = self.conv2(float_x, edge_index) bit_x = float2bit(x) float_x = bit2float(bit_x) return F.log_softmax(float_x, dim=1) def bit2float(b, num_e_bits=8, num_m_bits=23, bias=127.): #b = bit.clone().detach() """Turn input tensor into float. Args: b : binary tensor. The last dimension of this tensor should be the the one the binary is at. num_e_bits : Number of exponent bits. Default: 8. num_m_bits : Number of mantissa bits. Default: 23. bias : Exponent bias/ zero offset. Default: 127. Returns: Tensor: Float tensor. Reduces last dimension. """ expected_last_dim = num_m_bits + num_e_bits + 1 assert b.shape[-1] == expected_last_dim, "Binary tensors last dimension " \ "should be {}, not {}.".format( expected_last_dim, b.shape[-1]) # check if we got the right type dtype = torch.float32 if expected_last_dim > 32: dtype = torch.float64 if expected_last_dim > 64: warnings.warn("pytorch can not process floats larger than 64 bits, keep" " this in mind. Your result will be not exact.") s = torch.index_select(b, -1, torch.arange(0, 1)) e = torch.index_select(b, -1, torch.arange(1, 1 + num_e_bits)) m = torch.index_select(b, -1, torch.arange(1 + num_e_bits, 1 + num_e_bits + num_m_bits)) # SIGN BIT out = ((-1) ** s).squeeze(-1).type(dtype) # EXPONENT BIT exponents = -torch.arange(-(num_e_bits - 1.), 1.) exponents = exponents.repeat(b.shape[:-1] + (1,)) e_decimal = torch.sum(e * 2 ** exponents, dim=-1) - bias out *= 2 ** e_decimal # MANTISSA matissa = (torch.Tensor([2.]) ** ( -torch.arange(1., num_m_bits + 1.))).repeat( m.shape[:-1] + (1,)) out *= 1. + torch.sum(m * matissa, dim=-1) return out def float2bit(f, num_e_bits=8, num_m_bits=23, bias=127., dtype=torch.float32): #f = float.clone().detach() """Turn input tensor into binary. Args: f : float tensor. num_e_bits : Number of exponent bits. Default: 8. num_m_bits : Number of mantissa bits. Default: 23. bias : Exponent bias/ zero offset. Default: 127. dtype : This is the actual type of the tensor that is going to be returned. Default: torch.float32. Returns: Tensor: Binary tensor. Adds last dimension to original tensor for bits. """ ## SIGN BIT s = torch.sign(f) f = f * s # turn sign into sign-bit s = (s * (-1) + 1.) * 0.5 s = s.unsqueeze(-1) ## EXPONENT BIT e_scientific = torch.floor(torch.log2(f)) e_decimal = e_scientific + bias e = integer2bit(e_decimal, num_bits=num_e_bits) ## MANTISSA m1 = integer2bit(f - f % 1, num_bits=num_e_bits) m2 = remainder2bit(f % 1, num_bits=bias) m = torch.cat([m1, m2], dim=-1) dtype = f.type() idx = torch.arange(num_m_bits).unsqueeze(0).type(dtype) \ + (8. - e_scientific).unsqueeze(-1) idx = idx.long() m = torch.gather(m, dim=-1, index=idx) return torch.cat([s, e, m], dim=-1).type(dtype)
错误根因
- 核心触发点是
bit_x[super_node,:,2:] = 0这行代码,属于原位操作(inplace operation):直接修改了存在于PyTorch自动求导计算图中的张量值,反向传播时需要用到该张量前向传播时的原始值做梯度计算,但此时张量已经被修改、版本号不匹配,因此抛出错误。 - 辅助函数存在隐藏隐患:
bit2float、float2bit中用torch.arange、torch.Tensor([2.])生成的常量张量默认创建在CPU上,若模型迁移到GPU训练会触发设备不匹配报错;且生成常量时未对齐输入张量的数据类型,可能引发精度问题。 - 冗余性能损耗:
torch.autograd.set_detect_anomaly(True)是调试用的异常检测开关,放在forward函数中每次前向传播都会重复开启,会大幅降低训练速度,不需要在forward里调用。
修复方案
- 替换所有原位修改操作:将直接对
bit_x的索引赋值改为通过掩码生成新张量,不修改原计算图节点的值 - 对齐常量张量的设备与数据类型:所有函数内部生成的索引、权重常量,都和输入张量绑定相同的device与dtype
- 移除forward中的异常检测开关,将其放到训练脚本最外层仅调用一次
- 对输入到转换函数的张量先做clone(若需要修改输入值的场景),避免意外影响上游计算节点
修复后核心代码
# 训练脚本最外层开一次异常检测即可,不要放forward里 # torch.autograd.set_detect_anomaly(True) def forward(self, data): x, edge_index = data.x, data.edge_index x = self.conv1(x, edge_index) bit_x = float2bit(x) float_x = bit2float(bit_x) x = torch.sigmoid(float_x) bit_x = float2bit(x) # 替换原原位赋值操作,用掩码生成新张量 mask = torch.ones_like(bit_x) mask[super_node, :, 2:] = 0 bit_x = bit_x * mask float_x = bit2float(bit_x) x = self.conv2(float_x, edge_index) bit_x = float2bit(x) float_x = bit2float(bit_x) return F.log_softmax(float_x, dim=1) def bit2float(b, num_e_bits=8, num_m_bits=23, bias=127.): expected_last_dim = num_m_bits + num_e_bits + 1 assert b.shape[-1] == expected_last_dim, f"Binary tensors last dimension should be {expected_last_dim}, not {b.shape[-1]}." dtype = torch.float32 if expected_last_dim > 32: dtype = torch.float64 if expected_last_dim > 64: warnings.warn("pytorch can not process floats larger than 64 bits, result will be not exact.") # 所有生成的索引、常量对齐输入张量的设备和dtype s_idx = torch.arange(0, 1, device=b.device, dtype=torch.long) e_idx = torch.arange(1, 1 + num_e_bits, device=b.device, dtype=torch.long) m_idx = torch.arange(1 + num_e_bits, 1 + num_e_bits + num_m_bits, device=b.device, dtype=torch.long) s = torch.index_select(b, -1, s_idx) e = torch.index_select(b, -1, e_idx) m = torch.index_select(b, -1, m_idx) # SIGN BIT out = ((-1) ** s).squeeze(-1).to(dtype) # EXPONENT BIT exponents = -torch.arange(-(num_e_bits - 1.), 1., device=b.device, dtype=dtype) exponents = exponents.repeat(b.shape[:-1] + (1,)) e_decimal = torch.sum(e * (2 ** exponents), dim=-1) - bias out *= 2 ** e_decimal # MANTISSA matissa = (torch.tensor(2., device=b.device, dtype=dtype) ** (-torch.arange(1., num_m_bits + 1., device=b.device, dtype=dtype))).repeat(m.shape[:-1] + (1,)) out *= 1. + torch.sum(m * matissa, dim=-1) return out def float2bit(f, num_e_bits=8, num_m_bits=23, bias=127.): # 先clone避免修改原输入张量 f = f.clone() ## SIGN BIT s = torch.sign(f) f = f * s s = (s * (-1) + 1.) * 0.5 s = s.unsqueeze(-1) ## EXPONENT BIT e_scientific = torch.floor(torch.log2(f)) e_decimal = e_scientific + bias e = integer2bit(e_decimal, num_bits=num_e_bits) ## MANTISSA m1 = integer2bit(f - f % 1, num_bits=num_e_bits) m2 = remainder2bit(f % 1, num_bits=bias) m = torch.cat([m1, m2], dim=-1) dtype = f.dtype # 生成索引对齐设备 idx = torch.arange(num_m_bits, device=f.device, dtype=dtype).unsqueeze(0) + (8. - e_scientific).unsqueeze(-1) idx = idx.long() m = torch.gather(m, dim=-1, index=idx) return torch.cat([s, e, m], dim=-1).to(dtype)
内容的提问来源于stack exchange,提问作者zzll
相关产品推荐
相关产品推荐

