AI Briefing

vLLM CUDA 정수 오버플로우

·2026.03.25 17:24

핵심 내용

vLLM Mamba CUDA 커널의 32비트 오버플로우가 logprob 불일치를 만들었다.

1 / 2

자세히 보기

GRPO로 Jamba 3B를 학습하던 중, rollout이 계산한 logprob과 FSDP가 같은 시퀀스에서 다시 계산한 값이 맞지 않는 문제가 발생했다. 처음에는 training, inference, 동기화 중 어디가 원인인지 불명확했지만, 실제 원인은 vLLM 내부 CUDA 커널의 조용한 정수 오버플로우였다.

문제는 RL 시스템에서 흔히 쓰는 sanity check, 즉 rollout 엔진이 저장한 log π_old와 학습 모델이 같은 입력으로 재계산한 log π_train을 비교하는 과정에서 드러났다. 두 값은 같은 가중치와 같은 입력이라면 거의 같아야 하는데, 특정 step에서 주기적으로 차이가 튀면서 이상 징후가 보였다.

디버깅의 핵심 전환점은 rollouts_num이었다. 기본 설정에서는 대략 step 12 근처에서만 spike가 보였지만, rollout 수를 8, 16, 32, 64, 128로 늘리자 spike의 주기가 정확히 그 값에 맞춰 움직였다. 특히 128에서는 첫 rollout부터 차이가 나타나, 문제가 training dynamics가 아니라 rollout path 쪽에 있다는 점이 드러났다.

여기서부터 문제는 RL 학습 버그가 아니라 inference-engine 버그로 좁혀졌다. standalone vLLM 재현과 추가 ablation을 거친 뒤, GPU memory utilization을 **50%**로 낮추면 문제가 사라지고 그 이상에서는 재현된다는 사실을 확인했다. 이후 attention-only 모델에서는 문제가 나오지 않아, 원인은 KV-cache가 아니라 Mamba-1 selective scan CUDA kernel 쪽으로 수렴했다.

실제 결함은 커널의 포인터 연산에 있었다. SSMParamsBase의 index_t가 uint32_t로 정의돼 있었고, cache_index와 ssm_states_batch_stride의 곱이 32비트에서 계산되면서 UINT32_MAX를 넘는 순간 조용히 wrap-around 됐다. batch stride가 89,600일 때는 cache_index가 약 47,935를 넘으면 오버플로우가 발생했고, 총 69,776개의 cache slot 중 약 **31%**가 잘못된 메모리 위치에 쓰이면서 의도한 slot은 0으로 남았다.

수정은 단 두 글자였다.

  • using index_t = uint32_t;
  • using index_t = size_t;

이 변경으로 stride 곱셈이 64-bit arithmetic으로 승격돼 오버플로우가 사라졌다. 결론은 명확하다. 분산 RL에서는 겉으로 보이는 logprob mismatch가 학습 불안정처럼 보여도, 실제 원인은 그 아래의 저수준 CUDA 버그일 수 있으며, 문제를 푸는 가장 좋은 방법은 failure의 구조를 바꾸는 변수를 찾아 step-zero 재현까지 줄여내는 것이다.

이 한국어 요약은 AI가 자동으로 만들었습니다. 원문의 주장과 맥락은 원문에서 확인해 주세요. 저작권은 원저작자에게 있습니다.

AI 처리 방식을 확인하거나, 요약 오류와 출처 표기 문제, 삭제 요청을 문의 · 건의로 알려주세요.