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

如何使用KL散度完成8比特量化校准?KLD计算负值怎么解决

问题原因排查
  • 第一,KLDivLoss输入规则错误:PyTorch的torch.nn.KLDivLoss默认要求第一个输入是取过对数的预测分布(log_Q),第二个输入是真实分布(P),你的输入顺序完全相反,且输入的是原始张量而非概率分布,必然会得到非法的负值结果,同时你之前计算的P、Q分布没有实际传入损失函数,属于无效计算。
  • 第二,伪量化逻辑错误:直接对FP32张量调用.to(torch.int8)属于硬截断/溢出操作,不是常规INT8量化流程。正常INT8量化需要先计算缩放因子scale和零点zero_point,公式为A_int8 = clamp(round(A / scale + zero_point), -128, 127),直接转类型会导致超出INT8范围的值溢出,得到的A_int8分布完全失真,后续校准没有参考意义。
  • 第三,分布计算错误:调用norm.pdf(A)默认使用标准正态分布(均值0方差1),没有适配输入张量A的真实分布统计值,计算得到的P、Q不是实际数据的概率分布,无法用于校准。

修复后的实现方案
import torch
import numpy as np
from scipy.stats import norm

# 输入FP32张量,替换为你的实际输入
A = torch.randn(1, 4, 1024, 256)

# -------------------------- 1. 统计真实数据分布 --------------------------
max_val = A.abs().max().item()
# 对FP32数据做直方图分桶,桶数可根据精度需求调整
hist_p, bin_edges = np.histogram(A.numpy().flatten(), bins=2048, range=(-max_val, max_val))
# 归一化得到真实分布P
P = hist_p / hist_p.sum()
P = torch.from_numpy(P).float()

# -------------------------- 2. 遍历候选阈值计算最优KLD --------------------------
min_kld = float('inf')
best_threshold = max_val
# 对称量化使用[-127,127]范围,预留溢出裕度避免EOS等特殊值被截断
num_quant_bins = 254

for threshold in np.linspace(max_val*0.5, max_val, 100):
    # 计算当前阈值对应缩放因子
    scale = threshold / 127
    # 执行对称量化+反量化
    A_quant = torch.clamp(torch.round(A / scale), -127, 127)
    A_dequant = A_quant * scale
    # 统计反量化后的分布Q
    hist_q, _ = np.histogram(A_dequant.numpy().flatten(), bins=2048, range=(-max_val, max_val))
    Q = hist_q / hist_q.sum()
    Q = torch.from_numpy(Q).float()
    
    # 计算KLD,加极小值避免log(0)异常
    kld = torch.nn.KLDivLoss(reduction='sum')(torch.log(Q + 1e-10), P)
    if kld < min_kld:
        min_kld = kld
        best_threshold = threshold

# -------------------------- 3. 用最优参数完成最终量化 --------------------------
best_scale = best_threshold / 127
A_int8 = torch.clamp(torch.round(A / best_scale), -127, 127).to(torch.int8)
print(f"最优阈值: {best_threshold}, 最优KLD: {min_kld.item()}")

# 后续自注意力计算注意转INT32避免溢出
B_int8 = A_int8.clone()
AB = A_int8.to(torch.int32).matmul(B_int8.transpose(-1, -2).to(torch.int32))

注意事项
  • 对称量化选择[-127,127]而非[-128,127]是为了避免零点偏移问题,同时预留小范围溢出裕度,避免NLP任务中EOS、特殊token对应的激活值被截断,解决丢失EOS的问题。
  • INT8矩阵乘计算前要先转成INT32类型,避免累加过程中溢出INT8数值范围。
  • 如果激活值分布明显不对称,可以切换为非对称量化逻辑,额外计算zero_point适配分布特征。

内容的提问来源于stack exchange,提问作者esse non videri

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.29 10:09:03