MLX 빠른 시작 — 첫 배열과 간단한 모델

MLX 빠른 시작 — 첫 배열과 간단한 모델

MLX를 처음 쓰면 NumPy와 거의 같은 감각으로 배열을 만들 수 있어요. 핵심 모듈은 mlx.core(보통 mx로 import)이고, 신경망은 mlx.nn, 최적화는 mlx.optimizers에서 가져옵니다.

출처: https://ml-explore.github.io/mlx/usage/quick_start.html

배열 만들기

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합니다. 그래서 코드를 쓰는 순간보다 실제 평가가 일어날 때 연산이 실행되는 점을 기억해두면 디버깅이 편해져요.

더 알아보기