JAX 코어에는 nn.Linear, nn.Conv2d 같은 게 없어. 처음 보면 의아하지만, 의도된 설계야.
JAX의 핵심, 수치 기본 요소 (jax.numpy) + 변환 (jit, grad, vmap, pmap). 거기까지가 core. 신경망 추상화는 따로 골라 써.
왜 이런 설계?
1. 신경망 추상화는 "어떤 패러다임이 옳은가"에 답이 갈림
- 변경 가능한 상태 (PyTorch 식):
self.weight = ...같은 인스턴스 속성. 직관적이야. - 순수 함수형 (Haiku 식):
params와apply_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다운 방식 선호 |
| Haiku | transform 기반 (옛 Flax) | DeepMind의 옛 코드, AlphaFold |
| Penzai | 전체 모델 visualization | research 디버깅 |
| 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가 다 이해된 사람에게 신경망 라이브러리 선택은 비교적 작은 결정이 돼.