Pytorch training: after each layer, how can I make updates to the output and cast the updated output to next layer? I want to keep different bits

Viewed 21

I am doing node classification using Cora dataset in Pytorch. The model consists 2 GCN layers, I want to keep different precision of the output after each layer. Specificially, after each layer, I convert output (float32 tensor type) into binary representations (32 bits). Then I keep only a few bits of 32. Then I convert binary to the float32 and input to the next layer. I encountered an inplace operation, I wonder how to solve it ? Error: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)
0 Answers
Related