HIGGS
HIGGS
HIGGS는 Hadamard 전처리와 MSE-최적 양자화 그리드를 결합한 제로샷 양자화 알고리즘으로, 더 낮은 양자화 오차와 최신 수준의 성능을 달성합니다.
출처: 문서
본문
HIGGS는 Hadamard 전처리를 MSE-최적 양자화 그리드와 결합하여 더 낮은 양자화 오차와 최신 수준의 성능을 달성하는 제로샷 양자화 알고리즘입니다.
HIGGS의 런타임 지원은 FLUTE 라이브러리를 통해 구현됩니다. 현재 Llama 3와 Llama 3.0의 70B 및 405B 변형, Gemma 2의 8B 및 27B 변형만 지원됩니다. HIGGS는 현재 양자화 학습과 역전파를 일반적으로 지원하지 않습니다.
아래 명령으로 FLUTE를 설치하세요.
pip install flute-kernel
pip install flute-kernel -i https://flute-ai.github.io/whl/cu11.8
모델을 양자화할 비트 수를 지정해 HiggsConfig를 만드세요.
from transformers import AutoModelForCausalLM, AutoTokenizer, HiggsConfig
model = AutoModelForCausalLM.from_pretrained(
"google/gemma-2-9b-it",
quantization_config=HiggsConfig(bits=4),
device_map="auto",
)
[!TIP] 공식 ISTA-DASLab 컬렉션에서 HIGGS로 사전 양자화된 모델을 찾아볼 수 있습니다.
torch.compile
HIGGS는 torch.compile과 완전히 호환됩니다.
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer, HiggsConfig
model = AutoModelForCausalLM.from_pretrained(
"google/gemma-2-9b-it",
quantization_config=HiggsConfig(bits=4),
device_map="auto",
)
model = torch.compile(model)
RTX4090에서 Llama-3.1-8B-Instruct의 초당 forward 횟수 벤치마크는 아래 표를 참고하세요.
| Batch Size | BF16 (with torch.compile) |
HIGGS 4bit (without torch.compile) |
HIGGS 4bit (with torch.compile) |
|---|---|---|---|
| 1 | 59 | 41 | 124 |
| 4 | 57 | 42 | 123 |
| 16 | 56 | 41 | 120 |