EpistasisLab / EpistasisLab/tpot

Built-in Pytorch Classifiers Crash TPOT

Open
#1,149 5 comments 0 reactions 1 assignee Claimed by @JDRomano2 View on GitHub
being worked on bug enhancement
Dominant language
Jupyter Notebook
Stars
10.1k
Forks
1.6k
PR merge metrics
No merged PRs in 30d

Description

The two current pytorch-based TPOT classifiers are currently designed to support only binary targets, which I discovered in reviewing the GitHub repo. However, rather than failing safely, they crash TPOT when presented with a multi-class problem., e,g, MNIST digits.

## Context of the issue
I have been using TPOT with non-DL models and have been impressed with its ability to construct complex pipelines efficiently, with minimum guidance from the user. Much of my work involves deep learning, so I was happy to see that TPOT is supporting pytorch. After trying the TPOT NN example, which failed mightily on my system, I created a configuration dictionary (attached) that included only the two pytorch classifiers. These also caused a crash, but in doing so uncovered the specific issue, no current support for multi-class problems.

## Process to reproduce the issue

1. Load and run script Hale_Ex1_NN, making sure it can import from TPOT_NN_Cfg_T1. (Both scripts attached.)
2. Observe failures as shown below.

_Note: Duplicate _pre-test decorator error messages removed._

## Expected result

Ideally, the classifiers should support multi-class cases, but, failing that, avoid crashing the system, with an appropriate message, e.g., __multi-class_ support not yet _available.__

## Current result

