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

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)

错误根因

  1. 核心触发点是bit_x[super_node,:,2:] = 0这行代码,属于原位操作(inplace operation):直接修改了存在于PyTorch自动求导计算图中的张量值,反向传播时需要用到该张量前向传播时的原始值做梯度计算,但此时张量已经被修改、版本号不匹配,因此抛出错误。
  2. 辅助函数存在隐藏隐患:bit2float、float2bit中用torch.arange、torch.Tensor([2.])生成的常量张量默认创建在CPU上,若模型迁移到GPU训练会触发设备不匹配报错;且生成常量时未对齐输入张量的数据类型,可能引发精度问题。
  3. 冗余性能损耗: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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.29 03:09:38