MLX 빠른 시작 — 첫 배열과 간단한 모델
MLX 빠른 시작 — 첫 배열과 간단한 모델
MLX를 처음 쓰면 NumPy와 거의 같은 감각으로 배열을 만들 수 있어요. 핵심 모듈은 mlx.core(보통 mx로 import)이고, 신경망은 mlx.nn, 최적화는 mlx.optimizers에서 가져옵니다.
배열 만들기
import mlx.core as mx
x = mx.array([1, 2, 3, 4])
print(x) # array([1, 2, 3, 4], dtype=mlx.int64)
간단한 모델과 학습 루프
mlx.nn.Linear 같은 레이어를 쌓고, mlx.optimizers.SGD 같은 옵티마이저로 학습할 수 있어요. PyTorch에 익숙하다면 구조가 아주 비슷하다는 걸 느낄 수 있죠.
import mlx.core as mx
import mlx.nn as nn
import mlx.optimizers as optim
model = nn.Linear(2, 1)
optimizer = optim.SGD(learning_rate=0.01)
def loss_fn(model, X, y):
return mx.mean((model(X) - y) ** 2)
def step(model, X, y, optimizer):
loss_and_grad_fn = nn.value_and_grad(model, loss_fn)
loss, grads = loss_and_grad_fn(model, X, y)
optimizer.update(model, grads)
return loss
지연 계산 이해하기
MLX는 lazy(지연) 계산을 써요. 배열 연산을 바로 실행하지 않고 필요한 시점에 한꺼번에 materialize합니다. 그래서 코드를 쓰는 순간보다 실제 평가가 일어날 때 연산이 실행되는 점을 기억해두면 디버깅이 편해져요.