pyrfd
Pytorch implementation of RFD (see arXiv)
Covariance model
Provides an implementation of the SquaredExponential covariance model
with an auto_fit function, which requires only
- A
model_factorywhich returns the same but randomly initialized model every time it is called - A
lossfunction e.g.torch.nn.functional.nll_losswhich accepts a prediction and a true value - data, which can be passed to
torch.utils.DataLoaderwith different batch size parameters such that it returns(x,y)tuples when iterated on - a
csvfilename which acts as the cache for the covariance model ofthis unique (model, data, loss) combination.
Implementation of RFD
Such a covariance model can then be passed to RFD which implements the
pytorch optimizer interface. The end result can be used like torch.optim.Adam
Example usage
from benchmaking.classification.mnist.models.cnn3 import CNN3
import torch
import torchvision as tv
from pyrfd import RFD, SquaredExponential
cov_model = SquaredExponential()
cov_model.auto_fit(
model_factory=CNN3,
loss=torch.nn.functional.nll_loss,
data= tv.datasets.MNIST(
root="mnistSimpleCNN/data",
train=True,
transform=tv.transforms.ToTensor()
),
cache="cache/CNN3_mnist.csv",
# should be unique for (models, data, loss)
)
rfd = RFD(
CNN3().parameters(),
covariance_model=cov_model
)
How to cite
@inproceedings{benningRandomFunctionDescent2024,
title = {Random {{Function Descent}}},
booktitle = {Advances in {{Neural Information Processing Systems}}},
author = {Benning, Felix and D{\"o}ring, Leif},
year = {2024},
month = dec,
volume = {37},
primaryclass = {cs, math, stat},
publisher = {Curran Associates, Inc.},
address = {Vancouver, Canada},
}
Metadata
Release files for pyrfd 1.0.1
For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.
Source distribution (sdist)
| File | Size | Uploaded | |
|---|---|---|---|
| pyrfd-1.0.1.tar.gz | 15.3 kB | Details |
Built distribution (wheel)
| File | Interpreter | ABI | Platform | Reset |
|---|---|---|---|---|
| pyrfd-1.0.1-py3-none-any.whl | Python 3 | none | any | Details |
Total release size: 32.3 kB
Release files / pyrfd-1.0.1.tar.gz
| Download URL | pyrfd-1.0.1.tar.gz |
|---|---|
| Size | 15.3 kB |
| Tags | Source |
|
SHA-256 checksum How to use checksums |
69b9c05918ab524dec76fb9050e4146ee12751115b82d03d2148f0a35b6ed1f3
|
|
BLAKE2b-256 checksum How to use checksums |
98511e2bb8cd10c5316f99980c6d60a8f4a3ce91401fc1ae76a6d1711a2858c7
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
Yes |
| Uploaded via |
twine/5.1.1 CPython/3.12.7
|
Release files / pyrfd-1.0.1-py3-none-any.whl
| Download URL | pyrfd-1.0.1-py3-none-any.whl |
|---|---|
| Size | 17.0 kB |
| Tags | Python 3 |
|
SHA-256 checksum How to use checksums |
bc704ce7fc2ed5a57bbb2ba308460db216edba743da789adf21524e5ac54bc24
|
|
BLAKE2b-256 checksum How to use checksums |
747af27d012915ee4782562edd1a604156794709bf0f1f57e6f2e66edc446702
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
Yes |
| Uploaded via |
twine/5.1.1 CPython/3.12.7
|