跳转至

8.模型量化基本操作

实现对称线性量化的量化和反量化函数,输入 FP32 张量 x,比特数 bits=8,输出量化后的整数张量、scale,并验证反量化后的误差

对称线性量化 将浮点张量映射到有符号整数范围(比如 int8),其核心假设是张量的数值分布关于零点对称。公式如下:

image.png

- 其中 b 是比特数,max(|x|) 是张量绝对值的最大值。对于 b=8,2^{b-1} - 1 = 127

  • 量化:xq=round(x/s),结果应限制在 [-128, 127] 范围内(若使用 int8,但通常对称量化使用有符号整数,范围为 [-(2^{b-1}), 2^{b-1}-1],对于 b=8 是 [-128, 127];也有实现限制在 [-127, 127] 以对称,这里采用完整范围 [-128, 127],但 scale 用 127 是为了保证最大值映射到 127,而最小值映射到 -128 会略有不对称,但常见实现在对称量化时用 127 作为分母,并限制范围到 [-128,127][-127,127],这里采用 [-128,127],scale = max(|x|)/127)。

image.png

实现细节:

  • 需要处理 max(|x|)==0 的边界情况,此时 scale 设为 1.0 或 0 以避免除零。

  • 量化时,将浮点值除以 scale 后四舍五入,再 clamp 到整数范围。

  • 反量化时直接乘 scale 得到近似浮点值。

  • 误差可用均方误差 (MSE) 或平均绝对误差 (MAE) 衡量。

PyTorch 代码示例:

import torch

def symmetric_quantize(x: torch.Tensor, bits=8):
    """
    对称线性量化
    Args:
        x: FP32 张量
        bits: 量化位宽 (默认8)
    Returns:
        x_q: 量化后的整数张量 (torch.int8)
        scale: 浮点 scale
    """
    qmax = 2 ** (bits - 1) - 1   # 127 for 8-bit
    qmin = -qmax - 1             # -128 for 8-bit
    # 计算最大绝对值
    max_val = torch.max(torch.abs(x))
    if max_val == 0:
        scale = torch.tensor(1.0)
    else:
        scale = max_val / qmax
    # 量化
    x_q = torch.round(x / scale)
    x_q = torch.clamp(x_q, qmin, qmax)
    x_q = x_q.to(torch.int8)
    return x_q, scale

def symmetric_dequantize(x_q: torch.Tensor, scale: torch.Tensor):
    """
    对称反量化
    Args:
        x_q: int8 量化张量
        scale: 量化时使用的 scale
    Returns:
        x_fp: 反量化后的 FP32 张量
    """
    return x_q.float() * scale

# 验证
x = torch.randn(1000) * 0.5  # FP32 张量
x_q, scale = symmetric_quantize(x, bits=8)
x_recon = symmetric_dequantize(x_q, scale)

mse = torch.mean((x - x_recon) ** 2).item()
print(f"Scale: {scale.item():.6f}, MSE: {mse:.8f}")
print(f"Original max: {x.max().item():.4f}, min: {x.min().item():.4f}")
print(f"Quantized max: {x_q.max().item()}, min: {x_q.min().item()}")

误差验证: 由于对称量化将最大值映射到 ±127,量化误差主要由四舍五入产生,误差范围约为 ±0.5 * scale。MSE 大致为 (scale^2)/12(在均匀分布假设下)。实际测试中 MSE 会很小。


实现非对称量化的 scale 和 zero_point 计算,编写从 FP32→uint8 的量化及反量化函数

非对称量化 不假设零点对齐,而是将张量的最小值和最大值线性映射到无符号整数范围 [0, 2^b - 1](常用 uint8: 0~255)。需要计算 scale 和 zero_point(零点偏移)。公式:

image.png

实现细节:

  • 如果 min == max,则 scale=1.0,zero_point 取 0 或 128 等,避免除零。

  • zero_point 本身为整数(通常 uint8 或 int32),也需 clamp 到 [0, 2^b-1]

  • 反量化时减去 zero_point 后乘以 scale 恢复浮点值。

代码:

import torch

def asymmetric_quantize(x: torch.Tensor, bits=8):
    """
    非对称量化到 uint8
    Returns:
        x_q: uint8 tensor
        scale: float
        zero_point: int (uint8 范围内的整数)
    """
    qmin = 0
    qmax = 2 ** bits - 1   # 255
    min_val = torch.min(x)
    max_val = torch.max(x)
    if max_val == min_val:
        scale = torch.tensor(1.0)
        zero_point = torch.tensor(0)
    else:
        scale = (max_val - min_val) / (qmax - qmin)
        zero_point = torch.round(-min_val / scale)
        zero_point = torch.clamp(zero_point, qmin, qmax)
    # 量化
    x_q = torch.round(x / scale) + zero_point
    x_q = torch.clamp(x_q, qmin, qmax)
    x_q = x_q.to(torch.uint8)
    return x_q, scale, zero_point.to(torch.int32)

def asymmetric_dequantize(x_q: torch.Tensor, scale: torch.Tensor, zero_point: torch.Tensor):
    """
    非对称反量化
    """
    return (x_q.float() - zero_point.float()) * scale

# 验证
x = torch.randn(1000) * 0.5 + 0.2  # 非对称分布
x_q, scale, zp = asymmetric_quantize(x)
x_recon = asymmetric_dequantize(x_q, scale, zp)
mse = torch.mean((x - x_recon) ** 2).item()
print(f"Scale: {scale:.6f}, Zero Point: {zp.item()}, MSE: {mse:.8f}")

实现 per-channel 对称量化:对于权重 W[out_features, in_features],每个输出通道独立计算 scale,量化为 int8 并反量化

Per-channel 量化 常用于模型权重,因为不同输出通道的权重分布可能差异较大。对于形状为 [out_features, in_features] 的二维权重张量,我们对每个 out_features 行(即每个输出通道)独立计算 scale,得到 shape 为 [out_features] 的 scale 向量。量化时,每行使用对应的 scale,反量化同理。

实现:

  • 输入:W (FP32),shape [OC, IC]

  • 对每一行 i,计算 max_abs = max(|W[i,:]|)scale[i] = max_abs / 127

  • 量化:W_q[i, :] = round(W[i,:] / scale[i]) clamp 到 [-128,127]

  • 反量化:W_fp[i,:] = W_q[i,:] * scale[i]

代码:

def per_channel_symmetric_quantize(w: torch.Tensor, bits=8):
    """
    w: [OC, IC]
    Returns:
        w_q: int8 tensor [OC, IC]
        scale: FP32 tensor [OC]
    """
    qmax = 2 ** (bits - 1) - 1   # 127
    qmin = -qmax - 1             # -128
    # 每行的最大绝对值
    max_abs_per_ch = torch.max(torch.abs(w), dim=1).values  # [OC]
    # 避免除零
    max_abs_per_ch = torch.clamp(max_abs_per_ch, min=1e-8)
    scale = max_abs_per_ch / qmax   # [OC]
    # 量化:w / scale[:, None] 广播
    w_q = torch.round(w / scale.unsqueeze(1))
    w_q = torch.clamp(w_q, qmin, qmax)
    w_q = w_q.to(torch.int8)
    return w_q, scale

def per_channel_symmetric_dequantize(w_q: torch.Tensor, scale: torch.Tensor):
    return w_q.float() * scale.unsqueeze(1)

# 示例
w = torch.randn(64, 128) * 0.5
w_q, scale = per_channel_symmetric_quantize(w)
w_recon = per_channel_symmetric_dequantize(w_q, scale)
mse = torch.mean((w - w_recon) ** 2).item()
print(f"Per-channel MSE: {mse:.8f}")

实现 per-token 激活量化:给定激活张量 [B, L, D],对每个 token 的位置独立计算 scale,对称量化为 int8

