Flax — JAX를 위한 신경망 라이브러리

Flax

JAX로 신경망을 짜는데, 함수형뿐 아니라 익숙한 객체지향 방식으로도 모델을 만들고 싶다면 Flax를 써요. Flax는 "Neural Networks for JAX"의 약자로, JAX를 쓰는 연구자·개발자에게 유연한 end-to-end 경험을 제공해요. 최신 NNX API에서는 파라미터 같은 상태를 모듈이 직접 들고 있어서 PyTorch나 Keras 사용자도 편하게 적응할 수 있어요. MNIST 같은 전형적인 분류 모델부터 대규모 실험까지 폭넓게 쓰여요.

이 카테고리의 문서

  • 소개: JAX용 신경망 라이브러리
  • NNX 기초: 모듈과 변환
  • 실전: MNIST CNN 학습

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