```
/Users/david/opt/anaconda3/envs/TPOT_Torch/bin/python /Applications/PyCharm.app/Contents/plugins/python/helpers/pydev/pydevconsole.py --mode=client --port=49961
import sys; print('Python %s on %s' % (sys.version, sys.platform))
sys.path.extend(['/Users/david/PycharmProjects/KJStraddle'])
Python 3.8.3 (default, Jul 2 2020, 11:26:31)
Type 'copyright', 'credits' or 'license' for more information
IPython 7.19.0 -- An enhanced Interactive Python. Type '?' for help.
PyDev console: using IPython 7.19.0
Python 3.8.3 (default, Jul 2 2020, 11:26:31)
[Clang 10.0.0 ] on darwin
runfile('/Users/david/PycharmProjects/KJStraddle/Hale_Ex1_NN.py', wdir='/Users/david/PycharmProjects/KJStraddle')
Starting iteration 0
7 operators have been imported by TPOT.
_pre_test decorator: _random_mutation_operator: num_test=0 Non-binary targets not supported.
_pre_test decorator: _random_mutation_operator: num_test=1 Non-binary targets not supported.
_pre_test decorator: _random_mutation_operator: num_test=2 Non-binary targets not supported.
_pre_test decorator: _random_mutation_operator: num_test=3 Non-binary targets not supported.
_pre_test decorator: _random_mutation_operator: num_test=4 Non-binary targets not supported.
_pre_test decorator: _random_mutation_operator: num_test=5 Non-binary targets not supported.
_pre_test decorator: _random_mutation_operator: num_test=6 Non-binary targets not supported.
_pre_test decorator: _random_mutation_operator: num_test=7 Non-binary targets not supported.
_pre_test decorator: _random_mutation_operator: num_test=8 Non-binary targets not supported.
_pre_test decorator: _random_mutation_operator: num_test=9 Non-binary targets not supported.
_pre_test decorator: _mate_operator: num_test=0 Non-binary targets not supported.
_pre_test decorator: _mate_operator: num_test=1 Non-binary targets not supported.
_pre_test decorator: _mate_operator: num_test=2 Non-binary targets not supported.
_pre_test decorator: _mate_operator: num_test=3 Non-binary targets not supported.
_pre_test decorator: _mate_operator: num_test=4 Non-binary targets not supported.
_pre_test decorator: _mate_operator: num_test=5 Non-binary targets not supported.
_pre_test decorator: _mate_operator: num_test=6 Non-binary targets not supported.
_pre_test decorator: _mate_operator: num_test=7 Non-binary targets not supported.
_pre_test decorator: _mate_operator: num_test=8 Non-binary targets not supported.
_pre_test decorator: _mate_operator: num_test=9 Non-binary targets not supported.

Pipeline encountered that has previously been evaluated during the optimization process. Using the score from the previous evaluation.

Generation 1 - Current Pareto front scores:

-1 -inf PytorchMLPClassifier(input_matrix, PytorchMLPClassifier__batch_size=16, PytorchMLPClassifier__learning_rate=0.01, PytorchMLPClassifier__num_epochs=10, PytorchMLPClassifier__weight_decay=0.001)
/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/sklearn/model_selection/_split.py:670: UserWarning: The least populated class in y has only 1 members, which is less than n_splits=5.
warnings.warn(("The least populated class in y has only %d"
/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/sklearn/utils/validation.py:67: FutureWarning: Pass allow_nan=[8 0 6 4 2 4 0 3 3 3 9 3 6 8 4 4 1 5 1 6 7 6 6 1 0 9 2 6 4 0 5 5 9 4 2 6 5
9 4 8] as keyword args. From version 0.25 passing these as positional arguments will result in an error
warnings.warn("Pass {} as keyword args. From version 0.25 "
/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/sklearn/model_selection/_split.py:670: UserWarning: The least populated class in y has only 1 members, which is less than n_splits=5.
warnings.warn(("The least populated class in y has only %d"
/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/sklearn/utils/validation.py:67: FutureWarning: Pass allow_nan=[8 0 6 4 2 4 0 3 3 3 9 3 6 8 4 4 1 5 1 6 7 6 6 1 0 9 2 6 4 0 5 5 9 4 2 6 5
9 4 8] as keyword args. From version 0.25 passing these as positional arguments will result in an error
warnings.warn("Pass {} as keyword args. From version 0.25 "
/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/sklearn/model_selection/_split.py:670: UserWarning: The least populated class in y has only 1 members, which is less than n_splits=5.
warnings.warn(("The least populated class in y has only %d"
/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/sklearn/utils/validation.py:67: FutureWarning: Pass allow_nan=[8 0 6 4 2 4 0 3 3 3 9 3 6 8 4 4 1 5 1 6 7 6 6 1 0 9 2 6 4 0 5 5 9 4 2 6 5
9 4 8] as keyword args. From version 0.25 passing these as positional arguments will result in an error
warnings.warn("Pass {} as keyword args. From version 0.25 "
/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/sklearn/model_selection/_split.py:670: UserWarning: The least populated class in y has only 1 members, which is less than n_splits=5.
warnings.warn(("The least populated class in y has only %d"
/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/sklearn/utils/validation.py:67: FutureWarning: Pass allow_nan=[8 0 6 4 2 4 0 3 3 3 9 3 6 8 4 4 1 5 1 6 7 6 6 1 0 9 2 6 4 0 5 5 9 4 2 6 5
9 4 8] as keyword args. From version 0.25 passing these as positional arguments will result in an error
warnings.warn("Pass {} as keyword args. From version 0.25 "
/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/sklearn/model_selection/_split.py:670: UserWarning: The least populated class in y has only 1 members, which is less than n_splits=5.
warnings.warn(("The least populated class in y has only %d"
/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/sklearn/utils/validation.py:67: FutureWarning: Pass allow_nan=[8 0 6 4 2 4 0 3 3 3 9 3 6 8 4 4 1 5 1 6 7 6 6 1 0 9 2 6 4 0 5 5 9 4 2 6 5
9 4 8] as keyword args. From version 0.25 passing these as positional arguments will result in an error
warnings.warn("Pass {} as keyword args. From version 0.25 "
/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/sklearn/model_selection/_split.py:670: UserWarning: The least populated class in y has only 1 members, which is less than n_splits=5.
warnings.warn(("The least populated class in y has only %d"
/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/sklearn/utils/validation.py:67: FutureWarning: Pass allow_nan=[8 0 6 4 2 4 0 3 3 3 9 3 6 8 4 4 1 5 1 6 7 6 6 1 0 9 2 6 4 0 5 5 9 4 2 6 5
9 4 8] as keyword args. From version 0.25 passing these as positional arguments will result in an error
warnings.warn("Pass {} as keyword args. From version 0.25 "
/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/sklearn/model_selection/_split.py:670: UserWarning: The least populated class in y has only 1 members, which is less than n_splits=5.
warnings.warn(("The least populated class in y has only %d"
/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/sklearn/utils/validation.py:67: FutureWarning: Pass allow_nan=[8 0 6 4 2 4 0 3 3 3 9 3 6 8 4 4 1 5 1 6 7 6 6 1 0 9 2 6 4 0 5 5 9 4 2 6 5
9 4 8] as keyword args. From version 0.25 passing these as positional arguments will result in an error
warnings.warn("Pass {} as keyword args. From version 0.25 "
/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/sklearn/model_selection/_split.py:670: UserWarning: The least populated class in y has only 1 members, which is less than n_splits=5.
warnings.warn(("The least populated class in y has only %d"
/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/sklearn/utils/validation.py:67: FutureWarning: Pass allow_nan=[8 0 6 4 2 4 0 3 3 3 9 3 6 8 4 4 1 5 1 6 7 6 6 1 0 9 2 6 4 0 5 5 9 4 2 6 5
9 4 8] as keyword args. From version 0.25 passing these as positional arguments will result in an error
warnings.warn("Pass {} as keyword args. From version 0.25 "
/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/sklearn/model_selection/_split.py:670: UserWarning: The least populated class in y has only 1 members, which is less than n_splits=5.
warnings.warn(("The least populated class in y has only %d"
/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/sklearn/utils/validation.py:67: FutureWarning: Pass allow_nan=[8 0 6 4 2 4 0 3 3 3 9 3 6 8 4 4 1 5 1 6 7 6 6 1 0 9 2 6 4 0 5 5 9 4 2 6 5
9 4 8] as keyword args. From version 0.25 passing these as positional arguments will result in an error
warnings.warn("Pass {} as keyword args. From version 0.25 "
/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/sklearn/model_selection/_split.py:670: UserWarning: The least populated class in y has only 1 members, which is less than n_splits=5.
warnings.warn(("The least populated class in y has only %d"
/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/sklearn/utils/validation.py:67: FutureWarning: Pass allow_nan=[8 0 6 4 2 4 0 3 3 3 9 3 6 8 4 4 1 5 1 6 7 6 6 1 0 9 2 6 4 0 5 5 9 4 2 6 5
9 4 8] as keyword args. From version 0.25 passing these as positional arguments will result in an error
warnings.warn("Pass {} as keyword args. From version 0.25 "
Traceback (most recent call last):
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/tpot/base.py", line 730, in fit
self._pop, _ = eaMuPlusLambda(
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/tpot/gp_deap.py", line 281, in eaMuPlusLambda
per_generation_function(gen)
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/tpot/base.py", line 1052, in _check_periodic_pipeline
self._update_top_pipeline()
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/tpot/base.py", line 830, in _update_top_pipeline
cv_scores = cross_val_score(sklearn_pipeline,
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/sklearn/utils/validation.py", line 72, in inner_f
return f(**kwargs)
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/sklearn/model_selection/_validation.py", line 401, in cross_val_score
cv_results = cross_validate(estimator=estimator, X=X, y=y, groups=groups,
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/sklearn/utils/validation.py", line 72, in inner_f
return f(**kwargs)
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/sklearn/model_selection/_validation.py", line 242, in cross_validate
scores = parallel(
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/joblib/parallel.py", line 1041, in __call__
if self.dispatch_one_batch(iterator):
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/joblib/parallel.py", line 859, in dispatch_one_batch
self._dispatch(tasks)
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/joblib/parallel.py", line 777, in _dispatch
job = self._backend.apply_async(batch, callback=cb)
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/joblib/_parallel_backends.py", line 208, in apply_async
result = ImmediateResult(func)
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/joblib/_parallel_backends.py", line 572, in __init__
self.results = batch()
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/joblib/parallel.py", line 262, in __call__
return [func(*args, **kwargs)
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/joblib/parallel.py", line 262, in
return [func(*args, **kwargs)
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/sklearn/model_selection/_validation.py", line 531, in _fit_and_score
estimator.fit(X_train, y_train, **fit_params)
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/sklearn/pipeline.py", line 335, in fit
self._final_estimator.fit(Xt, y, **fit_params_last_step)
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/tpot/builtins/nn.py", line 122, in fit
self._init_model(X, y)
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/tpot/builtins/nn.py", line 315, in _init_model
X, y = self.validate_inputs(X, y)
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/tpot/builtins/nn.py", line 166, in validate_inputs
raise ValueError("Non-binary targets not supported")
ValueError: Non-binary targets not supported
During handling of the above exception, another exception occurred:
Traceback (most recent call last):
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/IPython/core/interactiveshell.py", line 3418, in run_code
exec(code_obj, self.user_global_ns, self.user_ns)
File "", line 1, in
runfile('/Users/david/PycharmProjects/KJStraddle/Hale_Ex1_NN.py', wdir='/Users/david/PycharmProjects/KJStraddle')
File "/Applications/PyCharm.app/Contents/plugins/python/helpers/pydev/_pydev_bundle/pydev_umd.py", line 197, in runfile
pydev_imports.execfile(filename, global_vars, local_vars) # execute the script
File "/Applications/PyCharm.app/Contents/plugins/python/helpers/pydev/_pydev_imps/_pydev_execfile.py", line 18, in execfile
exec(compile(contents+"\n", file, 'exec'), glob, loc)
File "/Users/david/PycharmProjects/KJStraddle/Hale_Ex1_NN.py", line 34, in
tpot.fit(X_train, y_train)
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/tpot/base.py", line 773, in fit
raise e
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/tpot/base.py", line 764, in fit
self._update_top_pipeline()
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/tpot/base.py", line 830, in _update_top_pipeline
cv_scores = cross_val_score(sklearn_pipeline,
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/sklearn/utils/validation.py", line 72, in inner_f
return f(**kwargs)
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/sklearn/model_selection/_validation.py", line 401, in cross_val_score
cv_results = cross_validate(estimator=estimator, X=X, y=y, groups=groups,
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/sklearn/utils/validation.py", line 72, in inner_f
return f(**kwargs)
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/sklearn/model_selection/_validation.py", line 242, in cross_validate
scores = parallel(
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/joblib/parallel.py", line 1041, in __call__
if self.dispatch_one_batch(iterator):
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/joblib/parallel.py", line 859, in dispatch_one_batch
self._dispatch(tasks)
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/joblib/parallel.py", line 777, in _dispatch
job = self._backend.apply_async(batch, callback=cb)
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/joblib/_parallel_backends.py", line 208, in apply_async
result = ImmediateResult(func)
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/joblib/_parallel_backends.py", line 572, in __init__
self.results = batch()
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/joblib/parallel.py", line 262, in __call__
return [func(*args, **kwargs)
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/joblib/parallel.py", line 262, in
return [func(*args, **kwargs)
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/sklearn/model_selection/_validation.py", line 531, in _fit_and_score
estimator.fit(X_train, y_train, **fit_params)
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/sklearn/pipeline.py", line 335, in fit
self._final_estimator.fit(Xt, y, **fit_params_last_step)
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/tpot/builtins/nn.py", line 122, in fit
self._init_model(X, y)
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/tpot/builtins/nn.py", line 315, in _init_model
X, y = self.validate_inputs(X, y)
File "/Users/david/opt/anaconda3/envs/TPOT_Torch/lib/python3.8/site-packages/tpot/builtins/nn.py", line 166, in validate_inputs
raise ValueError("Non-binary targets not supported")
ValueError: Non-binary targets not supported
```
## Possible fix

I volunteer to 1) add safe failure code and 2) extend classifiers to support multi-target cases (or assist as needed).

Contributor guide

No contributing guide indexed for this repository

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.