Per-token 量化 常用于 Transformer 激活,因为不同 token 的数值范围差异可能很大(如第一个 token 是 [CLS])。对形状 [B, L, D] 的激活,我们在 token 维度(即每个 (b, l) 位置)独立计算 scale。因此 scale 形状为 [B, L]。量化:对于每个 (b,l),将该 token 的 D 维向量使用对应 scale 量化为 int8。反量化同理。

实现:

  • 计算每个 token 的 max_abs:沿最后一维 Dmax(|x|) -> shape [B, L]

  • scale = max_abs / 127。

  • 量化:x_q = round(x / scale.unsqueeze(-1)),clamp。

  • 反量化:x_f = x_q * scale.unsqueeze(-1)

代码:

def per_token_symmetric_quantize(x: torch.Tensor, bits=8):
    """
    x: [B, L, D]
    Returns:
        x_q: int8 [B, L, D]
        scale: FP32 [B, L]
    """
    qmax = 127
    qmin = -128
    # 每个 token 的最大绝对值
    max_abs = torch.max(torch.abs(x), dim=-1).values  # [B, L]
    max_abs = torch.clamp(max_abs, min=1e-8)
    scale = max_abs / qmax  # [B, L]
    # 量化
    x_q = torch.round(x / scale.unsqueeze(-1))
    x_q = torch.clamp(x_q, qmin, qmax).to(torch.int8)
    return x_q, scale

def per_token_symmetric_dequantize(x_q: torch.Tensor, scale: torch.Tensor):
    return x_q.float() * scale.unsqueeze(-1)

实现分组量化(group-wise)的解包与反量化:权重以 int8 存储,每 group_size 个元素共享一个 scale,写出反量化到 FP32 的核函数

分组量化 介于 per-tensor 和 per-channel 之间,将权重矩阵按一定分组大小(如 128)分组,每组独立计算 scale。对于形状为 [K, N] 的权重,通常沿输入维度(列)或沿行分组。这里假设权重为二维 [OC, IC],分组大小 group_size,在 IC 维度上进行分组。存储格式:int8 的权重紧凑排列,另外存储 scale 数组,长度为 OC * ceil(IC/group_size)

反量化核函数需要根据分组索引取出对应的 scale,将 int8 转换为 FP32。实现时可以用循环或 PyTorch 的索引操作。

代码(PyTorch 实现,模拟反量化核函数):

def groupwise_dequantize(w_q: torch.Tensor, scale: torch.Tensor, group_size: int):
    """
    w_q: int8 tensor [OC, IC]  权重(以 int8 紧凑存储)
    scale: FP32 tensor [OC, num_groups]  其中 num_groups = ceil(IC/group_size)
    group_size: 每组的元素数
    Returns:
        w_fp: FP32 [OC, IC]
    """
    OC, IC = w_q.shape
    num_groups = (IC + group_size - 1) // group_size
    # 将 w_q reshape 为 [OC, num_groups, group_size] 但最后一组可能不足
    # 简单实现:循环
    w_fp = torch.empty_like(w_q, dtype=torch.float32)
    for g in range(num_groups):
        start = g * group_size
        end = min(start + group_size, IC)
        # 取出该组 scale,形状 [OC]
        s = scale[:, g].unsqueeze(1)  # [OC, 1]
        w_fp[:, start:end] = w_q[:, start:end].float() * s
    return w_fp

若权重的 int8 存储是按组交织的(如 [OC * num_groups, group_size]),则需要相应解包。这里演示简单形式。更高效的核函数可用 CUDA,逻辑同上。


编写 FakeQuantize 模块(PyTorch),前向时对输入进行对称量化再反量化,反向时使用直通估计器(STE),确保梯度正确传递

FakeQuantize 在训练时模拟量化误差,以帮助模型适应量化后的精度损失。前向过程:对输入进行量化(round)再反量化,相当于添加了近似量化噪声。反向传播时,由于 round 函数梯度为零(几乎处处为0),需要使用直通估计器(STE):直接将梯度原样通过量化模块,即忽略 round 的影响。

PyTorch 实现:可以继承 torch.autograd.Function 自定义正向和反向,或使用 torch.nn.Module 结合 detach 技巧。标准做法是实现一个 FakeQuantize 函数,前向做量化-反量化,反向直接返回上游梯度(对输入求导为 1)。

示例代码:

import torch
import torch.nn as nn

class FakeQuantizeSTE(torch.autograd.Function):
    @staticmethod
    def forward(ctx, x, scale, qmin, qmax):
        # 对称量化:round(x / scale) * scale
        x_q = torch.round(x / scale)
        x_q = torch.clamp(x_q, qmin, qmax)
        x_deq = x_q * scale
        # 保存用于反向(如果需要可以根据 x_deq 与 x 的关系传递梯度,但 STE 直接传)
        ctx.save_for_backward(scale)
        ctx.qmin = qmin
        ctx.qmax = qmax
        return x_deq

    @staticmethod
    def backward(ctx, grad_output):
        # STE: 直接传递梯度,忽略量化的影响
        return grad_output, None, None, None

# 封装为模块
class FakeQuantize(nn.Module):
    def __init__(self, bits=8, symmetric=True):
        super().__init__()
        self.bits = bits
        self.symmetric = symmetric
        if symmetric:
            self.qmax = 2 ** (bits - 1) - 1
            self.qmin = -self.qmax - 1
        else:
            self.qmax = 2 ** bits - 1
            self.qmin = 0
        # scale 可以作为可学习参数或由外部统计给出,这里假设在 forward 时传入 scale
        self.register_buffer('scale', torch.tensor(1.0))

    def forward(self, x, scale=None):
        if scale is None:
            scale = self.scale
        return FakeQuantizeSTE.apply(x, scale, self.qmin, self.qmax)

# 使用示例
x = torch.randn(5, requires_grad=True)
scale = torch.tensor(0.1)
fq = FakeQuantize(bits=8)
y = fq(x, scale)
loss = y.sum()
loss.backward()
print(x.grad)  # 全部为 1,因为 STE 梯度直通

说明:FakeQuantize 在 QAT(量化感知训练)中常插入到权重和激活之前,scale 可由校准得到或在线更新。


实现 MinMax 校准方法:给定一组校准数据,统计 activation 的全局最小值和最大值,计算对称量化的 scale(支持 per-tensor 和 per-token 选项)

MinMax 校准 是最简单的校准方式:在校准数据集上运行模型,收集每层激活(或权重)的 min 和 max,然后基于全局或 per-token 的极值计算 scale。

  • per-tensor:所有校准批次中激活张量的全局 min 和 max。

  • per-token:对每个 token 位置单独统计 min/max(常用于激活量化)。注意“per-token”可能指每个 token 位置(如序列中每个 token)或者每个 token 维度独立?通常指对形状 [B, L, D] 的激活,我们希望计算每个 token(即每个 (b,l))的 min/max,然后得出每个 token 的 scale,但校准期间需要将这些统计聚合为某种代表性尺度。常见实践是:对每个 token 位置,我们记录 running min/max,最终得到 shape [1, 1, D]? 或者 per-tensor 就是标量。本题要求支持 per-tensor 和 per-token 选项。per-tensor 使用单个全局 scale;per-token 为每个 token 位置独立计算 scale(如每个 token 的 D 维向量有自己的 scale),统计时需记录每个 token 的 max_abs 的滑动平均或全局最大值。

我们实现一个校准器类,输入校准数据(激活张量),逐步更新 min/max,最后计算 scale。

代码示例:

