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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.22 06:39:14