JAX-PRNN
Published:
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