class MinMaxCalibrator:
    def __init__(self, mode='per_tensor', bits=8):
        """
        mode: 'per_tensor' 或 'per_token'
        """
        self.mode = mode
        self.bits = bits
        self.qmax = 2 ** (bits - 1) - 1
        self.initialized = False
        # 统计值
        self.min_val = None
        self.max_val = None

    def update(self, x: torch.Tensor):
        """
        x: 激活张量,可以是任意维度,假设最后维度是特征维度
        """
        if self.mode == 'per_tensor':
            batch_min = torch.min(x)
            batch_max = torch.max(x)
            if not self.initialized:
                self.min_val = batch_min
                self.max_val = batch_max
                self.initialized = True
            else:
                self.min_val = torch.min(self.min_val, batch_min)
                self.max_val = torch.max(self.max_val, batch_max)
        elif self.mode == 'per_token':
            # 假设 token 维度为最后两维之前的维度,如 [B, L, D],per-token 意味着每个 token (B,L) 独立
            # 简化:对整个张量,保留每个 token 的 min/max 形状和输入匹配
            # 此处实现为保留每个 token 的 min/max(按元素)
            # 计算沿最后一维的 min/max
            token_min = torch.min(x, dim=-1, keepdim=True)[0]  # [..., 1]
            token_max = torch.max(x, dim=-1, keepdim=True)[0]
            if not self.initialized:
                self.min_val = token_min
                self.max_val = token_max
                self.initialized = True
            else:
                # 假设输入形状在 token 维度一致,逐元素取 min/max
                self.min_val = torch.min(self.min_val, token_min)
                self.max_val = torch.max(self.max_val, token_max)
        else:
            raise ValueError("mode 必须是 'per_tensor' 或 'per_token'")

    def compute_scale(self, symmetric=True):
        if not self.initialized:
            raise RuntimeError("未提供校准数据")
        if symmetric:
            max_abs = torch.max(torch.abs(self.min_val), torch.abs(self.max_val))
            scale = max_abs / self.qmax
        else:
            scale = (self.max_val - self.min_val) / (2 ** self.bits - 1)
        return scale

使用:

calib = MinMaxCalibrator(mode='per_tensor', bits=8)
for data in calib_loader:
    act = model.get_activation(data)  # 形状任意
    calib.update(act)
scale = calib.compute_scale()

使用 MSE 校准计算最优 scale:对每个通道的权重,在 MinMax 范围内搜索 scale,使得量化-反量化后的 MSE 最小,实现一维网格搜索

MSE 校准 通过最小化量化误差来寻找最佳 scale。对于权重张量,我们希望找到一个 scale,使得 Wquantize_dequantize(W, scale) 之间的 MSE 最小。搜索范围可以从 MinMax scale 出发,在一定区间内线性搜索。

对于对称量化,量化-反量化过程为:clamp(round(W/s), -128,127) * s。MSE = mean((W - W_hat)^2)。可以通过在 [alpha * s_minmax, beta * s_minmax] 范围内采样 s,计算 MSE,选择 MSE 最小的 s。常见做法:设置搜索步数(如 100),在 [0.5*s_minmax, 1.5*s_minmax][s_minmax/2, s_minmax*2] 之间线性搜索。

代码实现(针对 per-channel 权重):

def mse_calibration_per_channel(w: torch.Tensor, bits=8, search_steps=100, alpha=0.5, beta=1.5):
    """
    w: [OC, IC] 权重
    Returns:
        best_scale: [OC] 最优 scale
    """
    OC, IC = w.shape
    qmax = 2 ** (bits - 1) - 1
    qmin = -qmax - 1

    # 先计算 per-channel 的 minmax scale
    max_abs_per_ch = torch.max(torch.abs(w), dim=1).values  # [OC]
    base_scale = max_abs_per_ch / qmax  # [OC]

    best_scales = torch.empty_like(base_scale)
    # 为每个通道搜索
    for c in range(OC):
        w_c = w[c]  # [IC]
        base_s = base_scale[c].item()
        best_s = base_s
        best_mse = float('inf')
        # 在 [alpha*base_s, beta*base_s] 范围线性搜索
        for step in range(search_steps):
            s = base_s * (alpha + (beta - alpha) * step / (search_steps - 1))
            # 量化反量化
            w_q = torch.clamp(torch.round(w_c / s), qmin, qmax)
            w_hat = w_q * s
            mse = torch.mean((w_c - w_hat) ** 2).item()
            if mse < best_mse:
                best_mse = mse
                best_s = s
        best_scales[c] = best_s
    return best_scales

对于大型权重矩阵,逐个通道搜索较慢,可向量化部分操作或减小搜索范围。这种方法能找到比 MinMax 更好的 scale,提升量化精度。


实现百分位校准:对激活张量统计其 99.99% 分位数的绝对值,以此作为 max_val 计算对称量化的 scale

百分位校准 是为了减少离群值对量化范围的影响。直接使用最大值可能导致 scale 过大,大部分数值量化后分辨率不足。改用高分位数(如 99.99%)的绝对值作为有效最大值,能更好地平衡覆盖范围与精度。对称量化的 scale 公式为:

image.png

实现细节:

  • 使用 PyTorch 的 torch.quantile 计算分位数,但它只支持一维张量。所以先将激活张量展平为一维,再求分位数。

  • 如果分位数为 0,则回退到一个极小值(如 1e-8)以避免除零。

  • 计算得到的 scale 作为该激活的量化尺度,后续采用对称量化。

代码示例:

import torch

def percentile_calibrate_scale(x: torch.Tensor, percentile=0.9999, bits=8):
    """
    使用百分位校准计算对称量化的 scale
    Args:
        x: FP32 激活张量
        percentile: 分位数 (如 0.9999)
        bits: 量化位宽
    Returns:
        scale: 浮点 scale
    """
    qmax = 2 ** (bits - 1) - 1   # 127 for 8-bit
    # 展平并计算绝对值分位数
    x_abs_flat = torch.abs(x).flatten()
    # 若张量非空
    if x_abs_flat.numel() == 0:
        return torch.tensor(1.0)
    max_val = torch.quantile(x_abs_flat, percentile)
    if max_val == 0:
        max_val = torch.tensor(1e-8, device=x.device)
    scale = max_val / qmax
    return scale

# 使用示例
x = torch.randn(10000) * 0.5
# 加入一个离群值
x[0] = 100.0
scale = percentile_calibrate_scale(x, 0.9999)
print(f"Percentile scale: {scale:.6f}")
# 对比 minmax scale
max_abs = torch.max(torch.abs(x))
scale_minmax = max_abs / 127
print(f"MinMax scale: {scale_minmax:.6f}")

说明:百分位校准能显著减小离群值对量化精度的影响,尤其在激活分布存在长尾时效果很好,是量化部署中的常用策略。


编写将 FP16 权重打包为 INT4 格式的函数:每两个 4-bit 值存入一个 uint8,并实现对应的解包函数

INT4 打包 即将有符号 4-bit 整数(通常范围 [-8,7])两个一组拼成一个字节(uint8)。通常低 4 位存放第一个值,高 4 位存放第二个值。我们假设量化后的权重值在 [-8,7] 范围内。

打包函数:

  • 输入:w_q 为 PyTorch int8 张量,但值只使用低 4 位(范围 [-8,7])。

  • 将连续两个元素合并:packed = (w_q[2*i] & 0xF) | ((w_q[2*i+1] & 0xF) << 4),结果存储为 uint8。

  • 如果元素个数为奇数,最后一个字节的高 4 位补 0。

解包函数:

  • 将 uint8 拆分为两个 int8,恢复为 [N] 形状。需要将高 4 位右移 4 位并符号扩展:对于 4 位有符号数,若高 4 位最高位(bit3)为 1,表示负数,需进行符号扩展。解包逻辑:low = packed & 0x0F 然后 low = low - 16 if low > 7 else low,同理高 4 位:high = (packed >> 4) & 0x0F; high = high - 16 if high > 7 else high

代码实现:

def pack_int4(w_q: torch.Tensor) -> torch.Tensor:
    """
    将 int8(值在[-8,7])打包为 uint8,每两个元素一字节
    Args:
        w_q: int8 tensor,形状 [N],其中数值范围 [-8, 7]
    Returns:
        packed: uint8 tensor,形状 [ceil(N/2)]
    """
    # 确保值在有效范围内并转换为无符号整数的低4位表示
    w_q = w_q.to(torch.int8)
    # 转为无符号4位(0~15),负数加16
    w_u4 = torch.where(w_q < 0, w_q + 16, w_q).to(torch.uint8)  # 0~15
    N = w_q.numel()
    if N % 2 == 1:
        # 补齐到偶数长度,多出的一个高4位置0
        w_u4 = torch.cat([w_u4, torch.zeros(1, dtype=torch.uint8, device=w_q.device)])
    # 每两个元素:低4位来自偶数索引,高4位来自奇数索引
    low = w_u4[0::2]            # 低4位
    high = w_u4[1::2] << 4      # 高4位
    packed = low | high
    return packed.to(torch.uint8)

def unpack_int4(packed: torch.Tensor, original_length: int) -> torch.Tensor:
    """
    将打包的 uint8 解包为 int8,恢复形状 [original_length]
    """
    low = (packed & 0x0F).to(torch.int8)
    high = ((packed >> 4) & 0x0F).to(torch.int8)
    # 符号扩展:如果值 > 7 则是负数,需转为 [-8, -1]
    low = torch.where(low > 7, low - 16, low)
    high = torch.where(high > 7, high - 16, high)
    # 交替交织
    unpacked = torch.stack([low, high], dim=1).flatten()  # [2*len(packed)]
    # 截取到原始长度
    return unpacked[:original_length]

使用示例:

w_q = torch.tensor([1, -2, 7, -8, 3], dtype=torch.int8)
packed = pack_int4(w_q)
print(packed)  # 例如: tensor([0xEE?, ...]) 取决于顺序
unpacked = unpack_int4(packed, w_q.numel())
print(unpacked)  # 应该与 w_q 一致

实现 INT8 矩阵乘法的模拟:输入两个 FP32 矩阵,分别量化为 int8(带 scale),在 int32 下完成矩阵乘法,再反量化回 FP32,输出结果

INT8 矩阵乘法 是量化推理的核心。公式:假设有矩阵 A (M×K) 和 B (K×N),分别量化为 int8,scale 为 s_as_b(可以是标量或向量)。则:

image.png

其中 A_q, B_q 是 int8 矩阵,矩阵乘积在 int32 累积(避免溢出),然后乘以 scale 的乘积反量化。如果 scale 是 per-tensor 标量,直接相乘;如果 scale 是 per-channel 向量,则需要广播处理。

实现要点:

  • 使用对称 per-tensor 量化以简化,scale 为标量。

  • 矩阵乘法可以使用 PyTorch 的 torch.matmultorch.mm,输入需转换为 int8? PyTorch 对 int8 的 matmul 支持有限(通常需要通过 torch.int8 乘法后自动提升为 int32?)。可以直接使用 torch.matmul(A_q.float(), B_q.float()) 但这就不是整数计算了。为了模拟整数计算,我们依然用浮点计算乘积,只是用 int8 表示值,最终结果乘以 scale 乘积。这样在数学上等价。如果要真正使用整数运算,可利用 torch.int32 的矩阵乘法,但 PyTorch 的 torch.mm 不支持 int8。因此模拟时可以用 torch.matmul(A_q.to(torch.int32), B_q.to(torch.int32)) 在整数域得到 int32,然后转为 float 乘以 scale。这样可以验证整数计算的量化效果(精度损失)。

代码:

def int8_matmul_sim(A_fp: torch.Tensor, B_fp: torch.Tensor, bits=8):
    """
    模拟 INT8 矩阵乘法:A(M,K) @ B(K,N)
    采用 per-tensor 对称量化
    """
    # 量化 A, B
    A_q, scale_a = symmetric_quantize(A_fp, bits)  # 返回 int8, scale
    B_q, scale_b = symmetric_quantize(B_fp, bits)

    # 整数矩阵乘法(int32 避免溢出)
    C_int32 = torch.matmul(A_q.to(torch.int32), B_q.to(torch.int32))

    # 反量化
    scale_c = scale_a * scale_b
    C_fp = C_int32.float() * scale_c
    return C_fp, C_int32, scale_c

# 测试
A = torch.randn(4, 8) * 0.5
B = torch.randn(8, 4) * 0.5
C_orig = A @ B
C_sim, _, _ = int8_matmul_sim(A, B)
print("Max error:", (C_orig - C_sim).abs().max().item())

如需 per-channel,scale_a 为 [M] 或 [K],需要广播乘法。一般量化权重 B 使用 per-channel(输出通道),激活 A 用 per-tensor。代码可相应调整。


实现 Weight-Only 量化推理:权重存储为 int8 和 scale,激活为 FP32,在不事先量化激活的情况下模拟矩阵乘法(即反量化权重后与 FP32 激活相乘)

Weight-Only 量化 是一种在推理时只量化权重,而激活保持浮点的方案,减少内存占用和加载带宽,但计算仍在浮点域进行。这种方式避免了激活量化的复杂校准和精度损失,适用于大模型推理(如 llama.cpp 的某些模式)。

实现:

  • 权重 W 在推理前被量化为 int8,并保存 scale(可以是 per-channel)。

  • 推理时,对于输入 x (FP32),先将 W 反量化为 FP32,然后执行 FP32 矩阵乘法:y = x @ W_deq.Ty = F.linear(x, W_deq)

  • 为了模拟真实推理,也可以不在内存中完整恢复 W 的 FP32 矩阵,而是动态解量化和计算(如分块)。这里简单实现直接反量化整个权重矩阵,然后矩阵乘。

代码:

def weight_only_linear(x_fp: torch.Tensor, w_q: torch.Tensor, scale: torch.Tensor, bias=None):
    """
    权重仅量化推理
    Args:
        x_fp: 激活 FP32, shape [..., in_features]
        w_q: int8 权重 [out_features, in_features]
        scale: scale, 如果 per-channel 则为 [out_features]
    Returns:
        y_fp: FP32 输出
    """
    # 反量化权重到 FP32
    if scale.dim() == 1:  # per-channel
        w_deq = w_q.float() * scale.unsqueeze(1)
    else:  # per-tensor
        w_deq = w_q.float() * scale
    # 线性计算
    return torch.nn.functional.linear(x_fp, w_deq, bias)

注意:实际部署时,为了节省内存,不会显式重建整个 FP32 权重,而是分块解量化并直接进行乘加运算,但逻辑等价。


实现动态激活量化:对一个 Linear 层,激活输入为 FP32,在线计算 scale 并量化为 int8,然后与已量化的 int8 权重进行整数矩阵乘,输出 FP32

动态激活量化 在每次前向传播时,对激活进行量化(根据实际数值计算 scale),然后与 int8 权重进行整数矩阵乘法,最后反量化得到 FP32 输出。这种方式无需事先校准激活,但增加了在线量化的开销。

步骤:

  1. 输入 x (FP32),shape [..., in_features]

  2. 计算 x 的 scale(对称,per-tensor 或 per-token)。

  3. 量化激活:x_q = round(x / scale_act), clamp 到 int8 范围。

  4. 权重已经量化好为 int8,并有其 scale (per-channel 或 per-tensor)。

  5. 整数矩阵乘:y_int32 = x_q @ W_q^T(注意形状适配)。如果激活为 [M, K],权重为 [N, K](已转置),则乘法为 [M, N]。

  6. 反量化:y_fp = y_int32 * (scale_act * scale_w),其中 scale_w 需广播匹配输出维度。

代码实现:

