[ollama] Ollama, DFlash를 통한 추론 속도 향상: 블록 단위 추론의 힘
PR 링크: ollama/ollama#17571 상태: Merged | 변경: +1524 / -251
들어가며
최근 Ollama의 ollama/ollama 레포지토리에서는 모델 추론 속도를 획기적으로 개선할 수 있는 새로운 기법인 DFlash(Draft-only FlashAttention)를 도입하는 PR이 merge되었습니다. 기존의 LLM 추론 방식은 토큰 하나하나를 순차적으로 생성하는 방식이어서, 특히 코드 생성이나 편집과 같이 반복적인 패턴이 많은 작업에서 병목 현상이 발생하곤 했습니다. 이번 PR은 이러한 문제를 해결하기 위해, 한 번의 forward pass로 여러 개의 토큰을 동시에 예측하는 '블록 단위 추론(block-diffusion speculative decoding)' 방식을 도입하여 추론 속도를 크게 향상시켰습니다.
이 글에서는 DFlash가 무엇인지, 그리고 어떻게 기존의 추론 방식을 개선하는지에 대해 코드 변경사항을 중심으로 자세히 알아보겠습니다.
코드 분석
이번 PR의 핵심은 x/mlxrunner/dflash.go 파일의 추가와 x/mlxrunner/batch/batch.go 파일의 일부 수정입니다. DFlash는 기존의 추론 방식에 '드래프트 모델(draft model)'이라는 개념을 추가하여 작동합니다.
1. x/mlxrunner/batch/batch.go - Batch 구조체 변경
가장 먼저 눈에 띄는 변경은 Batch 구조체의 Hidden 필드에 대한 주석 수정입니다. 기존에는 Hidden 필드가 '타겟 모델이 드래프트 모델과 융합하는 데 사용하는 히든 상태'라고 설명되어 있었지만, 이제는 '드래프트 모델의 forward pass를 위한 드래프트-컨디셔닝(draft-conditioning) 상태'로 명확히 변경되었습니다.
Before:
// Hidden is the target hidden state a draft model fuses with its input
// embedding for this step. It is nil for ordinary forward passes.
+
After:
// Hidden is the draft-conditioning state for a draft model's forward.
// It is nil for ordinary forward passes.
+
이 변경은 DFlash의 작동 방식을 이해하는 데 중요한 단서가 됩니다. 즉, Hidden 필드는 일반적인 모델의 순방향 연산이 아닌, 드래프트 모델이 타겟 모델의 상태를 바탕으로 다음 토큰 블록을 예측하는 데 사용된다는 것을 의미합니다.
2. x/mlxrunner/dflash.go - DFlash 구현
이 파일은 DFlash의 핵심 로직을 담고 있습니다.
-
dflashDrafter구조체:- 드래프트 모델의 블록 크기(
blockSize)와 마스크 토큰(maskToken)을 저장합니다. newDFlashDrafter함수를 통해 초기화됩니다.draftLimit()메서드는 드래프트할 수 있는 최대 토큰 수를 반환합니다 (블록 크기 - 1).
- 드래프트 모델의 블록 크기(
-
dflashDraftSession구조체:- 현재 요청의 드래프팅 세션을 관리합니다.
ctxOffset: 컨텍스트 캐시에 기록된 마지막 피처 행의 위치를 나타냅니다.pendingFeatures: 아직 처리되지 않은 피처 행들을 버퍼링합니다.pendingCount: 버퍼링된 피처 행의 개수입니다.blockOutstanding: 현재 처리 중인 블록이 있는지 여부를 나타냅니다.
-
주요 메서드 분석:
committed(tokens, features, position):- 타겟 모델이 생성한 토큰과 피처를 받아 드래프트 세션에 커밋합니다.
dflashPendingFlushTokens(기본값 256) 크기만큼 피처가 쌓이면flush()를 호출하여 처리합니다.- 이전 실행에서 이미 처리된 부분은 건너뛰어 중복 처리를 방지합니다.
settle(_ *mlx.Array):- 버퍼링된 피처 행들을 최종적으로 처리합니다.
close():- 세션 종료 시 버퍼링된 내용을 모두 처리합니다.
takePending() *mlx.Array:- 버퍼링된 피처 행들을 가져와 하나의
mlx.Array로 합치고,ctxOffset을 업데이트합니다.
- 버퍼링된 피처 행들을 가져와 하나의
commitBlock():- 현재 블록의 예측 결과를 드래프트 캐시에서 제거합니다. 이는 새로운 컨텍스트 행이 캐시의 올바른 위치에 기록되도록 보장하기 위함입니다.
flush():commitBlock()을 호출하여 이전 블록을 정리한 후,takePending()으로 가져온 피처들을 사용하여 드래프트 모델의Forward를 호출합니다.mlx.AsyncEval(state...)를 통해 캐시 쓰기를 강제하여, 드래프트가 발생하지 않더라도 피처가 계속 누적되는 것을 방지합니다.
propose(current *mlx.Array, maxTokens int) *draftCandidates:- DFlash의 핵심 기능으로, 현재 토큰(
current)을 기준으로 최대maxTokens개수만큼의 토큰 블록을 예측합니다. blockSize - 1개의 마스크 토큰을 사용하여 블록을 생성합니다.- 드래프트 모델의
Forward메서드를 호출하여 히든 상태와 보조 히든 상태를 얻습니다. spec.draft.Unembed를 통해 로짓(logits)을 추출하고, 샘플러를 사용하여 다음 토큰 후보(draftCandidates)를 생성합니다.
- DFlash의 핵심 기능으로, 현재 토큰(
3. x/mlxrunner/dflash_test.go - 테스트 코드
fakeBlockDraft구조체:- 실제 DFlash 모델의 동작을 모방하는 테스트용 더미(dummy) 구현입니다.
Forward메서드에서 입력된 컨텍스트와 블록 토큰을 기록하고, 예측 결과를 반환합니다.LoadWeights,NewCaches,BlockParams,Unembed등의 인터페이스 메서드를 구현합니다.
newBlockTestSession함수:- 테스트를 위한
Runner,fakeBlockDraft,dflashDraftSession등을 설정합니다.
- 테스트를 위한
TestDFlashCommittedBuffersPastFlushCap:dflashPendingFlushTokens크기만큼의 데이터가 쌓였을 때flush가 올바르게 동작하는지, 그리고 그 이후 버퍼링된 데이터가settle을 통해 처리되는지를 검증합니다.
이 테스트 코드는 DFlash의 내부 동작, 특히 버퍼링 및 플러시 메커니즘이 예상대로 작동하는지 확인하는 데 중요한 역할을 합니다.
왜 이게 좋은가?
이번 PR에서 도입된 DFlash는 다음과 같은 이유로 매우 좋은 최적화/개선이라고 할 수 있습니다.
-
추론 속도 향상:
- 가장 큰 장점은 추론 속도 향상입니다. 기존의 토큰 단위 생성 방식은 각 토큰마다 모델의 forward pass가 필요했지만, DFlash는 한 번의 forward pass로 여러 토큰을 동시에 예측합니다. 이는 특히 코드 생성이나 편집과 같이 반복적인 패턴이 많은 작업에서 큰 성능 향상을 가져옵니다.
- PR 설명에 따르면,
laguna-s모델의 경우 코드 편집 작업에서 +59%의 속도 향상을 보였습니다. 이는 DFlash가 타겟 모델의 예측을 얼마나 자주 수용하는지에 따라 달라지며, 수용률이 높을수록 속도 향상 폭이 커집니다.
-
안정성 보장:
- DFlash는 '추측적 디코딩(speculative decoding)' 기법을 사용합니다. 이는 드래프트 모델이 예측한 토큰 블록을 먼저 생성한 후, 타겟 모델이 이 예측을 검증하는 방식입니다. 만약 드래프트 모델의 예측이 틀리면, 해당 부분은 타겟 모델이 직접 다시 생성하므로 결과의 정확성은 보장됩니다.
- PR 설명에 "DFlash never makes decode slower than not having it."라는 문구가 있는 것처럼, 최악의 경우에도 기존 방식보다 느려지지 않도록 설계되었습니다. 이는 드래프트 모델의 예측이 틀릴 경우, 드래프트 깊이(depth)를 0으로 줄여 일반적인 속도로 돌아가기 때문입니다.
-
유연한 아키텍처:
- DFlash는 특정 모델 아키텍처에 종속되지 않고, 다양한 타겟 모델과 함께 사용할 수 있도록 설계되었습니다.
x/models/dflash에 정의된 드래프트 모델은 어떤 타겟 모델과도 페어링될 수 있습니다. - 기존의 추론 메커니즘(rollback, adaptive depth)을 재사용하여 점진적인 개선을 가능하게 합니다.
- DFlash는 특정 모델 아키텍처에 종속되지 않고, 다양한 타겟 모델과 함께 사용할 수 있도록 설계되었습니다.
일반적 교훈:
- 병렬 처리의 힘: LLM 추론에서 병렬 처리는 여전히 큰 개선의 여지가 있는 영역입니다. DFlash는 '블록 단위 예측'이라는 아이디어를 통해 이를 효과적으로 활용했습니다.
- 추측적 실행(Speculative Execution): 예측 모델을 사용하여 계산량을 줄이는 것은 게임 개발, 컴파일러 최적화 등 다양한 분야에서 사용되는 강력한 기법입니다. LLM 추론에도 성공적으로 적용될 수 있음을 보여줍니다.
- 점진적 개선: 기존 시스템의 핵심 로직을 유지하면서 새로운 기능을 추가하고 성능을 개선하는 방식은 안정성과 효율성을 동시에 확보하는 좋은 전략입니다.
References
- mlx.Array documentation - MLX 라이브러리의 핵심 데이터 구조인
mlx.Array에 대한 문서입니다. - mlx.Slice documentation -
mlx.Array의 슬라이싱 연산에 대한 문서입니다. - mlx.Concatenate documentation - 여러
mlx.Array를 합치는 함수에 대한 문서입니다. - mlx.Pin/Unpin documentation - 메모리 관리를 위한
Pin및Unpin함수에 대한 문서입니다. - mlx.AsyncEval documentation - 비동기 연산 평가 함수에 대한 문서입니다.
- Ollama API Documentation - Ollama의 API에 대한 전반적인 문서입니다. (DFlash 자체의 API는 아니지만, Ollama의 추론 과정과 관련)
참고 자료
- https://github.com/ml-explore/mlx/blob/main/docs/python/mlx.md#mlxarray
- https://github.com/ml-explore/mlx/blob/main/docs/python/mlx.md#mlxslice
- https://github.com/ml-explore/mlx/blob/main/docs/python/mlx.md#mlxconcatenate
- https://github.com/ml-explore/mlx/blob/main/docs/python/mlx.md#mlxpinu
- https://github.com/ml-explore/mlx/blob/main/docs/python/mlx.md#mlxasync
- https://github.com/ollama/ollama/blob/main/docs/api.md
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [ollama] Ollama MLX Sampler 최적화: 성능 향상과 Logprobs 지원
- [sglang] Apple Silicon LLM 성능 향상: 슬라이딩 윈도우 KV 캐싱 및 인-그래프 샘플링 도입
- [ollama] Ollama MLX Gemma4 성능 최적화: Fused Operations를 통한 효율성 증대
- [vllm] vLLM, DeepSeek-V3.2/GLM-5.2 MTP 경로 최적화: All-Reduce 융합 및 로컬 Argmax 도입
- [vllm] vLLM에 Dots3 NOTE 모델 네이티브 지원 추가: 멀티모달 및 하이브리드 MLA 최적화
PR Analysis 의 다른글
- 이전글 [sglang] SGLang 캐시 시뮬레이터의 병목 해결: 직렬 전송에서 동시성 전송으로의 최적화
- 현재글 : [ollama] Ollama, DFlash를 통한 추론 속도 향상: 블록 단위 추론의 힘
- 다음글 [sglang] [Kimi K3] CPU 전송 이미지의 지연 전처리(Deferred Preprocessing)를 통한 VLM 성능 최적화
댓글