Hugging Face에서 mlx-image 사용하기
Hugging Face에서 mlx-image 사용하기
mlx-image는 Riccardo Musmeci가 개발한, Apple MLX 기반의 이미지 모델 라이브러리예요. 훌륭한 timm을 MLX 모델용으로 재현하려는 시도예요.
출처: 문서
본문
Hub에서 mlx-image 탐색하기
이 쿼리처럼 mlx-image 라이브러리 이름으로 필터링하면 mlx-image 모델을 찾을 수 있어요. MLX 형식으로 가중치를 변환·게시하는 기여자들의 공개 mlx-vision 커뮤니티도 있어요.
설치
pip install mlx-image
모델
모델 가중치는 HuggingFace의 mlx-vision 커뮤니티에서 제공돼요.
사전학습 가중치로 모델을 로드하려면:
from mlxim.model import create_model
# loading weights from HuggingFace (https://huggingface.co/mlx-vision/resnet18-mlxim)
model = create_model("resnet18") # pretrained weights loaded from HF
# loading weights from local file
model = create_model("resnet18", weights="path/to/resnet18/model.safetensors")
사용 가능한 모든 모델을 나열하려면:
from mlxim.model import list_models
list_models()
ImageNet-1K 결과
results-imagenet-1k.csv에서 mlx-image로 변환된 모든 모델과 서로 다른 설정에서의 ImageNet-1K 성능을 확인할 수 있어요.
TL;DR 성능은 PyTorch 구현의 원본 모델과 비교했을 때 동등한 수준이에요.
PyTorch 및 다른 익숙한 도구와의 유사성
mlx-image는 PyTorch에 최대한 가깝게 동작하려고 해요.
-
DataLoader-> 직접collate_fn을 정의하고 데이터 로딩을 빠르게 하는num_workers도 사용할 수 있어요 -
Dataset->mlx-image는 이미LabelFolderDataset(좋은 옛 PyTorchImageFolder)와FolderDataset(이미지가 담긴 일반 폴더)을 지원해요 -
ModelCheckpoint-> 최상의 모델을 추적해 디스크에 저장해요(PyTorchLightning과 유사). 조기 종료(early stopping)도 제안해 줘요
학습
학습은 PyTorch와 유사해요. 모델을 학습하는 예시는 다음과 같아요.
import mlx.nn as nn
import mlx.optimizers as optim
from mlxim.model import create_model
from mlxim.data import LabelFolderDataset, DataLoader
train_dataset = LabelFolderDataset(
root_dir="path/to/train",
class_map={0: "class_0", 1: "class_1", 2: ["class_2", "class_3"]}
)
train_loader = DataLoader(
dataset=train_dataset,
batch_size=32,
shuffle=True,
num_workers=4
)
model = create_model("resnet18") # pretrained weights loaded from HF
optimizer = optim.Adam(learning_rate=1e-3)
def train_step(model, inputs, targets):
logits = model(inputs)
loss = mx.mean(nn.losses.cross_entropy(logits, target))
return loss
model.train()
for epoch in range(10):
for batch in train_loader:
x, target = batch
train_step_fn = nn.value_and_grad(model, train_step)
loss, grads = train_step_fn(x, target)
optimizer.update(model, grads)
mx.eval(model.state, optimizer.state)
추가 리소스
연락처
질문이 있으면 [email protected]으로 이메일 보내주세요.
더 알아보기 (Learn more)
mlx-image는 PyTorch에 가까운 API를 갖춘 MLX용 이미지 모델 라이브러리예요. create_model("resnet18")처럼 Hub의 가중치를 직접 로드할 수 있고, 학습 루프도 PyTorch 스타일로 작성해요. 성능은 원본 PyTorch 구현과 동등하다고 알려져 있어요.