vLLM CUDA 정수 오버플로우
핵심 내용
vLLM Mamba CUDA 커널의 32비트 오버플로우가 logprob 불일치를 만들었다.
자세히 보기
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가 자동으로 만들었습니다. 원문의 주장과 맥락은 원문에서 확인해 주세요. 저작권은 원저작자에게 있습니다.