google-deepmind/dm-haiku
JAX-based neural network library
What it solves
JAX는 강력한 수치 계산 라이브러리이지만, 변환(jit 및 grad 등)을 사용하려면 함수가 순수해야 합니다. 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과 완전히 호환되어 여러 가속기에서 분산 학습이 가능합니다.