본문 바로가기
C.W.K.
Stream
Lesson 03 of 06 · published

MLX와 PyTorch MPS — 통역사와 원어민

~14 min · pytorch-mps, comparison, performance

Level 0호기심
0 XP0/51 lessons0/15 achievements
0/100 XP to next level100 XP to go0% complete

통역사와 원어민의 차이

PyTorch MPS는 NVIDIA 하드웨어를 겨냥해 작성한 PyTorch 코드를 Apple Silicon의 Metal 백엔드에서 돌려 주는 번역층이야. 각 PyTorch 연산을 대응하는 Metal Performance Shaders 연산에 맡기고, MPS 구현이 없으면 기본적으로 오류를 내고, PYTORCH_ENABLE_MPS_FALLBACK=1을 설정한 경우에만 CPU로 폴백하지. 번역은 잘 작동하지만 GPU를 별도 메모리를 가진 장치로 보는 PyTorch의 세계관도 그대로 물려받아.

MLX는 처음부터 Apple Silicon을 위해 만든 프레임워크야. 통합 메모리와 지연 실행 그래프, 함수 변환은 기존 API에 나중에 덧댄 기능이 아니라 설계의 중심이지.

코드와 실행에서 드러나는 차이

  • API — PyTorch MPS는 익숙한 torch.tensor, tensor.to('mps'), 테이프 방식 자동 미분을 그대로 써. MLX는 mx.array를 쓰고 .to()가 없으며, 함수 변환으로 미분해. 서로 다른 비용을 치르는 만큼 모양도 달라.
  • 성능 — 같은 모델에서 둘의 차이는 작거나 중간 정도고, 때로는 한쪽이 앞섰다가 다음 버전에는 뒤집혀. 구체적인 작업을 재 보지 않았다면 성능만으로 고르지 마.
  • 연산 지원 범위 — PyTorch MPS는 아직 모든 PyTorch 연산을 Metal에서 지원하지 않아. 버전마다 빈틈이 줄지만, CPU 폴백을 명시적으로 켜면 일부 연산이 CPU로 돌아가 성능을 크게 떨어뜨릴 수 있어. MLX는 지원 범위가 더 작지만 그 안에서는 일관되게 작동해.
  • 메모리 방식 — PyTorch MPS API에는 Apple Silicon에서 필요하지 않은 .to(device) 절차가 남아 있어. MLX API는 실제 하드웨어 방식과 맞아.

PyTorch MPS가 맞는 때

  • 기존 PyTorch 코드베이스가 있고 다시 쓰지 않은 채 Mac에서 개발하고 싶을 때. CUDA를 MPS로 바꾸는 최소 변경이 바로 PyTorch MPS의 존재 이유야.
  • PyTorch 전용 라이브러리에 의존할 때 — Hugging Face Transformers나 특정 torch 전용 모델이 아직 MLX로 모두 옮겨지지는 않았어.
  • 클라우드 GPU 학습과 세세한 동작까지 맞춰야 할 때 — Mac에서 시제품을 만든 뒤 예상 밖의 차이 없이 NVIDIA 환경에 배포하기 좋아.

MLX가 맞는 때

  • Apple Silicon에서 새로 시작하며 지켜야 할 PyTorch 유산이 없을 때.
  • 함수 변환 방식을 원할 때 — JAX와 닮은 mx.grad, mx.vmap, 간단한 컴파일을 쓸 수 있어.
  • GPU를 빌리지 않고 LLM을 로컬에서 파인튜닝할 때 — mlx-lm의 LoRA 작업 흐름은 PyTorch MPS 쪽보다 더 다듬어져 있어.
  • Mac 전용으로 배포하며 하드웨어에 자연스러운 API를 원할 때.

솔직한 중간 지점

외부 GPU 클러스터와 Mac을 오가는 연구라면 PyTorch MPS가 코드 하나를 유지하게 해 줘. Mac 전용 작업, 특히 로컬 LLM 작업이라면 다른 회사의 하드웨어를 먼저 생각하지 않고 만든 MLX가 더 자연스러워.

Code

PyTorch MPS와 MLX에서 같은 행렬 곱셈·python
# PyTorch MPS
import torch
device = torch.device("mps" if torch.backends.mps.is_available() else "cpu")
x = torch.randn(1024, 1024, device=device)
y = x @ x.T
print("PyTorch MPS:", tuple(y.shape), y.dtype, y.device)

# MLX
import mlx.core as mx
x_mlx = mx.random.normal((1024, 1024))
y_mlx = x_mlx @ x_mlx.T
mx.eval(y_mlx)
print("MLX        :", tuple(y_mlx.shape), y_mlx.dtype, mx.default_device())

# Verified outputs (2026-05-03):
#   PyTorch MPS: (1024, 1024) torch.float32 mps:0
#   MLX        : (1024, 1024) mlx.core.float32 Device(gpu, 0)

External links

Exercise

환경에 torch가 설치돼 있다면 PyTorch MPS와 MLX에서 같은 행렬 곱셈을 실행해. 이 레슨의 코드 블록이 둘 다 처리해. time.perf_counter()로 시간을 잰 다음 4096×4096 행렬로 바꿔 다시 재 봐. 차이가 커졌는지, 줄었는지, 비슷한지 두 문장으로 적어.

Progress

Progress is local-only — sign in to sync across devices.
이 페이지에서 버그를 발견하셨거나 피드백이 있으세요?문제 신고

댓글 0

🔔 답글 알림 (로그인 필요)
로그인댓글을 남기려면 로그인해 주세요.

아직 댓글이 없어요. 첫 댓글을 남겨보세요.