MLX 퀵스타트 가이드

MLX 퀵스타트 가이드

MLX를 처음 쓸 때 알아야 할 핵심을 다뤄요. 배열 만들기, 지연 계산, 함수·그래프 변환까지 한 번에 살펴볼게요. NumPy에 익숙하다면 대부분의 코드가 자연스럽게 읽혀요.

출처: MLX Quick Start Guide (공식)

기본: 배열 만들기

mlx.core를 import하고 array를 만들어요.

>>> import mlx.core as mx
>>> a = mx.array([1, 2, 3, 4])
>>> a.shape
(4,)
>>> a.dtype
int32
>>> b = mx.array([1.0, 2.0, 3.0, 4.0])
>>> b.dtype
float32

지연 계산

MLX의 연산은 지연돼요. 연산의 출력은 필요할 때까지 계산되지 않아요. 배열을 강제로 평가하려면 eval()을 사용해요. 스칼라를 array.item()으로 확인하거나, 배열을 print하거나, numpy.ndarray로 변환할 때는 자동으로 평가돼요.

>>> c = a + b   # c는 아직 평가 안 됨
>>> mx.eval(c)  # c 평가
>>> c = a + b
>>> print(c)    # 이 역시 c를 평가함
array([2, 4, 6, 8], dtype=float32)
>>> c = a + b
>>> import numpy as np
>>> np.array(c)   # 이것도 c를 평가함
array([2., 4., 6., 8.], dtype=float32)

함수와 그래프 변환

MLX는 grad()vmap() 같은 표준 함수 변환을 제공해요. 변환은 임의로 조합될 수 있어서 grad(vmap(grad(fn))) 같은 형태도 허용돼요.

>>> x = mx.array(0.0)
>>> mx.sin(x)
array(0, dtype=float32)
>>> mx.grad(mx.sin)(x)
array(1, dtype=float32)
>>> mx.grad(mx.grad(mx.sin))(x)
array(-0, dtype=float32)

다른 그래디언트 변환으로는 벡터-야코비안 곱을 위한 vjp(), 야코비안-벡터 곱을 위한 jvp()가 있어요.

함수의 출력과 입력에 대한 그래디언트를 함께 효율적으로 계산하려면 value_and_grad()를 사용해요.

더 알아보기