브레이커블 CUDA 그래프
브레이커블 CUDA 그래프 (Breakable CUDA Graph)
표준 CUDA 그래프는 전체 forward pass를 하나의 불투명한 그래프로 캡처하는데, 이것은 디버깅을 어렵게 하고 일부 연산과는 호환되지 않습니다. Breakable CUDA Graph는 특정 지점에 그래프 중단(graph break)을 삽입할 수 있게 해 두 문제를 모두 해결해요. CUDA 그래프 성능 이점의 대부분을 유지하면서 대상 연산만 그래프 밖에서 실행합니다.
출처: 문서
본문
동기 (Motivation)
표준 CUDA 그래프는 전체 forward pass를 단일의 불투명한 그래프로 캡처합니다. 성능에는 좋지만 두 가지 문제를 만듭니다:
-
디버깅이 어렵다 (Debugging is hard). 캡처된 그래프 안에서 무언가 잘못되면(잘못된 출력, 수치 불일치, 크래시) 그래프가 단일체(monolithic)로 재생되기 때문에 연산을 단계별로 살펴보거나 print 문을 삽입할 방법이 없습니다.
-
일부 연산은 호환되지 않는다 (Some ops are incompatible). 특정 연산들 — 동적 제어 흐름, 호스트-디바이스 동기화, JIT 컴파일, 또는 반복마다 동작이 바뀌는 연산 — 은 CUDA 그래프에 전혀 캡처할 수 없습니다. 오늘날 유일한 해결책은 CUDA 그래프를 완전히 비활성화하는 것인데, 그러면 모델의 나머지 부분에 대한 커널 런치 오버헤드 절감을 희생합니다.
Breakable CUDA Graph는 특정 지점에 그래프 중단을 삽입할 수 있게 해 두 문제를 모두 해결합니다. 계산이 여러 캡처된 그래프 세그먼트로 분할되고 그 사이에서 eager(비그래프) 실행이 이뤄집니다. 이는 CUDA 그래프 성능 이점의 대부분을 보존하면서 대상 연산을 그래프 밖에서 실행하게 합니다.
사용법 (Usage)
디버그 모드: 모든 것을 Eager로 실행 (Debug Mode: Run Everything Eagerly)
가장 간단한 사용 사례는 디버깅입니다. --debug-cuda-graph 플래그는 전체 decode forward pass를 그래프 중단으로 감싸 모든 연산을 eager로 실행하면서도 전체 CUDA 그래프 캡처/재생 코드 경로를 거칩니다. 이를 통해 모델 코드를 바꾸지 않고 CUDA 그래프 문제를 디버깅할 수 있어요.
python -m sglang.launch_server \
--model meta-llama/Llama-3.1-8B-Instruct \
--debug-cuda-graph
이 모드는 디버깅 전용입니다. 모든 연산이 eager로 실행되므로 CUDA 그래프의 성능 이점을 제거해요.
모델 코드의 선택적 그래프 중단 (Selective Graph Breaks in Model Code)
프로덕션에서는 @eager_on_graph 데코레이터를 사용해 특정 함수를 "non-graphable"로 표시할 수 있습니다. CUDA 그래프 캡처 중 이 함수들은 캡처된 그래프 세그먼트 사이에서 eager로 실행됩니다. 캡처 밖에서는 정상적으로 동작합니다.
from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import eager_on_graph
@eager_on_graph(enable=True)
def my_dynamic_op(x):
# This op is incompatible with CUDA graph capture
return some_dynamic_operation(x)
break_graph() 헬퍼로 계산 없는 그래프 중단도 삽입할 수 있습니다:
from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import break_graph
def forward(self, x):
x = self.layer1(x)
break_graph() # force a segment split here
x = self.layer2(x)
return x
환경 레벨에서(디버그 모드 없이) 브레이커블 CUDA 그래프를 활성화하려면 환경 변수를 설정하세요:
export SGLANG_USE_BREAKABLE_CUDA_GRAPH=1
python -m sglang.launch_server \
--model meta-llama/Llama-3.1-8B-Instruct
서버 인자 (Server Args)
| Argument | Default | Description |
|---|---|---|
--debug-cuda-graph |
False |
디버그/eager 모드 활성화. 전체 forward pass를 그래프 중단으로 감싸 모든 op가 capture/replay 경로를 통해 eager로 실행되게 함. |
SGLANG_USE_BREAKABLE_CUDA_GRAPH |
0 |
환경 변수. 디버그 모드 없이 브레이커블 CUDA 그래프 활성화. @eager_on_graph 데코레이터가 효과를 발휘하려면 필요. |
동작 방식 (How It Works)
캡처 (Capture)
Breakable CUDA 그래프는 PyTorch의 torch.cuda.CUDAGraph를 확장해 단일 캡처를 그래프 중단으로 구분된 여러 세그먼트로 나눕니다.
캡처 중 흐름은 다음과 같습니다:
Begin capture (segment 1)
... graphable ops ...
@eager_on_graph function encountered:
1. End current capture segment
2. Run the function eagerly (allocates output tensors)
3. Record the function for later replay
4. Begin new capture segment
... more graphable ops ...
End capture (segment N)
각 세그먼트는 독립적으로 CUDA 그래프 실행기(executable)로 인스턴스화됩니다. 비그래프 함수와 인자 참조는 재생을 위해 저장됩니다.
재생 (Replay)
재생 중:
For each segment i:
1. Launch CUDA graph segment i
2. Run the recorded non-graph function i eagerly
Launch final CUDA graph segment
비그래프 함수는 캡처 시점과 동일한 텐서 참조로 다시 호출됩니다. 이 참조들이 CUDA 그래프의 정적 입력/출력 버퍼를 가리키므로, 재생마다 갱신된 값을 봅니다.
출력 쓰기백 (Output Writeback)
비그래프 함수가 재생 중 출력을 만들 때, 결과는 다운스트림 그래프 세그먼트가 참조하는 동일한 텐서 버퍼에 다시 써야 합니다. 메커니즘은 다음을 처리합니다:
- 일반 텐서 (Plain tensors): 원래 버퍼로 인플레이스
copy_(). - 구조적 출력 (Structured outputs) (dataclass, 텐서 속성을 가진 객체): 텐서 필드는 인플레이스로 복사, 비텐서 필드는 교체.
- 텐서 딕셔너리 (Dicts of tensors): 텐서 값은 인플레이스 복사, 비텐서 값은 교체.
스트림 포크/조인 추적 (Stream Fork/Join Tracking)
일부 모델은 작업을 보조 CUDA 스트림으로 포크합니다(예: 겹치는 계산용). Breakable CUDA 그래프는 torch.cuda.Stream.wait_stream을 후킹해 캡처 스트림에서 어떤 스트림이 포크되었는지 추적합니다. 그래프 중단이 발생하면 모든 포크된 스트림이 세그먼트 캡처를 끝내기 전에 자동으로 조인되고, 다음 세그먼트 시작 후 다시 포크됩니다.
호환성 (Compatibility)
- CUDA 및 ROCm/HIP. Breakable CUDA 그래프는 NVIDIA와 AMD GPU 모두에서 동작합니다. 다른 플랫폼(NPU, CPU, MPS, XPU)은 지원되지 않으며, 거기서는
--debug-cuda-graph가 경고와 함께 자동 비활성화됩니다. - NVIDIA에서
cuda-python필요. 스트림 캡처 상태 조회는cuda.bindings로 CUDA 런타임을 사용합니다(pip install cuda-python). 이식 가능한torch.cuda.is_current_stream_capturing()은 CUDA에서 신뢰할 수 없는 것으로 드러났습니다.cuda-python이 없는 ROCm/HIP에서는torch.cudaAPI(HIP 런타임에 매핑)를 대신 사용합니다. - 메모리 세이버 모드와 호환되지 않음.
SGLANG_MEMORY_SAVER_CUDA_GRAPH와 함께 사용할 수 없습니다.
성능 (Performance)
그래프 중단이 삽입되지 않으면 브레이커블 CUDA 그래프는 표준 CUDA 그래프에 비해 오버헤드가 거의 없습니다. 캡처/재생 경로가 거의 동일해요.
각 그래프 중단은 다음을 추가합니다:
cudaGraphLaunch호출 하나(중단 전 세그먼트 재생용)- eager Python 함수 호출 하나
- 캡처 중
cudaStreamBeginCapture/cudaStreamEndCapture쌍 하나
그래프 중단 수가 적은 일반적인 사용 사례에서는 캡처된 세그먼트에서 절약되는 커널 런치 오버헤드와 비교해 그 오버헤드는 무시할 만합니다.
코드 참조 (Code Reference)
| File | Description |
|---|---|
python/sglang/srt/model_executor/runner_backend_utils/breakable_cuda_graph/breakable_cuda_graph.py |
핵심 구현: eager_on_graph, BreakableCUDAGraph, BreakableCUDAGraphCapture |
python/sglang/srt/model_executor/runner_backend_utils/breakable_cuda_graph/cuda_utils.py |
CUDA 런타임 바인딩 유틸 (NVIDIA stream-capture 조회) |
python/sglang/srt/model_executor/runner_backend/breakable_cuda_graph_backend.py |
CUDA 그래프 러너 백엔드와의 통합 |
python/sglang/srt/server_args.py |
--debug-cuda-graph 플래그와 환경 변수 처리 |
python/sglang/srt/environ.py |
SGLANG_USE_BREAKABLE_CUDA_GRAPH 환경 변수 정의 |