为啥大模型需要量化??如何量化
大模型量化这个话题已经不算新鲜了,但很多同学对它的底层机制还停留在“知道有这么回事”的阶段。今天我们就把它掰开揉碎,先聊聊量化是什么、为什么需要,再通过简单的数学推导看看它是怎么工作的,最后上点PyTorch代码,对LLM的权重参数实际做一遍量化和反量化。
一张图说明白:下面这个图很直观——Llama 3 8B的基础模型大小是32 GB,用Int8量化后缩到8 GB(直接砍掉75%),再用Int4量化还能降到4 GB(压缩近90%)。这么大的尺寸缩减,对显存和推理速度的改善是立竿见影的。
量化主要干两件事:一是大幅降低显存占用,让你能在有限的硬件上跑更大的模型;二是提升推理性能,毕竟计算量小了。而对精度的折损,在实践中往往可以接受,属于“不用白不用”的好东西。
什么是量化,为什么需要它?
量化说白了,就是一种把大模型压缩成小模型的技巧,核心是对模型的权重参数和激活值下手。计算一张模型大小的对比图,结果一目了然:FP16的8B模型,权重占14GB,Int8直接降到8GB,Int4再减半到4GB。省下来的显存可以装更多batch,或者直接上更大的模型。
量化不仅能省显存,还能加快推理——尤其是当模型在GPU上跑的时候,低精度整数运算比浮点运算快得多。所以现在主流部署几乎都离不开量化。
量化是如何工作的?简单的数学推导
从技术上讲,量化就是把模型权重从较高精度(比如FP32)映射到较低精度(比如FP16、BF16、INT8)。方法很多,这里选最常用的线性量化来拆解,它有两种模式:非对称量化和对称量化。我们一个一个看。
A. 非对称线性量化
非对称量化将原始张量的范围 [Wmin, Wmax] 映射到量化张量的范围 [Qmin, Qmax]。看图会更清楚。
几个关键参数:
- :原始张量的最小值和最大值(FP32)。大多数现代LLM默认权重是FP32。
Wmin, Wmax
- :量化张量的最小值和最大值(这里用INT8,范围-128到127)。也可以用INT4、FP16等。
Qmin, Qmax
- :FP32类型,量化时缩小原始值,反量化时放大量化值。
缩放值(S)
- :INT8类型,量化张量中一个非零值,直接映射到原始张量中的0。
零点(Z)
那么怎么从原始值得到量化值?其实就两步公式。建议对照上面的图来推导。
需要留意的两个细节:
- 如果Z超出范围,就把它强行拉到Qmin或Qmax,用if-else判断。
- 如果Q超出范围,用PyTorch的clamp函数把它截断到[-128, 127]。
B. 对称线性量化
对称量化里,原始张量的0映射到量化张量的0。所以没有零点(Z),映射发生在 (-Wmax, Wmax) 和 (-Qmax, Qmax) 之间。看图。
数学推导如下:
非对称 vs 对称:简单说,对称量化计算简单(少一个零点),但可能浪费部分量化范围(如果原始数据分布不对称)。非对称更灵活,但多一个零点补偿。实际用哪种看情况。
LLM权重参数进行量化和反量化
量化作用于所有权重、偏置和激活层。为了演示,我们只量化权重参数。先看一眼量化前后Transformer模型中的权重变化:
16个FP32权重(512位)量化成INT8后只有128位,内存减少了75%。在大模型上这个比例更惊人。下面是FP32、INT8、UINT8在内存中的实际位分布,已经用补码算过,你可以自己验证。
非对称量化代码实现(PyTorch):
先创建一个随机4x4的FP32权重张量:
import torch
original_weight = torch.randn((4,4))
print(original_weight)
定义量化和反量化函数:
def asymmetric_quantization(original_weight):
quantized_data_type = torch.int8
Wmax = original_weight.max().item()
Wmin = original_weight.min().item()
Qmax = torch.iinfo(quantized_data_type).max
Qmin = torch.iinfo(quantized_data_type).min
S = (Wmax - Wmin) / (Qmax - Qmin)
Z = Qmin - (Wmin / S)
if Z < Qmin:
Z = Qmin
elif Z > Qmax:
Z = Qmax
else:
Z = int(round(Z))
quantized_weight = (original_weight / S) + Z
quantized_weight = torch.clamp(torch.round(quantized_weight), Qmin, Qmax)
quantized_weight = quantized_weight.to(quantized_data_type)
return quantized_weight, S, Z
def asymmetric_dequantization(quantized_weight, scale, zero_point):
dequantized_weight = scale * (quantized_weight.to(torch.float32) - zero_point)
return dequantized_weight
运行量化函数:
quantized_weight, scale, zero_point = asymmetric_quantization(original_weight)
print(f"quantized weight: {quantized_weight}")
print(f"scale: {scale}")
print(f"zero point: {zero_point}")

然后反量化:
dequantized_weight = asymmetric_dequantization(quantized_weight, scale, zero_point)
print(dequantized_weight)

计算量化误差:
quantization_error = (dequantized_weight - original_weight).square().mean()
print(quantization_error)

对称量化代码
量化的基础原理基本就是这样。下一篇文章再深入讲讲TRT-LLM中的具体量化实现。
-
- 关于宇宙的好的网名有哪些
- 角色扮演 | 1
- 网名