Provide wrappers for popular ML libraries
- Dominant language
- Python
- Stars
- 951
- Forks
- 262
- PR merge metrics
- No merged PRs in 30d
Description
It'd be convenient to provide support for use of Keras or PyTorch models in model selection. There are two issues:
1. Keras/PyTorch models don't conform to the Scikit-learn API.
2. Keras models are not pickle-able.
I'm imaging this interface:
``` python
from torchvision.models import resnet18
import torch.optim as optim
from dask_ml.wrappers import PyTorchClassifier
pytorch_model = resnet18()
sklearn_model = SkorchClassifier(
model=pytorch_model,
model__alpha=1e-2, # if resnet18 had a kwarg `alpha`
optimizer=optim.SGD,
optimizer__lr=0.1,
)
```
[dhc]:https://github.com/stsievert/dask-hyperband-comparison/
**Related issues/PRs**
Same complaint in dask/distributed: https://github.com/dask/distributed/issues/3873
Contributor guide
Assessment
This issue has not been assessed yet.