Bamba-9B: 추론 효율적인 하이브리드 Mamba2 모델

TL;DR

Bamba-9B는 IBM, Princeton, CMU, UIUC가 발표한 추론 효율적인 하이브리드 Mamba2 모델입니다. vLLM에서 표준 트랜스포머에 비해 2.5배 높은 처리량과 2배 낮은 지연 시간을 보여주며, transformers, vLLM, TRL, llama.cpp에서 즉시 사용할 수 있습니다.

동기

트랜스포머 추론은 컨텍스트 길이가 증가함에 따라 커지는 KV‑cache 병목 현상에 제한됩니다. 하이브리드 Mamba2 아키텍처는 KV‑cache 크기를 일정하게 유지하여 이 병목을 해결합니다. Bamba‑9B는 완전 공개 데이터를 사용해 7B‑10B 규모에서 하이브리드 Mamba2 접근법을 검증하고, 커뮤니티 실험을 장려하기 위해 재현 가능한 체크포인트를 제공합니다.

transformers에서 사용하기

🤗 Transformers 라이브러리로 Bamba‑9B를 실행하려면 AutoModelForCausalLM과 AutoTokenizer로 모델과 토크나이저를 로드한 뒤 generate를 호출합니다. 예시 코드:

from transformers import AutoModelForCausalLM, AutoTokenizer
model = AutoModelForCausalLM.from_pretrained("ibm-fms/Bamba-9B
tokenizer = AutoTokenizer.from_pretrained("ibm-fms/Bamba-9B
txt = ["Mamba is a snake with following properties  
inputs = tokenizer(txt, return_tensors='pt', return_token_type_ids=False)
out = model.generate(**inputs, max_new_tokens=64)
print(tokenizer.batch_decode(out, skip_special_tokens=True)[0])

평가

최신(SOTA) 트랜스포머 모델과 비교

Bamba‑9B는 HF OpenLLM v1 리더보드에서 평균 62.31점을 기록했으며, Meta Llama 3.1 8B(63.51)보다 약간 낮지만 일부 지표에서는 Olmo2 7B(66.17)와 IBM Granite v3 8B(67.47)보다 높습니다. v2 리더보드에서는 Bamba‑9B가 평균 10.91로 Llama 3.1 8B(14.27)보다 낮지만, 수학 및 MMLU 과제를 제외하면 비슷한 수준이며, 이 경우 평균은 Llama 3.1 8B의 44.68에 근접한 45.53을 보입니다.

비슷한 토큰 예산으로 학습된 트랜스포머 모델과 비교

Bamba‑9B(2.2조 토큰)는 평균 62.31점으로 2조 토큰으로 학습된 Olmo1.5 7B(55.8)보다 우수합니다. 2조 토큰 Bamba‑9B 체크포인트(59.11)를 Llama2 7B(53.78)와 IBM Granite 7B(52.07)와 비교해도 Bamba‑9B가 더 높아, 동일한 데이터 조건에서 경쟁력이 있음을 보여줍니다.

다른 Mamba/Mamba2 모델과 비교

Bamba‑9B는 평균 62.31점으로 NVIDIA Mamba2 Hybrid 8B*(58.78)와 Zamba 7B(64.36)보다 일부 지표에서 높으며, Falcon Mamba 7B(65.31)보다 낮습니다. 표는 하이브리드 Mamba2 모델이 경쟁력 있는 결과를 제공하면서 이론적으로 최대 5배의 추론 효율성을 제공할 수 있음을 보여줍니다.

추론 효율성

Bamba‑9B는 NVIDIA H100 80GB GPU에서 vLLM을 사용해 Meta Llama 3.1 8B에 비해 최대 2.5배 높은 처리량과 2배 낮은 지연 시간을 달성했으며, 배치 크기와 시퀀스 길이 1K~64K 토큰에 걸쳐 측정되었습니다. 연산 강도 분석에 따르면 디코드 단계가 메모리 제한에 걸릴 경우 최대 5배 가속이 가능할 것으로 예측됩니다. 현재 vLLM 결과는 청크형 pre‑fill 지원 부족, 트랜스포머 방식 메모리 할당, H100용 최적화되지 않은 Mamba2 커널 등에 의해 제한됩니다.

모델 아키텍처

Bamba‑9B는 총 32개의 레이어를 사용합니다: 전체 어텐션 레이어 3개와 Mamba2 레이어 29개이며, MLP 확장 계수는 3.5, vocab 크기는 128k, RoPE 임베딩 및 GQA(8 KV‑heads, 32 heads)를 포함합니다. NVIDIA 하이브리드 Mamba2 8B 모델과 비교해 Bamba‑9B는 어텐션 레이어를 4개에서 3개로 줄이고 RoPE를 추가했습니다.

데이터

학습은 첫 번째 단계에서 Dolma v1.7을 사용하고, 이후 FineWeb‑edu와 Cosmopedia를 이용해 추가로 200B 토큰을 사용했습니다. 모든 데이터는 Ray 프레임워크를 활용해 내부 Red Hat OpenShift 클러스터에서 토크나이즈되었습니다. 1단계 데이터 혼합은 블로그 게시물에 시각화되어 있습니다.

사전 학습

사전 학습은 여러 단계로 진행되었습니다: 1.8B/100B 토큰에서의 절제 실험, 그 다음 Dolma와 함께 3B/2T 토큰, 이어 9B/2T 토큰 실행, 마지막으로 FineWeb‑edu와 Cosmopedia를 사용한 200B 토큰 미세조정 단계. 학습 하이퍼파라미터는 코사인 LR 스케줄, 피크 3e‑4, 2000 스텝에 걸친 2차 워밍업, 감쇠 0.033, 최종 LR 1e‑5, AdamW(β1=0.9, β2=0.95), 가중치 감쇠 0.1, 시퀀스 길이 4096, 전역 배치 크기 1.5M 토큰, IBM Cloud Vela에서 192개의 A100 GPU를 사용해 약 2개월 동안 진행되었습니다. 배포 오류와 하드웨어 고장으로 인해 작업이 세 번 중단되었으며, 이는 Autopilot 시스템에 의해 감지되었습니다.

데이터 로더

배포된 상태 저장 데이터 로더는 체크포인트 재개 가능, 자동 재조정, 오버헤드 없는 셔플 스트리밍, 피어‑투‑피어 트래픽 없는 비동기 분산 운영, 동적 데이터 혼합 및 실시간 토크나이제이션을 제공하며, PyTorch 네이티브, 모듈식, 확장 가능합니다. 수백 개의 학습 작업에서 실전 테스트를 거쳤으며 Torch Titan과 통합되었습니다.

양자화

FMS Model Optimizer 프레임워크와 llm‑compressor를 사용해 Bamba‑9B 체크포인트를 fp8로 양자화했으며, 정확도 손실은 거의 없습니다: OpenLLM v1 평균이 62.31에서 61.5(−0.1)로, v2 평균이 10.91에서 10.04(−0.9)로 감소했습니다. vLLM에서 fp8 추론을 활성화하려면 Mamba2 레이어용 커널 업데이트가 필요합니다.

컨텍스트 길이 확장

Full‑attention 레이어에 LongRoPE를 적용해 Bamba‑9B의 컨텍스트 길이를 확장했습니다. 초기 PhoneBook 검색 테스트에서는 튜닝 없이 16K 토큰까지 확장 모델이 기본 Bamba‑9B, Llama2‑7B, Llama3‑8B보다 우수하고, Llama3.1‑8B와 성능이 일치함을 보여줍니다. 32K 토큰에서는 Llama3.1‑8B가 앞섭니다.

요약

Bamba‑9B는 IBM, Princeton, CMU, UIUC가 2.2조 개의 공개 토큰으로 학습한 하이브리드 Mamba2 모델로, vLLM에서 Llama 3.1 8B 대비 2.5배 높은 처리량과 2배 낮은 지연 시간을 제공하며, transformers, vLLM, TRL, llama.cpp에서 즉시 사용할 수 있습니다. 또한 학습, 튜닝, 확장 사전 학습 레시피와 상태 저장 데이터 로더를 함께 제공합니다.

향후 작업

향후 계획에는 추가 데이터에 대한 지속적인 사전 학습, 커뮤니티가 제안한 혼합을 활용한 SFT, Tulu‑3, Orca‑AgentInstruct, Daring‑Anteater 데이터셋을 이용한 지도 미세조정, vLLM에서 청크형 pre‑fill 및 적절한 메모리 할당 지원, 빠른 추론을 위한 fp8 커널 추가, torch.compile 및 fp8 학습 적용, 그리고 컨텍스트 길이를 1M 토큰 이상으로 확장하는 것이 포함됩니다.

기여자

데이터 수집 및 정제: AllenAI (Dolma)와 Hugging Face (FineWeb‑edu, Cosmopedia). 데이터 전처리: IBM 팀원 Tuan Hoang Trong, Syed Zawad, Jay Gala, Ryan Gordon이 IBM Data Prep Kit을 사용. 모델 아키텍처: Tri Dao (Princeton), Albert Gu (CMU), Linsong Chu (IBM), Davis Wertheimer (IBM), Minjia Zhang (UIUC), Mudhakar Srivatsa (IBM), Raghu Ganti (IBM). 모델 학습: IBM 팀원 Linsong Chu, Divya Kumari, Davis Wertheimer, Raghu Ganti, Dakshi Agrawal. 모델 튜닝: IBM 팀원 Sukriti Sharma, Anh Uong (TRL을 통해). 모델 추론: IBM 및 커뮤니티 기여자 Fabian Lim, Antoni Viros i Martin, Adnan Hoque, Jamie Yang, Nelson Nimura Gonzalez, Joshua Rosenkranz, Nick Hill, Gabe Goodhart. 양자화: IBM 팀원 Naigang Wang, Charlie Liu. 평가: IBM 평가 팀 (Yotam Perlitz, Ofir Arviv, Michal Shmueli‑Scheuer, Haoechen Shen, Minjia Zhang (UIUC) 주도). 리더십 감사: Priya Nagpurkar, David Cox, Sriram Raghavan, Aya Soffer, Ruchir Puri, Mukesh Khare. 커뮤니티 감사: Pablo Montalvo‑Leroux, Aritra Roy Gosthipaty, Vaibhav Srivastav (Hugging Face), Stas Bekman (Contextual AI), Tyler Michael Smith (Neural Magic). 또한 Meta PyTorch, AllenAI, Hugging Face에 오픈소스 기여에 감사드립니다.

부록: 연산 강도

부록에서는 어텐션 및 Bamba 모델에 대한 연산과 메모리 방정식을 도출하여, Bamba‑9B의 디코드 단계 메모리 이점이 16K 토큰 이상의 긴 시퀀스에서 Llama 대비 최대 5배의 속도 향상을 가져올 수 있음을 보여줍니다. 현재 vLLM에서 측정된 2.5배 처리량 및 2배 지연은 청크형 pre‑fill 부재, 트랜스포머 방식 메모리 할당, H100용 최적화되지 않은 Mamba2 커널에 의해 제한됩니다.

Sources