google-deepmind/dm-haiku

JAX-based neural network library

What it solves

JAX는 강력한 수치 계산 라이브러리이지만, 변환(jitgrad 등)을 사용하려면 함수가 순수해야 합니다. Haiku는 개발자가 익숙한 객체 지향 프로그래밍 모델을 사용해 신경망을 정의하도록 허용하고, 이러한 "불순" 객체 지향 정의를 JAX가 처리할 수 있는 순수 함수로 자동 변환함으로써 이 문제를 해결합니다.

How it works

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

  • hk.Module: 네트워크 레이어와 컴포넌트를 정의하는 Python 객체. 이 모듈들은 파라미터와 메서드에 대한 참조를 보유하여, 사용자가 표준 신경망 라이브러리와 같은 코드를 작성할 수 있게 합니다.
  • hk.transform: hk.Module을 사용하는 함수를 순수 함수 쌍인 init(초기 파라미터 값을 수집)와 apply(그 파라미터를 함수에 주입하여 계산)로 변환하는 함수 변환기.

배치 정규화와 같이 내부 가변 상태가 필요한 모델의 경우, Haiku는 hk.transform_with_state를 제공하여 파라미터와 상태를 별도로 관리합니다.

Who it’s for

객체 지향 API의 생산성을 활용하면서도 JAX의 함수 변환 및 하드웨어 가속의 전체 기능을 필요로 하는 연구자와 개발자를 위한 것입니다.

Highlights

  • DeepMind Scale: DeepMind 연구원들이 대규모 이미지, 언어, 강화 학습 실험에서 테스트했습니다.
  • Library, Not Framework: 경량 라이브러리로 설계되어 파라미터와 상태 관리에 집중하고, 커스텀 옵티마이저나 체크포인트 포맷을 강제하지 않습니다.
  • hk.next_rng_key(): 결정적인 키 시퀀스를 제공하여 JAX 내에서 난수 생성을 단순화합니다.
  • JAX Compatibility: jax.pmap과 완전히 호환되어 여러 가속기에서 분산 학습이 가능합니다.