Apple Silicon MPS에서 F.linear 디코드가 느렸던 이유와 PyTorch 2.14의 변경점 | DAKER 커뮤니티

맥에서 자동회귀 디코드가 유독 느리게 느껴졌다면, 원인이 꼭 모델 구조에만 있는 것은 아닙니다. Apple Silicon의 MPS 환경에서는 F.linear에 들어가는 입력 shape가 커널 선택에 영향을 주었고, 그 차이가 실제 지연으로 이어지기도 했습니다.
특히 활성값이 [B, 1, K] 형태로 들어가는 디코드 구간은 겉보기에는 사소해 보여도, 예전 버전에서는 빠른 경로를 벗어날 수 있었습니다. PyTorch 2.14는 바로 이 지점을 손보며, 시퀀스 길이 1 디코드에서 생기던 간극을 줄이도록 경로와 커널을 보강했습니다.
Apple Silicon MPS에서 자동회귀 디코드는 종종 활성값을[B, 1, K]형태로F.linear에 넣는데, 예전에는 이 모양이 빠른 경로를 벗어나 bf16·fp16에서 크게 느려질 수 있었습니다.
왜 shape 하나가 성능을 갈랐나
맥에서 디코드만 유난히 답답했다면, 선형층의 입력 shape를 먼저 보는 것이 좋습니다. 배치·시퀀스·특징 차원의 표기가 한 칸만 달라도 내부적으로 선택되는 커널 경로가 달라질 수 있기 때문입니다.
문제의 핵심은 자동회귀 디코드에서 자주 등장하는 L=1 상황입니다. 이때 활성값이 [B, 1, K]로 F.linear에 들어가면, 같은 연산량처럼 보여도 예전에는 빠른 경로를 타지 못하는 경우가 있었습니다. 특히 bf16·fp16에서 그 차이가 더 크게 드러날 수 있었습니다.
PyTorch 2.14에서 달라진 점
공식 정리 기준으로, 시퀀스 길이 1이 빠른 경로를 벗어나던 경우가 수정되었고, 벡터-행렬 형태를 받는 GEMV 커널이 보강되었습니다. 이 변경은 CUDA 대비 남아 있던 디코드 병목 가운데 큰 축 하나를 줄이는 방향의 개선으로 볼 수 있습니다.
PyTorch 2.14는 시퀀스 길이 1 라우팅과 GEMV 커널을 보강해 이 간극을 줄입니다.
재현할 때 확인할 조건
이 차이는 비교적 재현하기 쉬운 편입니다. MPS 장치에서 F.linear 입력 shape를 [B, 1, K]로 두고, dtype을 bf16 또는 fp16으로 맞춘 뒤, 같은 연산량을 2.13과 2.14에서 각각 재어 커널과 지연을 비교하면 됩니다.
이때 중요한 것은 프리필과 디코드를 한 덩어리로 보지 않는 일입니다. 둘은 겉으로는 같은 모델 실행처럼 보여도 병목이 다를 수 있습니다.
실무에서 같이 점검할 부분
벤치를 볼 때는 프리필(L>1)과 디코드(L=1)를 섞지 않는 것이 좋습니다. 병목이 다르기 때문에 한 그래프에 묶어 보면 원인을 놓치기 쉽습니다.
또한 MPS 관련 API에는 실험적 표시가 있으므로, 배포 노트에는 PyTorch 버전을 함께 고정해 두는 편이 안전합니다. 텐서의 디바이스와 dtype 같은 기초 조건이 흔들리면, shape 이슈를 모델 버그로 오인하기도 쉽습니다.
관련 학습 자료
기초 점검이 필요하다면 아래 자료를 함께 보면 됩니다.
참고 자료
https://daker.ai/public/learning/materials/pytorch-beginner-part1-tensors-devices-ops
https://daker.ai/public/learning/tracks/pytorch-practice-beginner-track
맥에서 디코드 성능을 볼 때, 여러분은 입력 shape와 PyTorch 버전을 얼마나 먼저 확인하는 편인가요?