如何在PyTorch中创建低精度整数(int8)线性层?
实现仅推理的int8单层网络
问题根源在于PyTorch的nn.Parameter默认开启梯度追踪,但int8类型张量不支持梯度运算——既然你只做推理不需要训练,绕开Parameter的梯度机制就能解决,给你两种实用方案:
方案一:手动实现线性层逻辑
直接创建普通int8张量作为权重/偏置,自己写线性运算逻辑,完全避开Parameter的限制:
import torch from torch import nn class Int8InferenceModel(nn.Module): def __init__(self, in_feat, out_feat): super().__init__() # 生成随机int8权重(范围可自定义,这里取int8全区间-127到127) self.weight = torch.randint(low=-127, high=128, size=(out_feat, in_feat), dtype=torch.int8) # 生成int8偏置(不需要的话可以删掉这行) self.bias = torch.randint(low=-127, high=128, size=(out_feat,), dtype=torch.int8) def forward(self, x): # 强制输入转为int8 x = x.to(torch.int8) # 手动实现线性计算:y = x @ weight.T + bias # 注:int8矩阵乘法会自动转int32计算避免溢出,要输出int8的话可以加截断 output = torch.matmul(x, self.weight.t()) + self.bias # 可选:截断并转回int8 # output = torch.clamp(output, -128, 127).to(torch.int8) return output # 测试用例 in_feat = 10 out_feat = 5 int8_model = Int8InferenceModel(in_feat, out_feat) input_int8 = torch.randint(low=-127, high=128, size=(3, in_feat), dtype=torch.int8) output = int8_model(input_int8) print("输入类型:", input_int8.dtype) print("权重类型:", int8_model.weight.dtype) print("输出类型:", output.dtype)
方案二:改造nn.Linear适配int8
如果习惯用nn.Linear的结构,可以先创建浮点型Linear层,再替换为int8权重并关闭梯度:
import torch from torch import nn class Int8LinearModel(nn.Module): def __init__(self, in_feat, out_feat): super().__init__() # 先建浮点型Linear,后续替换权重 self.layer1 = nn.Linear(in_feat, out_feat, bias=True) # 关闭梯度追踪(因为不需要训练) self.layer1.weight.requires_grad = False self.layer1.bias.requires_grad = False # 替换为随机int8权重和偏置 self.layer1.weight.data = torch.randint(low=-127, high=128, size=self.layer1.weight.shape, dtype=torch.int8) self.layer1.bias.data = torch.randint(low=-127, high=128, size=self.layer1.bias.shape, dtype=torch.int8) def forward(self, x): x = x.to(torch.int8) # nn.Linear会自动处理类型转换,运算时自动提升精度避免溢出 output = self.layer1(x) return output # 测试用例 in_feat = 10 out_feat = 5 int8_model = Int8LinearModel(in_feat, out_feat) input_int8 = torch.randint(low=-127, high=128, size=(3, in_feat), dtype=torch.int8) output = int8_model(input_int8) print("权重类型:", int8_model.layer1.weight.dtype) print("输入类型:", input_int8.dtype)
额外说明
- 两种方案都彻底避开了int8张量作为Parameter的问题,完全适配推理场景
- PyTorch中int8运算默认会提升到int32计算,防止溢出;如果必须输出int8,可通过
torch.clamp截断到[-128,127]后转换类型
内容的提问来源于stack exchange,提问作者roadrev
相关产品推荐
相关产品推荐

