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 ADDED
@@ -0,0 +1,6 @@
1
+ #
2
+ # from . import layers
3
+ # from .layers import InputLayer, ReservoirLayer, ReadoutLayer
4
+ #
5
+ # from . import custom_models
6
+ # from .custom_models import RC
@@ -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