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

더 알아보기 (Learn more)