JAX-PRNN

Published:

JAX-PRNN

Physically Recurrent Neural Networks using JAX

This is a JAX-based implementation of PRNNs. This version accelerates training and inference time by >10x. For more information and the code see the GitHub repository: https://github.com/JoepStorm/jax-prnn

Features:

  • Just-In-Time compilation
  • Scan function instead of for loop
  • Encoder and Decoder run only once per sequence, instead of once per time step
  • Einsum notation
  • Jupyter notebook that demonstrates a simple PRNN training