def dynamic_quantize_linear(x_fp: torch.Tensor, w_q: torch.Tensor, w_scale: torch.Tensor,
                            bias=None, bits=8):
    """
    动态激活量化线性层
    x_fp: [..., in_features]  FP32 激活
    w_q: [out_features, in_features] int8 权重
    w_scale: [out_features] 或标量,权重量化 scale
    """
    # 动态量化激活(per-tensor 对称量化)
    x_q, act_scale = symmetric_quantize(x_fp, bits)  # x_q: int8, act_scale: 标量

    # 整数矩阵乘法: [M, K] x [N, K]^T = [M, N]
    # 将 x_q 转换为 int32 防止溢出,与 w_q 相乘
    y_int32 = torch.matmul(x_q.to(torch.int32), w_q.to(torch.int32).t())

    # 反量化
    if w_scale.dim() == 1:
        # per-channel: 形状 [out_features]
        combined_scale = act_scale * w_scale  # [out_features]
    else:
        combined_scale = act_scale * w_scale  # 标量
    y_fp = y_int32.float() * combined_scale
    if bias is not None:
        y_fp += bias
    return y_fp

要点:symmetric_quantize 前文已实现,注意激活量化可能使用 per-token 以获得更高精度,这里为简单用 per-tensor。


实现 SmoothQuant 的平滑因子计算:给定激活 X 和权重 W,分别求 X 列的最大绝对值、W 行最大绝对值,按公式 s = (max(|X_j|)^α / max(|W_j|)^(1-α)) 计算平滑因子,并应用到 X 和 W

SmoothQuant 是一种缓解激活异常值导致量化困难的技巧。它通过将量化难度从激活转移到权重来实现。具体来说,对于激活 X 和权重 W(形状分别为 [M, K][N, K],注意 W 的输入维度与 X 的列维度相同,均为 K),我们计算每个输入通道 j (0 ≤ j < K) 的平滑因子 s_j:

image.png

其中 α 通常为 0.5。然后将激活除以 s_j,权重乘以 s_j(在对应的输入通道维度上),从而保持乘积结果不变。这样使激活变得更平滑,易于量化;权重的异常值相应变化。

计算步骤:

  1. 对于激活 X: shape [M, K],计算每个通道的最大绝对值 max_x = max(|X|, dim=0).values -> [K]

  2. 对于权重 W: shape [N, K],沿 N 维度(输出通道)求每个输入通道的最大绝对值?注意公式中 max(|W_j|) 应为权重在输入通道 j 上的最大绝对值,即 max(|W[:, j]|) -> [K]

  3. 计算 s_j = (max_x[j] ** alpha) / (max_w[j] ** (1 - alpha)),避免除零加 epsilon。

  4. 应用平滑:X_new = X / s (广播);W_new = W * s (广播,在输入通道维度)。

代码:

def smoothquant_scale(X: torch.Tensor, W: torch.Tensor, alpha=0.5, eps=1e-8):
    """
    X: 激活 [M, K]  FP32
    W: 权重 [N, K]  FP32  (注意第二维是输入通道)
    alpha: 平滑系数
    Returns:
        X_smooth, W_smooth, s
    """
    # 每个输入通道的最大绝对值
    max_x, _ = torch.max(torch.abs(X), dim=0)  # [K]
    max_w, _ = torch.max(torch.abs(W), dim=0)  # [K]

    # 平滑因子
    s = (max_x.pow(alpha) + eps) / (max_w.pow(1 - alpha) + eps)  # [K]

    # 应用
    X_smooth = X / s  # 广播除
    W_smooth = W * s  # 广播乘

    return X_smooth, W_smooth, s

注意:公式中的 α 调整激活与权重之间的量化难度分配。α=1 时等价于仅平滑激活(相当于激活量化友好),α=0 时仅平滑权重。经验值 0.5 适用于多数场景。


实现 LLM.int8() 混合精度矩阵乘法的简化版:识别输入 X 中绝对值大于阈值(如 6.0)的特征列,这些列用 FP16 计算,其余列量化为 INT8 计算,最后合并结果

LLM.int8() 观察到大语言模型激活中存在少量离群特征维度(异常值),它们集中出现在几个特定列上。将这些离群列用 FP16 计算,其余大部分列用 INT8 量化计算,可以同时保持速度与精度。

实现简化版:

  • 输入 X: [M, K],权重 W: [N, K](注意形状,LLM 通常是 X @ W^T)。

  • 统计 X 每一列的最大绝对值,若某列的最大绝对值 > threshold(如 6.0),则标记该列为离群列 outlier_cols

  • 将 X 分为两部分:X_out 包含离群列(FP16),X_in 为其余列。

  • 权重 W 相应分为 W_outW_in

  • X_out @ W_out^T 使用 FP16 矩阵乘法(高精度)。

  • X_inW_in 进行 INT8 量化(动态或静态)和整数矩阵乘,得到 INT32 结果后反量化。

  • 两部分结果相加得到最终输出。

代码:

def llm_int8_matmul_simple(X: torch.Tensor, W: torch.Tensor, threshold=6.0, bits=8):
    """
    X: [M, K] FP32 激活
    W: [N, K] FP32 权重
    threshold: 离群值阈值
    Returns:
        Y: [M, N] FP32
    """
    M, K = X.shape
    N = W.shape[0]

    # 识别离群列:沿 M 维度取每个特征列的绝对最大值
    col_max, _ = torch.max(torch.abs(X), dim=0)  # [K]
    outlier_mask = col_max > threshold            # [K] bool
    inlier_mask = ~outlier_mask

    # 分离离群和正常部分
    X_out = X[:, outlier_mask]  # FP32 -> 后续转 FP16
    W_out = W[:, outlier_mask]  # FP32 -> FP16

    X_in = X[:, inlier_mask]
    W_in = W[:, inlier_mask]

    # FP16 计算离群部分
    Y_out = torch.matmul(X_out.half(), W_out.t().half()).float()

    # INT8 计算正常部分
    if X_in.shape[1] > 0:
        # 使用动态量化方法
        Y_in = dynamic_quantize_linear(X_in, W_in.contiguous(),
                                       # 需要提前量化权重或在线量化。这里简化:在线量化 W_in
                                       # 我们可以直接调用之前实现的 int8_matmul_sim 等
                                       )
        # 为简便,我们调用通用 int8 矩阵乘法
        Y_in_sim, _, _ = int8_matmul_sim(X_in, W_in.t())  # 注意转置
        Y_in = Y_in_sim
    else:
        Y_in = torch.zeros(M, N, device=X.device)

    return Y_out + Y_in

说明:真实的 LLM.int8() 中离群特征的维度是固定的(例如在特定注意力头),且离群检测与混合计算在内核级别做了优化。这里简化为列级别的阈值检测,按列拆分并分别用 FP16 和 INT8 计算,合并结果。可动态量化 W_in 使用 per-channel 量化以获得更好精度。


实现跨层均衡(CLE):给定两个连续的 Linear 层 weight1 和 weight2,通过计算通道缩放因子,均衡两个权重的范围,以便后续量化

