Mercurial > repos > bgruening > keras_train_and_eval
annotate train_test_eval.py @ 0:03f61bb3ca43 draft
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
| author | bgruening | 
|---|---|
| date | Mon, 16 Dec 2019 05:36:53 -0500 | 
| parents | |
| children | 3866911c93ae | 
| rev | line source | 
|---|---|
| 0 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 1 import argparse | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 2 import joblib | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 3 import json | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 4 import numpy as np | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 5 import os | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 6 import pandas as pd | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 7 import pickle | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 8 import warnings | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 9 from itertools import chain | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 10 from scipy.io import mmread | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 11 from sklearn.base import clone | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 12 from sklearn import (cluster, compose, decomposition, ensemble, | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 13 feature_extraction, feature_selection, | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 14 gaussian_process, kernel_approximation, metrics, | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 15 model_selection, naive_bayes, neighbors, | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 16 pipeline, preprocessing, svm, linear_model, | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 17 tree, discriminant_analysis) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 18 from sklearn.exceptions import FitFailedWarning | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 19 from sklearn.metrics.scorer import _check_multimetric_scoring | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 20 from sklearn.model_selection._validation import _score, cross_validate | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 21 from sklearn.model_selection import _search, _validation | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 22 from sklearn.utils import indexable, safe_indexing | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 23 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 24 from galaxy_ml.model_validations import train_test_split | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 25 from galaxy_ml.utils import (SafeEval, get_scoring, load_model, | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 26 read_columns, try_get_attr, get_module) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 27 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 28 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 29 _fit_and_score = try_get_attr('galaxy_ml.model_validations', '_fit_and_score') | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 30 setattr(_search, '_fit_and_score', _fit_and_score) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 31 setattr(_validation, '_fit_and_score', _fit_and_score) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 32 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 33 N_JOBS = int(os.environ.get('GALAXY_SLOTS', 1)) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 34 CACHE_DIR = os.path.join(os.getcwd(), 'cached') | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 35 del os | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 36 NON_SEARCHABLE = ('n_jobs', 'pre_dispatch', 'memory', '_path', | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 37 'nthread', 'callbacks') | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 38 ALLOWED_CALLBACKS = ('EarlyStopping', 'TerminateOnNaN', 'ReduceLROnPlateau', | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 39 'CSVLogger', 'None') | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 40 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 41 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 42 def _eval_swap_params(params_builder): | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 43 swap_params = {} | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 44 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 45 for p in params_builder['param_set']: | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 46 swap_value = p['sp_value'].strip() | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 47 if swap_value == '': | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 48 continue | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 49 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 50 param_name = p['sp_name'] | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 51 if param_name.lower().endswith(NON_SEARCHABLE): | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 52 warnings.warn("Warning: `%s` is not eligible for search and was " | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 53 "omitted!" % param_name) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 54 continue | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 55 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 56 if not swap_value.startswith(':'): | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 57 safe_eval = SafeEval(load_scipy=True, load_numpy=True) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 58 ev = safe_eval(swap_value) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 59 else: | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 60 # Have `:` before search list, asks for estimator evaluatio | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 61 safe_eval_es = SafeEval(load_estimators=True) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 62 swap_value = swap_value[1:].strip() | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 63 # TODO maybe add regular express check | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 64 ev = safe_eval_es(swap_value) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 65 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 66 swap_params[param_name] = ev | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 67 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 68 return swap_params | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 69 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 70 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 71 def train_test_split_none(*arrays, **kwargs): | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 72 """extend train_test_split to take None arrays | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 73 and support split by group names. | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 74 """ | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 75 nones = [] | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 76 new_arrays = [] | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 77 for idx, arr in enumerate(arrays): | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 78 if arr is None: | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 79 nones.append(idx) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 80 else: | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 81 new_arrays.append(arr) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 82 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 83 if kwargs['shuffle'] == 'None': | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 84 kwargs['shuffle'] = None | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 85 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 86 group_names = kwargs.pop('group_names', None) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 87 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 88 if group_names is not None and group_names.strip(): | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 89 group_names = [name.strip() for name in | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 90 group_names.split(',')] | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 91 new_arrays = indexable(*new_arrays) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 92 groups = kwargs['labels'] | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 93 n_samples = new_arrays[0].shape[0] | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 94 index_arr = np.arange(n_samples) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 95 test = index_arr[np.isin(groups, group_names)] | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 96 train = index_arr[~np.isin(groups, group_names)] | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 97 rval = list(chain.from_iterable( | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 98 (safe_indexing(a, train), | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 99 safe_indexing(a, test)) for a in new_arrays)) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 100 else: | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 101 rval = train_test_split(*new_arrays, **kwargs) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 102 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 103 for pos in nones: | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 104 rval[pos * 2: 2] = [None, None] | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 105 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 106 return rval | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 107 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 108 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 109 def main(inputs, infile_estimator, infile1, infile2, | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 110 outfile_result, outfile_object=None, | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 111 outfile_weights=None, groups=None, | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 112 ref_seq=None, intervals=None, targets=None, | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 113 fasta_path=None): | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 114 """ | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 115 Parameter | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 116 --------- | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 117 inputs : str | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 118 File path to galaxy tool parameter | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 119 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 120 infile_estimator : str | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 121 File path to estimator | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 122 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 123 infile1 : str | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 124 File path to dataset containing features | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 125 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 126 infile2 : str | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 127 File path to dataset containing target values | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 128 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 129 outfile_result : str | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 130 File path to save the results, either cv_results or test result | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 131 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 132 outfile_object : str, optional | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 133 File path to save searchCV object | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 134 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 135 outfile_weights : str, optional | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 136 File path to save deep learning model weights | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 137 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 138 groups : str | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 139 File path to dataset containing groups labels | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 140 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 141 ref_seq : str | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 142 File path to dataset containing genome sequence file | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 143 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 144 intervals : str | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 145 File path to dataset containing interval file | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 146 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 147 targets : str | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 148 File path to dataset compressed target bed file | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 149 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 150 fasta_path : str | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 151 File path to dataset containing fasta file | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 152 """ | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 153 warnings.simplefilter('ignore') | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 154 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 155 with open(inputs, 'r') as param_handler: | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 156 params = json.load(param_handler) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 157 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 158 # load estimator | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 159 with open(infile_estimator, 'rb') as estimator_handler: | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 160 estimator = load_model(estimator_handler) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 161 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 162 # swap hyperparameter | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 163 swapping = params['experiment_schemes']['hyperparams_swapping'] | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 164 swap_params = _eval_swap_params(swapping) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 165 estimator.set_params(**swap_params) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 166 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 167 estimator_params = estimator.get_params() | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 168 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 169 # store read dataframe object | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 170 loaded_df = {} | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 171 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 172 input_type = params['input_options']['selected_input'] | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 173 # tabular input | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 174 if input_type == 'tabular': | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 175 header = 'infer' if params['input_options']['header1'] else None | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 176 column_option = (params['input_options']['column_selector_options_1'] | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 177 ['selected_column_selector_option']) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 178 if column_option in ['by_index_number', 'all_but_by_index_number', | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 179 'by_header_name', 'all_but_by_header_name']: | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 180 c = params['input_options']['column_selector_options_1']['col1'] | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 181 else: | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 182 c = None | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 183 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 184 df_key = infile1 + repr(header) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 185 df = pd.read_csv(infile1, sep='\t', header=header, | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 186 parse_dates=True) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 187 loaded_df[df_key] = df | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 188 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 189 X = read_columns(df, c=c, c_option=column_option).astype(float) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 190 # sparse input | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 191 elif input_type == 'sparse': | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 192 X = mmread(open(infile1, 'r')) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 193 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 194 # fasta_file input | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 195 elif input_type == 'seq_fasta': | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 196 pyfaidx = get_module('pyfaidx') | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 197 sequences = pyfaidx.Fasta(fasta_path) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 198 n_seqs = len(sequences.keys()) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 199 X = np.arange(n_seqs)[:, np.newaxis] | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 200 for param in estimator_params.keys(): | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 201 if param.endswith('fasta_path'): | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 202 estimator.set_params( | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 203 **{param: fasta_path}) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 204 break | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 205 else: | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 206 raise ValueError( | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 207 "The selected estimator doesn't support " | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 208 "fasta file input! Please consider using " | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 209 "KerasGBatchClassifier with " | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 210 "FastaDNABatchGenerator/FastaProteinBatchGenerator " | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 211 "or having GenomeOneHotEncoder/ProteinOneHotEncoder " | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 212 "in pipeline!") | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 213 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 214 elif input_type == 'refseq_and_interval': | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 215 path_params = { | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 216 'data_batch_generator__ref_genome_path': ref_seq, | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 217 'data_batch_generator__intervals_path': intervals, | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 218 'data_batch_generator__target_path': targets | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 219 } | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 220 estimator.set_params(**path_params) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 221 n_intervals = sum(1 for line in open(intervals)) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 222 X = np.arange(n_intervals)[:, np.newaxis] | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 223 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 224 # Get target y | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 225 header = 'infer' if params['input_options']['header2'] else None | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 226 column_option = (params['input_options']['column_selector_options_2'] | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 227 ['selected_column_selector_option2']) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 228 if column_option in ['by_index_number', 'all_but_by_index_number', | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 229 'by_header_name', 'all_but_by_header_name']: | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 230 c = params['input_options']['column_selector_options_2']['col2'] | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 231 else: | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 232 c = None | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 233 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 234 df_key = infile2 + repr(header) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 235 if df_key in loaded_df: | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 236 infile2 = loaded_df[df_key] | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 237 else: | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 238 infile2 = pd.read_csv(infile2, sep='\t', | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 239 header=header, parse_dates=True) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 240 loaded_df[df_key] = infile2 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 241 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 242 y = read_columns( | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 243 infile2, | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 244 c=c, | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 245 c_option=column_option, | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 246 sep='\t', | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 247 header=header, | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 248 parse_dates=True) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 249 if len(y.shape) == 2 and y.shape[1] == 1: | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 250 y = y.ravel() | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 251 if input_type == 'refseq_and_interval': | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 252 estimator.set_params( | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 253 data_batch_generator__features=y.ravel().tolist()) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 254 y = None | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 255 # end y | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 256 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 257 # load groups | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 258 if groups: | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 259 groups_selector = (params['experiment_schemes']['test_split'] | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 260 ['split_algos']).pop('groups_selector') | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 261 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 262 header = 'infer' if groups_selector['header_g'] else None | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 263 column_option = \ | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 264 (groups_selector['column_selector_options_g'] | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 265 ['selected_column_selector_option_g']) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 266 if column_option in ['by_index_number', 'all_but_by_index_number', | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 267 'by_header_name', 'all_but_by_header_name']: | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 268 c = groups_selector['column_selector_options_g']['col_g'] | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 269 else: | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 270 c = None | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 271 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 272 df_key = groups + repr(header) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 273 if df_key in loaded_df: | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 274 groups = loaded_df[df_key] | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 275 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 276 groups = read_columns( | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 277 groups, | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 278 c=c, | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 279 c_option=column_option, | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 280 sep='\t', | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 281 header=header, | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 282 parse_dates=True) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 283 groups = groups.ravel() | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 284 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 285 # del loaded_df | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 286 del loaded_df | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 287 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 288 # handle memory | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 289 memory = joblib.Memory(location=CACHE_DIR, verbose=0) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 290 # cache iraps_core fits could increase search speed significantly | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 291 if estimator.__class__.__name__ == 'IRAPSClassifier': | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 292 estimator.set_params(memory=memory) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 293 else: | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 294 # For iraps buried in pipeline | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 295 new_params = {} | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 296 for p, v in estimator_params.items(): | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 297 if p.endswith('memory'): | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 298 # for case of `__irapsclassifier__memory` | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 299 if len(p) > 8 and p[:-8].endswith('irapsclassifier'): | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 300 # cache iraps_core fits could increase search | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 301 # speed significantly | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 302 new_params[p] = memory | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 303 # security reason, we don't want memory being | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 304 # modified unexpectedly | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 305 elif v: | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 306 new_params[p] = None | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 307 # handle n_jobs | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 308 elif p.endswith('n_jobs'): | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 309 # For now, 1 CPU is suggested for iprasclassifier | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 310 if len(p) > 8 and p[:-8].endswith('irapsclassifier'): | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 311 new_params[p] = 1 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 312 else: | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 313 new_params[p] = N_JOBS | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 314 # for security reason, types of callback are limited | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 315 elif p.endswith('callbacks'): | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 316 for cb in v: | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 317 cb_type = cb['callback_selection']['callback_type'] | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 318 if cb_type not in ALLOWED_CALLBACKS: | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 319 raise ValueError( | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 320 "Prohibited callback type: %s!" % cb_type) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 321 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 322 estimator.set_params(**new_params) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 323 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 324 # handle scorer, convert to scorer dict | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 325 scoring = params['experiment_schemes']['metrics']['scoring'] | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 326 scorer = get_scoring(scoring) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 327 scorer, _ = _check_multimetric_scoring(estimator, scoring=scorer) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 328 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 329 # handle test (first) split | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 330 test_split_options = (params['experiment_schemes'] | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 331 ['test_split']['split_algos']) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 332 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 333 if test_split_options['shuffle'] == 'group': | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 334 test_split_options['labels'] = groups | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 335 if test_split_options['shuffle'] == 'stratified': | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 336 if y is not None: | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 337 test_split_options['labels'] = y | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 338 else: | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 339 raise ValueError("Stratified shuffle split is not " | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 340 "applicable on empty target values!") | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 341 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 342 X_train, X_test, y_train, y_test, groups_train, groups_test = \ | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 343 train_test_split_none(X, y, groups, **test_split_options) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 344 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 345 exp_scheme = params['experiment_schemes']['selected_exp_scheme'] | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 346 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 347 # handle validation (second) split | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 348 if exp_scheme == 'train_val_test': | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 349 val_split_options = (params['experiment_schemes'] | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 350 ['val_split']['split_algos']) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 351 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 352 if val_split_options['shuffle'] == 'group': | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 353 val_split_options['labels'] = groups_train | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 354 if val_split_options['shuffle'] == 'stratified': | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 355 if y_train is not None: | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 356 val_split_options['labels'] = y_train | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 357 else: | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 358 raise ValueError("Stratified shuffle split is not " | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 359 "applicable on empty target values!") | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 360 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 361 X_train, X_val, y_train, y_val, groups_train, groups_val = \ | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 362 train_test_split_none(X_train, y_train, groups_train, | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 363 **val_split_options) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 364 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 365 # train and eval | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 366 if hasattr(estimator, 'validation_data'): | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 367 if exp_scheme == 'train_val_test': | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 368 estimator.fit(X_train, y_train, | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 369 validation_data=(X_val, y_val)) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 370 else: | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 371 estimator.fit(X_train, y_train, | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 372 validation_data=(X_test, y_test)) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 373 else: | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 374 estimator.fit(X_train, y_train) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 375 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 376 if hasattr(estimator, 'evaluate'): | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 377 scores = estimator.evaluate(X_test, y_test=y_test, | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 378 scorer=scorer, | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 379 is_multimetric=True) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 380 else: | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 381 scores = _score(estimator, X_test, y_test, scorer, | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 382 is_multimetric=True) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 383 # handle output | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 384 for name, score in scores.items(): | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 385 scores[name] = [score] | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 386 df = pd.DataFrame(scores) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 387 df = df[sorted(df.columns)] | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 388 df.to_csv(path_or_buf=outfile_result, sep='\t', | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 389 header=True, index=False) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 390 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 391 memory.clear(warn=False) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 392 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 393 if outfile_object: | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 394 main_est = estimator | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 395 if isinstance(estimator, pipeline.Pipeline): | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 396 main_est = estimator.steps[-1][-1] | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 397 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 398 if hasattr(main_est, 'model_') \ | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 399 and hasattr(main_est, 'save_weights'): | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 400 if outfile_weights: | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 401 main_est.save_weights(outfile_weights) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 402 del main_est.model_ | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 403 del main_est.fit_params | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 404 del main_est.model_class_ | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 405 del main_est.validation_data | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 406 if getattr(main_est, 'data_generator_', None): | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 407 del main_est.data_generator_ | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 408 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 409 with open(outfile_object, 'wb') as output_handler: | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 410 pickle.dump(estimator, output_handler, | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 411 pickle.HIGHEST_PROTOCOL) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 412 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 413 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 414 if __name__ == '__main__': | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 415 aparser = argparse.ArgumentParser() | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 416 aparser.add_argument("-i", "--inputs", dest="inputs", required=True) | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 417 aparser.add_argument("-e", "--estimator", dest="infile_estimator") | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 418 aparser.add_argument("-X", "--infile1", dest="infile1") | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 419 aparser.add_argument("-y", "--infile2", dest="infile2") | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 420 aparser.add_argument("-O", "--outfile_result", dest="outfile_result") | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 421 aparser.add_argument("-o", "--outfile_object", dest="outfile_object") | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 422 aparser.add_argument("-w", "--outfile_weights", dest="outfile_weights") | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 423 aparser.add_argument("-g", "--groups", dest="groups") | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 424 aparser.add_argument("-r", "--ref_seq", dest="ref_seq") | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 425 aparser.add_argument("-b", "--intervals", dest="intervals") | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 426 aparser.add_argument("-t", "--targets", dest="targets") | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 427 aparser.add_argument("-f", "--fasta_path", dest="fasta_path") | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 428 args = aparser.parse_args() | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 429 | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 430 main(args.inputs, args.infile_estimator, args.infile1, args.infile2, | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 431 args.outfile_result, outfile_object=args.outfile_object, | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 432 outfile_weights=args.outfile_weights, groups=args.groups, | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 433 ref_seq=args.ref_seq, intervals=args.intervals, | 
| 
03f61bb3ca43
"planemo upload for repository https://github.com/bgruening/galaxytools/tree/master/tools/sklearn commit 5b2ac730ec6d3b762faa9034eddd19ad1b347476"
 bgruening parents: diff
changeset | 434 targets=args.targets, fasta_path=args.fasta_path) | 
