google-deepmind/dm-haiku

JAX-based neural network library

해결하는 문제

Haiku는 개발자가 (Sonnet 또는 TensorFlow에서 볼 수 있는 것과 같은) 익숙한 객체 지향 프로그래밍 모델을 사용하면서도 JAX의 순수 함수 변환에 대한 완전한 액세스를 유지할 수 있도록 해주는 JAX용 신경망 라이브러리입니다. 초기화를 위해 방대한 보일러플레이트 코드를 작성할 필요 없이 모델 파라미터와 내부 상태 관리를 단순화합니다.

작동 방식

Haiku는 객체 지향 설계와 함수적 순수성 사이의 간극을 메우기 위해 두 가지 주요 도구를 제공합니다:

  • hk.Module: 신경망 레이어 및 컴포넌트를 정의하는 데 사용되는 Python 객체입니다. 이러한 모듈은 파라미터 및 메서드에 대한 참조를 보유하지만, 정의 중에는 함수적으로 "불순(impure)"하게 취급됩니다.
  • hk.transform: 이러한 불순한 모듈 기반 함수를 init(초기 파라미터 값을 수집) 및 apply(계산을 위해 해당 파라미터를 함수에 주입)라는 한 쌍의 순수 함수로 변환하는 함수 변환입니다.

가변 상태(Batch Normalization 이동 평균과 같은)가 필요한 모델의 경우, Haiku는 파라미터와 상태를 별도로 관리하는 hk.transform_with_state를 제공합니다.

대상

신경망 구축을 위해 객체 지향 API의 생산성을 원하면서도 JAX의 성능 및 변환 기능(jax.jit, jax.grad, jax.pmap 등)이 필요한 연구자 및 개발자.

주요 특징

  • JAX 통합: JAX의 자동 미분 및 GPU/TPU 지원과 완전히 호환됩니다.
  • Sonnet 스타일 API: Sonnet 2 API와 거의 일치하도록 설계되어 TensorFlow/Sonnet로부터의 전환이 쉽습니다.
  • 단순화된 RNG 관리: 변환된 함수 내에서 결정론적으로 난수 생성을 처리하기 위해 hk.next_rng_key()를 제공합니다.
  • 확장성: DeepMind 연구자들에 의해 이미지, 언어 및 강화 학습 작업에 대해 대규모로 테스트되었습니다.
  • 프레임워크가 아닌 라이브러리: 파라미터 및 상태 관리에만 엄격하게 집중하며, 옵티마이저 및 체크포인팅은 다른 라이브러리에 맡깁니다.

관련

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