google/flax

Flax is a neural network library for JAX that is designed for flexibility.

Flax란 무엇인가요?

Flax는 JAX (고성능 수치 계산 프레임워크) 위에 구축된 오픈소스 신경망 라이브러리입니다. 연구자들이 일반적인 파이썬 객체로 모델을 작성할 수 있도록 하는 유연하고 파이썬 친화적인 API(최신 Flax NNX API)를 제공합니다. 참조 공유 및 가변성도 지원합니다. 라이브러리는 일반적인 레이어(Linear, Conv, BatchNorm, Attention, LSTM/GRU, Dropout 등), 복제된 훈련, 체크포인트, 메트릭을 위한 유틸리티, MNIST 및 Gemma 언어 모델 데모와 같은 교육용 예제를 포함합니다.

누가 유지보수하고 있나요?

Google DeepMind의 엔지니어와 연구자들이 JAX 팀과 긴밀히 협력하여 개발했습니다. 공식 Google 제품은 아니지만, GitHub 디스커션과 풀 리퀘스트를 통해 커뮤니티 기여를 적극 수용하며 지속적으로 유지보수되고 있습니다.

어떻게 시작하나요?

  1. JAX 설치 (JAX CPU/GPU/TPU 가이드 따르기).
  2. PyPI를 통해 Flax 설치:
    pip install flax
    
    옵션: pip install "flax[all]"로 Matplotlib 등 추가 종속성 설치 가능.
  3. nnx.Module를 상속하여 모델을 작성하고 제공된 레이어를 사용한 후, 표준 JAX 코드로 훈련합니다.

예제 코드 (README에서)

class MLP(nnx.Module):
  def __init__(self, din, dmid, dout, *, rngs):
    self.linear1 = nnx.Linear(din, dmid, rngs=rngs)
    self.dropout = nnx.Dropout(rate=0.1, rngs=rngs)
    self.bn = nnx.BatchNorm(dmid, rngs=rngs)
    self.linear2 = nnx.Linear(dmid, dout, rngs=rngs)

  def __call__(self, x):
    x = nnx.gelu(self.dropout(self.bn(self.linear1(x))))
    return self.linear2(x)

더 배우고 싶다면?

  • 문서 사이트: https://flax.readthedocs.io/
  • 튜토리얼: MNIST 튜토리얼, Gemma LM 추론 예제, "Flax NNX 기초" 가이드.
  • 디스커션 및 지원: GitHub Discussions, 이슈 트래커, flax-dev@google.com 메일 주소.

언제 Flax를 사용해야 하나요?

JAX를 이미 사용 중이고 다음 조건을 충족하는 신경망 라이브러리가 필요하다면, Flax는 자연스러운 선택입니다:

  • JAX의 완전한 유연성 유지 (숨겨진 그래프 컴파일 단계 없음).
  • 일반적인 파이썬 문법으로 모델 작성 가능.
  • 미리 준비된 레이어, 훈련 유틸리티, 예제 코드 제공.

Citation: README에는 학술적 인용을 위한 BibTeX 항목이 제공됩니다.

관련

  • 프로젝트
  • 프로젝트
  • Dispatch
  • 프로젝트
  • 프로젝트