google-deepmind/kfac-jax
Second Order Optimization and Curvature Estimation with K-FAC in JAX.
What it solves
KFAC-JAX는 사용하기 쉬운 K-FAC (Kronecker-factored Approximate Curvature) 옵티마이저와 곡률 추정기를 구현합니다. 2차 최적화는 일반적으로 1차 방법보다 계산 비용이 훨씬 더 많이 들지만, KFAC-JAX는 확장 가능한 곡률 행렬 근사치를 제공하여 이 문제를 해결합니다.
How it works
JAX 기반으로 구축된 이 라이브러리는 연구자들이 손실 함수의 곡률을 근사하여 신경망을 최적화할 수 있도록 합니다. 사용자는 라이브러리가 곡률 행렬을 정확하게 근사할 수 있도록 특정 손실 함수를 등록해야 합니다 (예: register_factored-cross_entropy_loss - Note: The user'نماది기 위해 register_cross_entropy_loss를 등록해야 합니다.). 옵티마이저 내부적으로 스테이징 (JIT and PMAP)을 자동으로 처리하므로, 사용자는 옵티마이저의 step function이나 손실 함수에 jax.jit 또는 jax.pmap을 수동적으로 적용하지 않아야 합니다.
Who it’s for
JAX를 사용하고 있으며, 신경망 학습을 개선하기 위해 2차 최적화 기술을 implement したい
Highlights
JAX-native: 가속기에서의 고성능을 위해 라이브러리 전용으로 구축되었습니다.
Automatic Staging: 옵티마이저 step에 대해
jitandpmap을 자동으로 처리합니다.Adaptive Parameters: 적응형 학습률, 적응형 모멘텀, 적응형 댐핑을 지원합니다.
Curvature Estimation: 신경망 학습을 위한 확장 가능한 곡률 근사치를 제공합니다。
관련
- 프로젝트
- 프로젝트
- 프로젝트
- 프로젝트
- 프로젝트