Downloads · 30 days
0
flax-community/NeuralODE_SDE
NeuralODE_SDE is a machine learning model from flax-community. Use it for the machine learning task on the model card, and read the license before you ship it in a product.
This is the result of project ["Reproduce Neural ODE and SDE"][projectlink] in [HuggingFace Flax/JAX community week][comweeklink].
Downloads · 30 days
0
Access
Public
Updated Jul 19, 2021
Repo size
—
Likes
0
Public
Click a slice to open those files.
.jpg27.1 MB · 53%
From the Hugging Face model README
This is the result of project "Reproduce Neural ODE and SDE" in HuggingFace Flax/JAX community week.
<code>main.py</code> will execute training of ResNet or OdeNet for MNIST dataset.
For JAX installation, please follow here.
or simply, type
pip install jax jaxlib
For Flax installation,
pip install flax
Tensorflow-datasets will download MNIST dataset to environment.
For (small) ResNet training,
python main.py --model=resnet --lr=1e-4 --n_epoch=20 --batch_size=64
For Neural ODE training,
python main.py --model=odenet --lr=1e-4 --n_epoch=20 --batch_size=64
For Continuous Normalizing Flow,
python main.py --model=cnf --sample_dataset=circles
Sample datasets can be chosen as circles, moons, or scurve.

These are the codes for the bird call generation score sde model.
<code>core-sde-sampler.py</code> will execute the sampler. The sampler uses pretrained weight to generate bird calls. The weight can be found here
For using different sample generation parameters change the argument values. For example,
python main.py --sigma=25 --num_steps=500 --signal_to_noise_ratio=0.10 --etol=1e-5 --sample_batch_size = 128 --sample_no = 47
In order to generate the audios, these dependencies are required,
pip install librosa
pip install soundfile
In order to train the model from scratch, please generate the dataset using this link. The dataset is generated in kaggle. Therefore, during training your username and api key is required in the specified section inside the script.
python main.py --sigma=35 --n_epochs=1000 --batch_size=512 --lr=1e-3 --num_steps=500 --signal_to_noise_ratio=0.15 --etol=1e-5 --sample_batch_size = 64 --sample_no = 23