Efficient implementation of tensorized kernel methods for different hardware.
pip install --upgrade pip
pip install wheel
pip install --upgrade "jax[cpu]"
https://github.com/google/jax/blob/main/README.md#pip-installation-gpu-cuda
pip install "jax[cuda11_cudnn805]" -f https://storage.googleapis.com/jax-releases/jax_cuda_releases.html --> !!! WARNING: jax 0.3.14 does not provide the extra 'cuda11_cudnn811' !!!
Jax on apple silicon: https://developer.apple.com/metal/jax/
pip install -e ".[dev]"
https://jax.readthedocs.io/en/latest/profiling.html?highlight=gpu#gpu-profiling
Build: docker build -t jtsch/tkm .
Build and Push: docker build -t jtsch/tkm . && docker push jtsch/tkm
Pull: docker pull jtsch/tkm
Run with terminal: docker run -it jtsch/tkm