dask / dask/dask-ml

Provide wrappers for popular ML libraries

Open
#696 14 comments 0 reactions 0 assignees View on GitHub
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

Open the contributing guide

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.