Flax 실전 — MNIST CNN 학습

Flax 실전 — MNIST CNN 학습

Flax NNX를 실전에서 익히려면 가장 전형적인 MNIST 필기 숫자 분류 CNN을 만들고 학습하는 튜토리얼을 따라가면 돼요. 데이터 로드부터 모델 정의, 학습 스텝, 추론까지 한 흐름으로 볼 수 있어요.

출처: https://flax.readthedocs.io/en/latest/mnist_tutorial.html

설치

# !pip install -U flax

Hugging Face datasets로 MNIST를 불러와 이미지를 0~1로 정규화하고 배치로 만드는 헬퍼를 준비해요.

import numpy as np
from datasets import load_dataset

train_steps = 1200
eval_every = 200
batch_size = 32
dataset = load_dataset('mnist')
train_ds = dataset['train'].shuffle(seed=0)
test_ds = dataset['test']

CNN 모델 정의

nnx.Module을 상속해 CNN을 정의해요. nnx.Conv, nnx.BatchNorm, nnx.Dropout, nnx.Linear 같은 레이어를 조합해요.

from flax import nnx
from functools import partial

class CNN(nnx.Module):
    # A simple CNN model.
    def __init__(self, *, rngs: nnx.Rngs):
        self.conv1 = nnx.Conv(1, 32, kernel_size=(3, 3), rngs=rngs)
        self.batch_norm1 = nnx.BatchNorm(32, rngs=rngs)
        self.dropout1 = nnx.Dropout(rate=0.025)
        self.conv2 = nnx.Conv(32, 64, kernel_size=(3, 3), rngs=rngs)
        self.batch_norm2 = nnx.BatchNorm(64, rngs=rngs)
        self.avg_pool = partial(nnx.avg_pool, window_shape=(2, 2), strides=(2, 2))
        self.linear1 = nnx.Linear(3136, 256, rngs=rngs)
        ...

# Instantiate the model.
model = CNN(rngs=nnx.Rngs(0))

각 레이어에 rngs를 넘겨 파라미터를 초기화하고, nnx.Optimizer + nnx.value_and_grad로 학습 스텝을 구성해요. 테스트 셋 추론 함수를 정의해 학습 후 정확도를 확인해요.

결과

이렇게 해서 Flax NNX로 CNN을 MNIST에 end-to-end로 학습하고 분류하는 방법을 배울 수 있어요. 저장소의 mnist_tutorial.ipynb를 Colab에서 그대로 실행해 볼 수도 있어요.

더 알아보기