Flax NNX 기초 — 모듈과 변환

Flax NNX 기초 — 모듈과 변환

Flax NNX는 네트워크를 만들고, 검사하고, 디버깅하고, 분석하기 쉽게 만든 새 단순화 API예요. **파이썬 레퍼런스 의미론(refresh semantics)**을 일급 지원해서, PyTorch나 Keras 사용자에게 익숙한 방식으로 모델을 표현할 수 있어요.

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

Module 시스템

NNX에서는 대부분의 것이 **명시적(explicit)**이에요. 모듈이 상태(파라미터 등)를 직접 들고, PRNG 상태는 사용자가 전달하며, 모든 모양 정보는 초기화 시 지정해요(모양 추론 없음). 동적 상태는 보통 Param에, 정적 상태(정수·문자열)는 그대로 저장해요.

from flax import nnx
import jax
import jax.numpy as jnp

class Linear(nnx.Module):
    def __init__(self, din: int, dout: int, *, rngs: nnx.Rngs):
        self.w = nnx.Param(rngs.params.uniform((din, dout)))
        self.b = nnx.Param(jnp.zeros((dout,)))
        self.din, self.dout = din, dout
    def __call__(self, x: jax.Array):
        return x @ self.w + self.b[None]

model = Linear(2, 5, rngs=nnx.Rngs(params=0))
y = model(x=jnp.ones((1, 2)))
nnx.display(model)

Variable의 내부 값은 .value로 접근하지만, 편의상 산술 연산에서 직접 쓸 수도 있어요.

상태 갱신 (예: Counter)

정방향 중 상태를 갱신해야 하는 레이어라면 Variable을 만들고 forward에서 .value를 바꾸면 돼요.

class Count(nnx.Variable): pass
class Counter(nnx.Module):
    def __init__(self):
        self.count = Count(jnp.array(0))
    def __call__(self):
        self.count[...] += 1

학습 스텝과 변환

nnx.value_and_grad로 손실·그래디언트를 구하고 옵티마이저로 갱신해요. nnx.jit로 자동 상태 전파를 켜면 BatchNorm·Dropout 상태 갱신이 안쪽에서 바깥쪽 모델 참조까지 전파돼요.

model = MLP(2, 16, 10, rngs=nnx.Rngs(0))
optimizer = nnx.Optimizer(model, optax.adam(1e-3), wrt=nnx.Param)

@nnx.jit
def train_step(model, optimizer, x, y, rngs):
    def loss_fn(model: MLP, rngs: nnx.Rngs):
        y_pred = model(x, rngs)
        return jnp.mean((y_pred - y) ** 2)
    loss, grads = nnx.value_and_grad(loss_fn)(model, rngs)
    optimizer.update(model, grads)
    return loss

함수형 API

JAX 변환 경계를 넘을 때는 nnx.split, nnx.merge, nnx.update로 상태를 명시적으로 다뤄요. nnx.split은 모듈을 GraphDef + State(pytree)로 분리해서 어떤 JAX 변환이든 승격시킬 수 있어요.

graphdef, state = nnx.split(model)
model = nnx.merge(graphdef, state)
nnx.update(model, state)

더 알아보기