跨层均衡(Cross-Layer Equalization, CLE) 用来处理两个线性层之间激活函数的缩放等效性。例如,对于 Linear1 -> ReLU -> Linear2 这样的结构,可以将权重按通道缩放,同时保证输出不变。设 Weight1 形状 [N1, K],Weight2 形状 [N2, K](注意两个权重的公共维度是 K,即中间特征维度)。我们可以引入缩放因子向量 s (形状 [K]),将 Weight1 的第 j 列乘以 s_j,Weight2 的第 j 列除以 s_j(因为第二个权重通常与中间激活相乘,中间激活是 Weight1 的输出,受 s 影响;实际上 CLE 应用于卷积或线性层,需保证等价性。对于 Linear1 -> Linear2,若中间无非线性或仅有 ReLU,满足 ReLU(s*x) = s * ReLU(x) 当 s>0,因此可缩放。公式:W1_new[:, j] = W1[:, j] * s_jW2_new[:, j] = W2[:, j] / s_j。这样输出不变。

计算缩放因子 常通过均衡两个权重的极值范围,使得它们具有相近的量化范围。一种常见方法:对每个通道 j,计算 r1_j = max(|W1[:, j]|)r2_j = max(|W2[:, j]|),然后 s_j = sqrt(r2_j / r1_j) 或更一般的 s_j = r1_j^(-α) * r2_j^(1-α)?经典做法是 s_j = (r1_j / r2_j)^{1/2},即几何平均的倒数。目的是让两个权重的范围接近。

实现:

def cross_layer_equalization(W1: torch.Tensor, W2: torch.Tensor):
    """
    W1: [out1, K] 第一层权重
    W2: [out2, K] 第二层权重 (注意中间维度 K 必须一致)
    Returns:
        W1_eq, W2_eq, s: 均衡后的权重和缩放因子
    """
    # 计算每个通道的最大绝对值
    r1 = torch.max(torch.abs(W1), dim=0).values  # [K]
    r2 = torch.max(torch.abs(W2), dim=0).values  # [K]

    # 防止除零
    eps = 1e-8
    s = torch.sqrt(r2 / (r1 + eps))

    # 应用缩放
    W1_eq = W1 * s  # 广播乘,每列乘对应 s
    W2_eq = W2 / s  # 每列除
    return W1_eq, W2_eq, s

说明:跨层均衡后,两个权重在各通道上的范围更均衡,有利于统一量化。通常作为量化前预处理步骤,与偏置修正等方法结合。在实际部署中,可将缩放因子吸收到权重中,而无需额外操作。


实现量化感知训练(QAT)的一个小训练循环:对一个 SimpleNet 插入 FakeQuantize,使用 SGD 优化器,训练一步并验证权重梯度是否正确。

原理

QAT 在前向传播中模拟量化误差,反向传播使用 STE,使得模型参数能够适应量化噪声。下面构造一个简单的多层感知机,插入 FakeQuantize,训练一步,并对比权重更新与不使用 QAT 时的差异。

import torch.nn as nn
import torch.optim as optim

class SimpleNet(nn.Module):
    def __init__(self):
        super().__init__()
        self.fc1 = nn.Linear(10, 20)
        self.fc2 = nn.Linear(20, 5)
        self.fq1 = FakeQuantize(bits=8, symmetric=True)
        self.fq2 = FakeQuantize(bits=8, symmetric=True)

    def forward(self, x):
        x = self.fc1(x)
        x = self.fq1(x)      # 量化激活
        x = torch.relu(x)
        x = self.fc2(x)
        x = self.fq2(x)      # 量化输出(若需要)
        return x

# 初始化
model_qat = SimpleNet()
optimizer = optim.SGD(model_qat.parameters(), lr=0.01)

# 随机输入和目标
x = torch.randn(4, 10)
target = torch.randn(4, 5)

# 前向
out = model_qat(x)
loss = torch.nn.functional.mse_loss(out, target)
optimizer.zero_grad()
loss.backward()
# 检查 fc1 权重的梯度是否存在且不为零
print("fc1 weight grad norm:", model_qat.fc1.weight.grad.norm().item())
optimizer.step()
print("QAT step completed.")

验证梯度:由于使用了 FakeQuantize 模块,其中的 STE 确保梯度能够回传到前面的层,权重会得到更新。可以通过对比 QAT 和普通训练若干步后的精度来评估效果。


实现将 FP32 模型参数按 per-channel 量化为 int8 并生成量化参数字典的函数,同时输出 scale 和 zero_point。

原理

对模型的每一个可量化权重(如 nn.Linear 的 weight),按输出通道进行对称量化,得到 int8 权重和 per-channel scale。将这些信息打包成字典,便于后续部署或反量化。

def quantize_model_weights(model, bits=8):
    quant_state = {}
    for name, module in model.named_modules():
        if isinstance(module, nn.Linear):
            W = module.weight.data  # [out_features, in_features]
            # per-channel 对称量化
            qmax = 2**(bits-1) - 1
            max_per_ch = torch.max(torch.abs(W), dim=1).values
            scales = max_per_ch / qmax
            scales[scales==0] = 1e-8
            W_int = torch.round(W / scales.unsqueeze(1)).clamp(-qmax, qmax).to(torch.int8)
            # 存储
            quant_state[name + '.weight'] = {'int': W_int, 'scale': scales, 'bits': bits}
            # 对于 bias,通常不量化,保留原样
            if module.bias is not None:
                quant_state[name + '.bias'] = module.bias.data.clone()
    return quant_state

# 使用
model = SimpleNet()
q_state = quantize_model_weights(model)

反量化重建:

def dequantize_model(model, q_state):
    for name, module in model.named_modules():
        if isinstance(module, nn.Linear):
            w_info = q_state[name + '.weight']
            W_fp = w_info['int'].float() * w_info['scale'].unsqueeze(1)
            module.weight.data.copy_(W_fp)

这样即可将模型权重转换为 int8 并存储为字典,用于后续的量化推理或微调。


实现量化的矩阵乘法的 GPU 模拟(使用 Triton 或 CUDA 伪代码):将输入矩阵分块,每块完成加载、反量化、乘加并写回,重点写出反量化与点积的融合逻辑。

原理

在 GPU 上实现量化矩阵乘法时,通常采用 Block-wise 计算,融合反量化。这里给出 Triton 伪代码(简化),展示如何在每个线程块内加载一个 BLOCK_M x BLOCK_K 的激活块和 BLOCK_K x BLOCK_N 的权重块,然后对每个元素反量化并累加到输出。

# Triton 伪代码(不可直接运行,仅示意)
@triton.jit
def int8_matmul_kernel(
    A_ptr, B_ptr, C_ptr,
    scale_A, scale_B,
    M, N, K,
    stride_am, stride_ak, ...
    BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr
):
    pid = tl.program_id(0)
    # 计算当前块在 M/N 维度的起始位置
    # 加载 A 块 [BLOCK_M, BLOCK_K] int8
    # 加载 B 块 [BLOCK_K, BLOCK_N] int8
    # 反量化与乘加:
    acc = tl.zeros([BLOCK_M, BLOCK_N], dtype=tl.float32)
    for k in range(0, K, BLOCK_K):
        a = tl.load(A_ptr + offsets, mask=...)  # int8
        b = tl.load(B_ptr + offsets, mask=...)  # int8
        # 反量化到 FP32(融合)
        a_fp = a.to(tl.float32) * scale_A  # scale_A 可以是一个标量或向量
        b_fp = b.to(tl.float32) * scale_B
        # 点积累加
        acc += tl.dot(a_fp, b_fp)
    # 写回 C

实际实现中,为了高效利用 Tensor Core,权重可能以 INT4/INT8 直接参与矩阵乘指令(如 mma.sync),反量化由硬件或指令自动处理,无需手动转为 FP32。但上述伪代码清晰展示了反量化与点积的融合思想。


实现量化缩放因子的存储压缩:对于 per-channel 的 scale 量化为 FP8 或共享部分指数,写出压缩与解压缩代码。

原理

在 per-channel 量化中,scale 数量等于输出通道数,对于大模型(如 70B),scale 存储也占据一定空间。可以对这些 scale 进一步压缩,例如将 FP32 scale 量化为 FP8(8位浮点),或者利用 scale 之间数值相近的特性,共享部分指数位。这里演示将 FP32 scale 量化为 FP8(E4M3 或 E5M2)的过程,以及简单的共享指数压缩方案。

方案一:FP32 → FP8 量化

我们模拟 FP8 E4M3 格式(4位指数,3位尾数)。PyTorch 目前没有原生 FP8 张量,但可以用整数模拟存储,并实现反量化。

import numpy as np

def fp32_to_fp8(scale_fp32):
    """
    将 FP32 张量转换为 FP8 E4M3 格式的整数表示。
    这里简单使用截断方法,不实现完整舍入逻辑。
    """
    # 提取符号、指数、尾数
    # 为简化,使用 numpy 的 float16 作为中间,再转为自定义 FP8。
    # 实际硬件中,FP8 有专门的转换指令。
    # 我们仅模拟:将 FP32 钳制到 FP8 表示范围,并返回整数编码。
    pass

方案二:共享指数压缩

如果多个 scale 数值量级相近,可以让它们共享同一个指数,仅存储各自的尾数。例如,对于一组 scale,找出一个公共的 exponent_base,然后将每个 scale 表示为 mantissa * 2^exponent_base,存储 int8 的尾数即可。

def compress_scales_shared_exp(scales, base_exp=None):
    """
    scales: [N] FP32
    返回: mantissas (int8), base_exp (int)
    """
    max_val = torch.max(scales).item()
    # 选择一个公共指数,使得缩放后的尾数落在 int8 范围内
    # 例如 base_exp = floor(log2(max_val/127))
    base_exp = int(math.floor(math.log2(max_val / 127.0)))
    # 对每个 scale,计算 mantissa = round(scale / 2^base_exp)
    mantissas = torch.round(scales / (2.0 ** base_exp))
    mantissas = mantissas.clamp(-128, 127).to(torch.int8)
    return mantissas, base_exp

def decompress_scales_shared_exp(mantissas, base_exp):
    return mantissas.float() * (2.0 ** base_exp)

# 测试
scales = torch.rand(128) * 0.01
m, exp = compress_scales_shared_exp(scales)
scales_rec = decompress_scales_shared_exp(m, exp)
print("压缩误差:", (scales - scales_rec).abs().mean().item())

应用:这类压缩常用于进一步减小模型文件尺寸,如 GGUF 格式中的 scale 压缩。


实现简单的 AdaRound 自适应舍入

原理

AdaRound (Adaptive Rounding) 是一种训练后量化的优化方法,旨在通过微调权重的舍入方向(向上或向下取整)来最小化量化误差。标准量化使用“四舍五入”(round-to-nearest),但这在均方误差意义上并不总是全局最优。AdaRound 将每个权重的舍入决策视为一个二值变量,并利用少量校准数据,通过优化层输出重构误差来学习最佳的舍入方式。

这里实现一个简化版:给定一个线性层的权重 W、初始量化 scale,以及一小批校准输入 X,我们对每个权重逐一决策:尝试向上取整(ceil)和向下取整(floor),计算两种情况下该层输出(对校准数据)的均方误差(MSE),选择误差较小的那个舍入方向。这本质上是一种贪婪的、基于数据影响的局部优化。

import torch
import torch.nn.functional as F

def ada_round_weight(W, scales, X, bits=8):
    """
    自适应舍入:对权重 W 进行逐元素量化,通过校准数据 X 优化舍入方向。
    Args:
        W: FP32 权重,形状 [out_features, in_features]
        scales: per-channel scale,形状 [out_features]
        X: 校准输入,形状 [N, in_features]   (假设 bias=None)
        bits: 量化比特数
    Returns:
        W_int: 经过自适应舍入的 int8 权重
    """
    qmax = 2 ** (bits - 1) - 1
    out_features, in_features = W.shape

    # 首先计算每个权重的两个候选整数值:floor 和 ceil
    W_div = W / scales.unsqueeze(1)          # [out, in]
    W_floor = torch.floor(W_div)
    W_ceil = torch.ceil(W_div)
    # 钳制到量化范围
    W_floor = torch.clamp(W_floor, -qmax, qmax)
    W_ceil = torch.clamp(W_ceil, -qmax, qmax)

    # 初始量化输出(使用四舍五入作为基准)
    # 我们将逐个权重决定最终采用 floor 还是 ceil

    # 为了计算每个权重对输出的影响,我们可以预计算 X 与单位权重的乘积,但这不太现实。
    # 替代方案:利用线性层的性质:输出 Y = X W^T。
    # 权重矩阵中位置 (i,j) 的改变对输出 Y 的第 i 列(特征)的影响为 delta_W[i,j] * X[:, j]。
    # 因此,若我们将 W[i,j] 的舍入从 floor 改为 ceil,输出变化为:
    #   delta_Y = (ceil_val - floor_val) * scales[i] * X[:, j:j+1]   (添加到第 i 个输出通道)
    # 我们可以维护一个基准输出 Y_base(使用 floor 作为基准),然后对每个权重,判断改为 ceil 是否能降低 MSE。

    # 计算基准输出(全部使用 floor)
    W_base = W_floor * scales.unsqueeze(1)   # [out, in] 浮点
    Y_base = torch.matmul(X, W_base.T)        # [N, out]

    # 真实浮点输出作为目标
    Y_true = torch.matmul(X, W.T)

    # 初始化最佳舍入选择为 floor
    best_choice = torch.zeros_like(W_div, dtype=torch.bool)   # False表示floor, True表示ceil

    # 因为计算量很大(out * in 个权重),这里仅演示思路,实际可能采用迭代或分块。
    # 为简化,我们逐个权重评估(仅适用于教学)。
    for i in range(out_features):
        for j in range(in_features):
            if W_floor[i,j] == W_ceil[i,j]:
                continue   # 整数无需选择
            # 当前 base 下,尝试改用 ceil 这个权重
            delta_w = (W_ceil[i,j] - W_floor[i,j]) * scales[i]
            # 计算 Y 的改变:仅第 i 个输出通道受影响
            Y_candidate = Y_base.clone()
            Y_candidate[:, i] += delta_w * X[:, j]
            # 计算 MSE
            mse_floor = F.mse_loss(Y_base, Y_true)
            mse_ceil  = F.mse_loss(Y_candidate, Y_true)
            if mse_ceil < mse_floor:
                best_choice[i,j] = True   # 选择 ceil
                Y_base = Y_candidate      # 更新基准

    # 根据 best_choice 生成最终量化整数
    W_final_int = torch.where(best_choice, W_ceil, W_floor).to(torch.int8)
    return W_final_int

# 测试
if __name__ == "__main__":
    torch.manual_seed(0)
    out_f, in_f = 4, 8
    W = torch.randn(out_f, in_f) * 0.5
    X_calib = torch.randn(16, in_f)  # 16个校准样本
    # 先计算初始 per-channel scale(minmax)
    max_per_ch = torch.max(torch.abs(W), dim=1).values
    scales = max_per_ch / 127.0
    # 应用 AdaRound
    W_int_ada = ada_round_weight(W, scales, X_calib, bits=8)
    # 对比四舍五入的量化权重
    W_int_nearest = torch.round(W / scales.unsqueeze(1)).clamp(-128,127).to(torch.int8)
    # 计算反量化后与原始权重的 MSE
    W_ada_fp = W_int_ada.float() * scales.unsqueeze(1)
    W_near_fp = W_int_nearest.float() * scales.unsqueeze(1)
    mse_ada = F.mse_loss(W_ada_fp, W)
    mse_near = F.mse_loss(W_near_fp, W)
    print(f"四舍五入 MSE: {mse_near.item():.6f}")
    print(f"AdaRound MSE: {mse_ada.item():.6f}")

说明:上述代码逐个权重地评估舍入选择,计算复杂度高,仅用于演示。实际 AdaRound 通常使用泰勒展开近似海森矩阵,并通过迭代优化连续松弛的舍入变量,然后舍入。但此实现清晰展示了“基于重构误差决定舍入”的核心思想。


实现一个量化配置解析器

目标 输入字符串如 "W8A8_PER_CHANNEL""W4A16_GROUP128",自动解析出权重和激活的量化比特数、量化粒度(per_tensor / per_channel / per_group)、组大小等信息,生成对应的量化 scheme,并能应用到给定的线性层上(包括权重量化和激活量化)。

import re
from dataclasses import dataclass
from typing import Optional

@dataclass
class QuantizationConfig:
    weight_bits: int
    activation_bits: int
    weight_granularity: str      # "per_tensor", "per_channel", "per_group"
    activation_granularity: str  # "per_tensor", "per_token"
    group_size: Optional[int] = None   # for weight per_group

def parse_quant_config(config_str: str) -> QuantizationConfig:
    """
    解析形如 "W8A8_PER_CHANNEL", "W4A16_GROUP128", "W8A8_PER_TENSOR" 等。
    默认激活粒度为 per_tensor,除非显示指定激活粒度(目前只支持权重粒度的解析)。
    """
    # 正则匹配
    pattern = r'^W(\d+)A(\d+)_(PER_TENSOR|PER_CHANNEL|PER_TOKEN|GROUP(\d+))$'
    match = re.match(pattern, config_str)
    if not match:
        raise ValueError(f"Invalid config string: {config_str}")

    weight_bits = int(match.group(1))
    activation_bits = int(match.group(2))
    gran_str = match.group(3)

    weight_gran = "per_tensor"
    activation_gran = "per_tensor"
    group_size = None

    if gran_str == "PER_CHANNEL":
        weight_gran = "per_channel"
    elif gran_str == "PER_TOKEN":
        # 如果指定 PER_TOKEN,通常指激活粒度,这里简单处理
        activation_gran = "per_token"
    elif gran_str.startswith("GROUP"):
        weight_gran = "per_group"
        group_size = int(match.group(4))

    return QuantizationConfig(
        weight_bits=weight_bits,
        activation_bits=activation_bits,
        weight_granularity=weight_gran,
        activation_granularity=activation_gran,
        group_size=group_size
    )

# 应用配置到层
def apply_quant_config_to_linear(layer: torch.nn.Linear, config: QuantizationConfig):
    """
    根据配置对线性层的权重进行量化(模拟),并返回一个量化后的权重副本。
    """
    W = layer.weight.data
    out_f, in_f = W.shape
    qmax = 2 ** (config.weight_bits - 1) - 1

    if config.weight_granularity == "per_tensor":
        max_val = torch.max(torch.abs(W))
        scale = max_val / qmax
        W_int = torch.round(W / scale).clamp(-qmax, qmax).to(torch.int8)
        # 反量化并返回(模拟)
        W_q = W_int.float() * scale
    elif config.weight_granularity == "per_channel":
        max_per_ch = torch.max(torch.abs(W), dim=1).values
        scales = max_per_ch / qmax
        scales[scales == 0] = 1e-8
        W_int = torch.round(W / scales.unsqueeze(1)).clamp(-qmax, qmax).to(torch.int8)
        W_q = W_int.float() * scales.unsqueeze(1)
    elif config.weight_granularity == "per_group":
        group_size = config.group_size
        assert in_f % group_size == 0
        num_groups = in_f // group_size
        W_q = torch.zeros_like(W)
        for g in range(num_groups):
            start = g * group_size
            end = start + group_size
            group = W[:, start:end]
            max_val = torch.max(torch.abs(group))
            scale = max_val / qmax if max_val > 0 else 1.0
            W_q[:, start:end] = torch.round(group / scale).clamp(-qmax, qmax).float() * scale
    else:
        raise NotImplementedError

    return W_q

# 示例
config_str = "W4A16_GROUP128"
config = parse_quant_config(config_str)
print(f"Parsed: W{config.weight_bits}A{config.activation_bits} "
      f"weight_gran={config.weight_granularity}, group={config.group_size}")

layer = torch.nn.Linear(256, 512)
W_quant = apply_quant_config_to_linear(layer, config)
print("量化后权重均值:", W_quant.mean().item())

说明:该解析器支持多种配置字符串,可根据实际需要扩展。apply_quant_config_to_linear 展示了如何基于解析出的配置实现 weight 的量化模拟。


实现 KV Cache 的对称 per-token 量化

原理 在推理时对 KV Cache 进行低精度(如 INT8)存储可以大幅降低显存占用。per-token 量化对每个新生成的 token 的 Key 和 Value 向量(形状 [num_heads, head_dim])在线计算 scale,量化为 INT8,并存入缓存,同时保存每个 token 的 scale。注意力计算时,先反量化缓存的 K、V 回浮点,再进行点积。

class QuantizedKVCache:
    def __init__(self, num_layers, num_heads, head_dim, max_seq_len, bits=8):
        self.num_layers = num_layers
        self.num_heads = num_heads
        self.head_dim = head_dim
        self.max_seq_len = max_seq_len
        self.bits = bits
        self.qmax = 2 ** (bits - 1) - 1

        # 存储量化后的 K 和 V,以及每 token 的 scale
        # 形状:[num_layers][2] -> key/value,每个为 [num_heads, max_seq_len, head_dim] int8
        self.k_cache = [torch.zeros(num_heads, max_seq_len, head_dim, dtype=torch.int8) for _ in range(num_layers)]
        self.v_cache = [torch.zeros(num_heads, max_seq_len, head_dim, dtype=torch.int8) for _ in range(num_layers)]
        # scale 存储:[num_layers][2] -> [num_heads, max_seq_len] FP32
        self.k_scales = [torch.zeros(num_heads, max_seq_len, dtype=torch.float32) for _ in range(num_layers)]
        self.v_scales = [torch.zeros(num_heads, max_seq_len, dtype=torch.float32) for _ in range(num_layers)]
        self.seq_len = 0

    def append_token(self, k_new, v_new, layer_idx):
        """
        k_new, v_new: [num_heads, head_dim] 单个 token 的 K/V
        将新 token 量化为 int8 并存储到缓存的 seq_len 位置。
        """
        idx = self.seq_len
        # 计算 per-token scale (对称)
        k_scale = torch.max(torch.abs(k_new)) / self.qmax
        v_scale = torch.max(torch.abs(v_new)) / self.qmax
        # 避免除零
        if k_scale == 0: k_scale = 1.0
        if v_scale == 0: v_scale = 1.0
        # 量化
        k_int = torch.round(k_new / k_scale).clamp(-self.qmax, self.qmax).to(torch.int8)
        v_int = torch.round(v_new / v_scale).clamp(-self.qmax, self.qmax).to(torch.int8)
        # 存储
        self.k_cache[layer_idx][:, idx, :] = k_int
        self.v_cache[layer_idx][:, idx, :] = v_int
        self.k_scales[layer_idx][:, idx] = k_scale
        self.v_scales[layer_idx][:, idx] = v_scale

    def get_all_kv(self, layer_idx, seq_len=None):
        """
        获取反量化后的完整 K 和 V。
        """
        if seq_len is None:
            seq_len = self.seq_len
        # 取出 int8 张量并反量化
        k_int = self.k_cache[layer_idx][:, :seq_len, :].float()
        v_int = self.v_cache[layer_idx][:, :seq_len, :].float()
        k_scales = self.k_scales[layer_idx][:, :seq_len].unsqueeze(-1)  # [H, L, 1]
        v_scales = self.v_scales[layer_idx][:, :seq_len].unsqueeze(-1)
        k_fp = k_int * k_scales
        v_fp = v_int * v_scales
        return k_fp, v_fp

使用示例

num_layers=1; H=2; D=4; max_len=10
cache = QuantizedKVCache(num_layers, H, D, max_len)
k_tok = torch.randn(H, D)
v_tok = torch.randn(H, D)
cache.append_token(k_tok, v_tok, 0)
cache.seq_len = 1
k_all, v_all = cache.get_all_kv(0)
print(k_all.shape)  # [2,1,4]

该实现展示了 KV Cache 量化的核心思路:在线计算 scale、量化存储、按需反量化。实际引擎中为提升性能,反量化通常融合在注意力 kernel 内。