LLM을 1.58비트로 파인튜닝: BitNet을 이용한 극단적인 양자화
Hugging Face는 BitNet 아키텍처를 사용하여 기존 대형 언어 모델(LLM)을 1.58비트 정밀도로 파인튜닝하는 방법을 개발했습니다. 이 접근 방식은 모델이 매개변수를 세 가지 값(-1, 0, 1)만으로 표현하도록 하여, 1비트 모델을 처음부터 사전 학습하는 데 일반적으로 필요한 막대한 예산 없이도 계산 및 에너지 비용을 크게 줄일 수 있습니다.
BitNet 아키텍처와 1.58비트 양자화
BitNet은 멀티헤드 어텐션 및 피드포워드 네트워크의 표준 Linear 레이어를 BitLinear 레이어로 교체합니다. 이러한 레이어는 가중치에 삼진 정밀도, 활성화에 8비트 정밀도를 활용합니다.
계산 패러다임
표준 LLM(예: Llama)이 FP16 덧셈 및 곱셈에 의존하는 것과 달리, BitNet b1.58은 행렬 곱셈에 INT8 덧셈을 사용합니다. 이 계산 방식의 전환은 Llama 기준 대비 행렬 곱셈의 에너지 소비를 이론적으로 71.4배 감소시킵니다.
Straight-Through Estimator (STE)를 이용한 학습
삼진 양자화에 사용되는 round() 함수는 미분이 불가능하기 때문에, BitNet은 Straight Through Estimator (STE) 를 사용합니다. STE는 반올림 연산의 그래디언트를 1로 근사하여, 마치 항등 함수인 것처럼 연산을 통과하도록 함으로써 표준 그래디언트 기반 최적화를 가능하게 합니다.
양자화 메커니즘
- Weights: 대칭 per-tensor 양자화를 사용합니다. 스케일은 가중치 행렬의 평균 절대값의 역수입니다. 가중치는 스케일링, 반올림, -1과 1 사이로 클램프된 뒤 다시 스케일링됩니다.
- Activations: absmax per-token 양자화를 사용해 8비트 정밀도로 양자화되며, 값은 [-128, 127] 범위로 스케일링됩니다. 레이어 정규화(LN)는 활성화 양자화 전에 적용되어 출력 분산을 유지합니다.
기존 모델을 1.58비트로 파인튜닝
Hugging Face는 Llama 3 8B 모델을 1.58비트 정밀도로 성공적으로 파인튜닝했습니다. 초기 실험에서는 BitLinear 레이어를 급격히 도입하면 모델이 사전 학습된 정보를 거의 모두 잃어버려 손실이 급증하는 현상이 나타났습니다.
동적 워밍업 양자화
사전 지식 손실을 방지하기 위해 Hugging Face는 양자화를 점진적으로 도입하기 위해 동적 $\lambda$ (lambda) 값을 구현했습니다:
$$\lambda = \min\left(\frac{\text{training_step}}{\text{total_training_steps}}, 1\right)$$
원본 값과 양자화된 값 사이의 차이를 $\lambda$ 로 스케일링함으로써 모델은 전체 정밀도($\lambda=0$)에서 전체 양자화($\lambda=1$)로 전환됩니다. 이 선형 스케줄러는 더 나은 수렴을 이끌어냈으며 TinyStories 데이터셋에서 약 4의 퍼플렉시티를 달성했습니다.
스케일링 및 일반화
모델이 일반 지식을 유지하고 작은 데이터셋에 과적합되지 않도록 팀은 학습을 FineWeb-edu 데이터셋으로 확장했습니다. 학습률 1e-4, 배치 크기 200만 토큰, 총 100억 토큰을 사용했을 때 모델은 WikiText 퍼플렉시티 12.2를 기록했습니다.
추가로 1000억 토큰까지 스케일링한 결과, 일부 메트릭에서는 원본 Llama 3 8B와 거의 동등한 성능을 보였지만, 전반적으로는 전체 정밀도 기준보다 약간 뒤처지는 경향을 보였습니다.
성능 벤치마크 및 결과
1.58비트 아키텍처로 파인튜닝된 모델은 HF1BitLLM 조직 아래 공개되었습니다.
주요 결과
- Competitive Performance: 100억 토큰으로 파인튜닝한 1.58비트 Llama 3 8B 모델은 BitNet 7B 모델(1000억 토큰 학습) 및 FBI LLM(1.26조 토큰 디스틸링)보다 뛰어난 성능을 보였습니다.
- MMLU Benchmarks: 개발된 8B 모델은 MMLU 벤치마크에서 Llama 1 7B 모델을 능가했습니다.
- Model Size: 가중치를
int8텐서에 패킹함으로써 파라미터 수를 80억에서 28억으로 감소시켰습니다.
추론 최적화 및 커스텀 커널
1.58비트 가중치의 속도 및 메모리 이점을 실현하기 위해 Hugging Face는 행렬 곱셈 중 가중치를 실시간으로 언패킹하는 맞춤형 CUDA 및 Triton 커널을 구현했습니다.
타일링 매트릭스 곱셈
메모리 대역폭 병목과 중복 데이터 접근을 극복하기 위해 팀은 tiling을 사용했습니다. 이 기법은 행렬을 GPU의 빠른 공유 메모리에 맞는 작은 서브 행렬(타일)로 나누어, 느린 전역 메모리 접근 빈도를 줄입니다.
커널 벤치마킹
- Triton vs. Torch: 맞춤형 Triton 커널은 BF16 정밀도로
@torch.compile과 거의 동등한 성능을 달성했습니다. - BitBlas: 팀은 혼합 정밀도 소프트웨어 라이브러리인 BitBlas가 맞춤 Triton 커널과 Torch의
matmul함수를 저정밀도에서 모두 능가한다는 것을 발견했지만, 커널 컴파일로 인한 초기 로딩 시간이 더 길어졌습니다.
Transformers와의 통합
통합은 transformers 라이브러리의 새로운 "bitnet" 양자화 방법을 통해 처리됩니다. 표준 Linear 레이어는 특수화된 BitLinear 레이어로 교체됩니다. API는 변경되지 않아 사용자는 AutoModelForCausalLM.from_pretrained를 사용해 모델을 로드할 수 있습니다.