pyreco 0.0.1__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- pyreco/__init__.py +6 -0
- pyreco/cross_validation.py +86 -0
- pyreco/custom_models.py +672 -0
- pyreco/layers.py +139 -0
- pyreco/metrics.py +50 -0
- pyreco/models.py +208 -0
- pyreco/network_prop_extractor.py +24 -0
- pyreco/optimizers.py +42 -0
- pyreco/plotting.py +54 -0
- pyreco/remove_transients.py +17 -0
- pyreco/utils_data.py +290 -0
- pyreco/utils_networks.py +164 -0
- pyreco-0.0.1.dist-info/METADATA +150 -0
- pyreco-0.0.1.dist-info/RECORD +16 -0
- pyreco-0.0.1.dist-info/WHEEL +4 -0
- pyreco-0.0.1.dist-info/licenses/LICENSE +201 -0
pyreco/__init__.py
ADDED
|
@@ -0,0 +1,86 @@
|
|
|
1
|
+
# Comment for cross val check
|
|
2
|
+
import numpy as np
|
|
3
|
+
from typing import List, Tuple
|
|
4
|
+
import warnings
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
def cross_val(model, X: np.ndarray, y: np.ndarray, n_splits: int, metric: str = 'mse') -> Tuple[
|
|
8
|
+
List[float], float, float]:
|
|
9
|
+
'''
|
|
10
|
+
Performs k-fold cross-validation on a given model.
|
|
11
|
+
|
|
12
|
+
Parameters:
|
|
13
|
+
model : object
|
|
14
|
+
The RC model to be validated.
|
|
15
|
+
X : np.ndarray
|
|
16
|
+
Feature matrix.
|
|
17
|
+
y : np.ndarray
|
|
18
|
+
Target vector.
|
|
19
|
+
n_splits : int
|
|
20
|
+
Number of folds.
|
|
21
|
+
metrics : str
|
|
22
|
+
Metric to evaluate.
|
|
23
|
+
|
|
24
|
+
Returns:
|
|
25
|
+
tuple
|
|
26
|
+
A tuple containing:
|
|
27
|
+
- list of metric's value for each fold
|
|
28
|
+
- mean of the metric values
|
|
29
|
+
- standard deviation of the metric values
|
|
30
|
+
'''
|
|
31
|
+
|
|
32
|
+
# issue warning if more than one metric is specified
|
|
33
|
+
if type(metric) is list:
|
|
34
|
+
if len(metric) > 1:
|
|
35
|
+
metric = metric[0]
|
|
36
|
+
warnings.warn('Only a single metric should be specified. Using the first metric in the list.')
|
|
37
|
+
|
|
38
|
+
# get indices for splitting the data
|
|
39
|
+
indices = np.arange(X.shape[0])
|
|
40
|
+
# shuffle the indices
|
|
41
|
+
np.random.shuffle(indices)
|
|
42
|
+
|
|
43
|
+
# split the indices into n_splits parts (floor division)
|
|
44
|
+
fold_sizes = np.full(n_splits, X.shape[0] // n_splits)
|
|
45
|
+
# distribute remaining data (caused by floor division) as far as possible
|
|
46
|
+
fold_sizes[:X.shape[0] % n_splits] += 1
|
|
47
|
+
|
|
48
|
+
# get shuffled indices for folds
|
|
49
|
+
current = 0
|
|
50
|
+
fold_indices = []
|
|
51
|
+
for fold_size in fold_sizes:
|
|
52
|
+
start, stop = current, current + fold_size
|
|
53
|
+
fold_indices.append(indices[start:stop])
|
|
54
|
+
current = stop
|
|
55
|
+
|
|
56
|
+
# perform cross-validation
|
|
57
|
+
metric_folds_values = []
|
|
58
|
+
mean_metric_value = []
|
|
59
|
+
std_dev_metric_value = []
|
|
60
|
+
|
|
61
|
+
for i in range(n_splits):
|
|
62
|
+
# select the test indices for the current fold
|
|
63
|
+
test_indices = fold_indices[i]
|
|
64
|
+
|
|
65
|
+
# select the train indices (test_indices are removed from indices)
|
|
66
|
+
train_indices = np.setdiff1d(indices, test_indices)
|
|
67
|
+
|
|
68
|
+
# split the data into training and testing sets
|
|
69
|
+
X_train, X_test = X[train_indices], X[test_indices]
|
|
70
|
+
y_train, y_test = y[train_indices], y[test_indices]
|
|
71
|
+
|
|
72
|
+
# train the model
|
|
73
|
+
model.fit(X_train, y_train)
|
|
74
|
+
|
|
75
|
+
# calculate metric value for fold
|
|
76
|
+
metric_value = model.evaluate(X_test, y_test, metric)
|
|
77
|
+
|
|
78
|
+
# append metric value of fold
|
|
79
|
+
metric_folds_values.append(metric_value)
|
|
80
|
+
|
|
81
|
+
# get mean accuracy and standard deviation of metric values of all folds
|
|
82
|
+
mean_metric_value = float(np.mean(metric_folds_values))
|
|
83
|
+
std_dev_metric_value = float(np.std(metric_folds_values))
|
|
84
|
+
|
|
85
|
+
# output the results
|
|
86
|
+
return metric_folds_values, mean_metric_value, std_dev_metric_value
|