본문 바로가기
C.W.K.
Stream
Lesson 01 of 06 · published

JAX가 신경망 라이브러리를 안 갖는 이유

~8 min · neural-nets, jax, tutorial

Level 0호기심
0 XP0/73 lessons0/17 achievements
0/100 XP to next level100 XP to go0% complete

JAX 코어에는 nn.Linear, nn.Conv2d 같은 게 없어. 처음 보면 의아하지만, 의도된 설계야.

JAX의 핵심, 수치 기본 요소 (jax.numpy) + 변환 (jit, grad, vmap, pmap). 거기까지가 core. 신경망 추상화는 따로 골라 써.

왜 이런 설계?

1. 신경망 추상화는 "어떤 패러다임이 옳은가"에 답이 갈림

  • 변경 가능한 상태 (PyTorch 식): self.weight = ... 같은 인스턴스 속성. 직관적이야.
  • 순수 함수형 (Haiku 식): paramsapply_fn 분리해. JAX다운 방식.
  • pytree 모듈 (Equinox 식): 모델 자체가 pytree, jit/grad가 직접 변환해.
  • NNX (Flax의 새 API): PyTorch식 변경 가능한 모델과 JAX 변환을 잘 결합해.

한 표준을 강요하면 다른 패러다임이 막혀. JAX는 의도적으로 비워 둬.

2. 다양한 분야가 다른 추상을 원함

  • RL, 상태 머신 추상화 강조
  • 과학, ODE solver, simulator와 자연스러운 연동
  • 비전, CNN, transformer 표준
  • NLP, transformer, attention 표준

각 분야가 자기 라이브러리를 만들 자유.

3. 유지보수 부담 분리

신경망 API는 변화가 빠름 (transformer 변형, attention 종류, normalization 변형). core에 두면 JAX의 안정성과 신경망의 진화 사이에 충돌이 생겨. 분리하면 JAX 코어는 천천히 안정되고, 신경망 라이브러리는 빠르게 진화해.

현재 주요 라이브러리

이름스타일주력 사용처
Flax NNX변경 가능한 Python 상태Google, DeepMind의 새 표준
Equinox모델 = pytree, pure 함수형학술 연구, JAX다운 방식 선호
Haikutransform 기반 (옛 Flax)DeepMind의 옛 코드, AlphaFold
Penzai전체 모델 visualizationresearch 디버깅
Levanter학습 scale 특화대규모 학습

🎯 어느 걸 골라야 하나

(1) 새 프로젝트 + Google/DeepMind 영향권, Flax NNX. 가장 적극적 발전. (2) 학술 연구 / 함수형 선호, Equinox. JAX의 정신과 가장 일치. (3) AlphaFold 같은 구식 코드 봐야 함, Haiku. 그러나 새 코드는 안 추천해. (4) 처음 배우면 Flax NNX가 PyTorch와 가장 비슷한 사용감이라 진입 장벽 낮아.

이 퀘스트는 Flax NNX와 Equinox 둘 다 다룸 (10-2, 10-3). 같은 모델을 두 라이브러리로 작성해서 차이를 직접 봐.

중요한 점은 어느 라이브러리를 선택하든 JAX의 핵심인 jit, grad, vmap, pytree는 그대로라는 거야. 신경망 라이브러리는 그 위에 얹는 편의 문법이야. 이 퀘스트의 Track 1~9가 다 이해된 사람에게 신경망 라이브러리 선택은 비교적 작은 결정이 돼.

Code

# PyTorch: one way to define a model
# class Model(nn.Module):
#     def __init__(self):
#         super().__init__()
#         self.linear = nn.Linear(784, 10)
#     def forward(self, x):
#         return self.linear(x)

# JAX: you choose your library
# Flax NNX version:
from flax import nnx
class Model(nnx.Module):
    def __init__(self, rngs):
        self.linear = nnx.Linear(784, 10, rngs=rngs)
    def __call__(self, x):
        return self.linear(x)

# Equinox version:
import equinox as eqx
class Model(eqx.Module):
    linear: eqx.nn.Linear
    def __init__(self, key):
        self.linear = eqx.nn.Linear(784, 10, key=key)
    def __call__(self, x):
        return self.linear(x)

External links

Exercise

공식 자료 세 곳에서 JAX 코어에 신경망 라이브러리가 없는 이유에 관한 서로 다른 설명을 읽어. 'Flax나 Equinox를 왜 따로 배워야 하지?'라고 묻는 동료에게 100단어 안팎으로 답을 써. 이 관점이 라이브러리 선택에 어떤 영향을 주는지도 덧붙여.

Progress

Progress is local-only — sign in to sync across devices.
이 페이지에서 버그를 발견하셨거나 피드백이 있으세요?문제 신고

댓글 0

🔔 답글 알림 (로그인 필요)
로그인댓글을 남기려면 로그인해 주세요.

아직 댓글이 없어요. 첫 댓글을 남겨보세요.