zeroth-learn 0.2.3__tar.gz → 0.2.4__tar.gz

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.
Files changed (49) hide show
  1. {zeroth_learn-0.2.3/zeroth_learn.egg-info → zeroth_learn-0.2.4}/PKG-INFO +1 -1
  2. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/pyproject.toml +1 -1
  3. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/abstract/model.py +24 -41
  4. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/abstract/summary.py +4 -8
  5. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/experiment.py +10 -10
  6. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/plot_losses.py +5 -5
  7. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/zeroth_order/gradient_estimators.py +15 -15
  8. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/zeroth_order/optimizers.py +12 -12
  9. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4/zeroth_learn.egg-info}/PKG-INFO +1 -1
  10. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/MANIFEST.in +0 -0
  11. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/README.md +0 -0
  12. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/setup.cfg +0 -0
  13. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/__init__.py +0 -0
  14. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/abstract/__init__.py +0 -0
  15. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/abstract/activation.py +0 -0
  16. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/abstract/blackbox.py +0 -0
  17. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/abstract/data_creator.py +0 -0
  18. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/abstract/loss.py +0 -0
  19. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/abstract/metric.py +0 -0
  20. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/abstract/neural_network.py +0 -0
  21. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/abstract/optimizer.py +0 -0
  22. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/abstract/perturbation_matrix.py +0 -0
  23. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/all.py +0 -0
  24. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/data.py +0 -0
  25. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/first_order/__init__.py +0 -0
  26. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/first_order/layer.py +0 -0
  27. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/first_order/model.py +0 -0
  28. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/first_order/neural_network.py +0 -0
  29. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/first_order/optimizers.py +0 -0
  30. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/losses/__init__.py +0 -0
  31. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/losses/cross_entropy.py +0 -0
  32. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/losses/mse.py +0 -0
  33. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/paths.py +0 -0
  34. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/types.py +0 -0
  35. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/utils/__init__.py +0 -0
  36. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/utils/activation_functions.py +0 -0
  37. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/utils/dataclasses_utils.py +0 -0
  38. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/utils/metrics.py +0 -0
  39. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/utils/perturbation_matrices.py +0 -0
  40. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/zeroth_order/__init__.py +0 -0
  41. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/zeroth_order/model.py +0 -0
  42. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/zeroth_order/neural_network/__init__.py +0 -0
  43. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/zeroth_order/neural_network/neural_network.py +0 -0
  44. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/zeroth_order/neural_network/parameter_manager.py +0 -0
  45. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth/zeroth_order/zeroth_order_blackbox.py +0 -0
  46. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth_learn.egg-info/SOURCES.txt +0 -0
  47. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth_learn.egg-info/dependency_links.txt +0 -0
  48. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth_learn.egg-info/requires.txt +0 -0
  49. {zeroth_learn-0.2.3 → zeroth_learn-0.2.4}/zeroth_learn.egg-info/top_level.txt +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: zeroth-learn
3
- Version: 0.2.3
3
+ Version: 0.2.4
4
4
  Requires-Dist: numpy
5
5
  Requires-Dist: pandas
6
6
  Requires-Dist: matplotlib
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
4
4
 
5
5
  [project]
6
6
  name = "zeroth-learn"
7
- version = "0.2.3"
7
+ version = "0.2.4"
8
8
  dependencies = [
9
9
  "numpy",
10
10
  "pandas",
@@ -1,10 +1,9 @@
1
1
  from __future__ import annotations
2
2
 
3
- import json
4
- import os
5
3
  import pickle
6
4
  from abc import ABC, abstractmethod
7
5
  from dataclasses import dataclass
6
+ from pathlib import Path
8
7
  from typing import Callable
9
8
 
10
9
  import matplotlib.pyplot as plt
@@ -20,35 +19,6 @@ from ..plot_losses import plot_losses
20
19
  from ..types import Array
21
20
 
22
21
 
23
- @dataclass
24
- class ModelRecord:
25
- name: str
26
- id: dict
27
- training_loss: Array
28
-
29
- @classmethod
30
- def load(cls, model_dir_path: str):
31
- config_path = os.path.join(model_dir_path, "config.json")
32
- with open(config_path, "r") as f:
33
- config = json.load(f)
34
-
35
- loss_path = os.path.join(model_dir_path, LOSS_FILE)
36
- df_loss = pd.read_csv(loss_path)
37
-
38
- return cls(
39
- id=config.get("id", {}),
40
- name=config.get("name", "Unknown"),
41
- training_loss=df_loss["training_loss"].values
42
- )
43
-
44
- def plot_loss(self, smooth_fraction: float = 0.05) -> plt.Figure:
45
- fig = plot_losses(dimension=0,
46
- models=[self],
47
- title=self.name,
48
- smooth_fraction=smooth_fraction)
49
- plt.close(fig)
50
- return fig
51
-
52
22
  @dataclass(frozen=True, kw_only=True)
53
23
  class ModelConfig(ABC, Summary):
54
24
  """
@@ -71,7 +41,7 @@ class ModelConfig(ABC, Summary):
71
41
  ...
72
42
 
73
43
 
74
- class Model(ABC, ModelRecord):
44
+ class Model(ABC):
75
45
  """
76
46
  Base class orchestrating the training and testing loop.
77
47
 
@@ -124,6 +94,7 @@ class Model(ABC, ModelRecord):
124
94
  print(f" epoch n°{epoch_idx + 1} out of {self.nb_epochs}")
125
95
  self.data.permutation()
126
96
  self.data.batch_size = self.batch_size
97
+
127
98
  for batch_idx, (X_train, Y_train) in enumerate(self.data):
128
99
  avg_loss = self.optimizer.do_descent(self.neural_network, self.loss, X_train, Y_train)
129
100
  self.training_loss[epoch_idx * nb_batches + batch_idx] = avg_loss
@@ -131,8 +102,17 @@ class Model(ABC, ModelRecord):
131
102
  if batch_idx in print_indexes:
132
103
  print(f" batch n°{batch_idx + 1} out of {nb_batches}, "
133
104
  f"loss : {np.round(self.training_loss[epoch_idx * nb_batches + batch_idx], 3)}")
105
+
134
106
  self.test()
135
107
 
108
+ def plot_loss(self, smooth_fraction: float = 0.05) -> plt.Figure:
109
+ fig = plot_losses(dimension=0,
110
+ models=[self],
111
+ title=self.name,
112
+ smooth_fraction=smooth_fraction)
113
+ plt.close(fig)
114
+ return fig
115
+
136
116
  def test(self) -> None:
137
117
  X_test, Y_true = self.data.X_test, self.data.Y_test # (in, batch), (out, batch)
138
118
  Y_pred = self.neural_network(X_test) # (out, batch)
@@ -142,24 +122,27 @@ class Model(ABC, ModelRecord):
142
122
 
143
123
  print(f" {self.id} accuracy : {self.test_accuracy}, loss : {self.test_loss}")
144
124
 
145
- def save_loss(self, save_dir: str) -> None:
146
- save_path = os.path.join(save_dir, self.LOSS_FILE)
147
- os.makedirs(os.path.dirname(save_dir), exist_ok=True)
125
+ def save_loss(self, save_dir: Path) -> None:
126
+ save_dir.mkdir(parents=True, exist_ok=True)
127
+ save_path = save_dir / self.LOSS_FILE
128
+
148
129
  df = pd.DataFrame({
149
130
  'training_loss': self.training_loss
150
131
  })
151
132
  df.to_csv(save_path)
152
133
 
153
- def save_weights(self, save_dir: str) -> None:
154
- save_path = os.path.join(save_dir, self.WEIGHTS_FILE)
134
+ def save_weights(self, save_dir: Path) -> None:
135
+ save_dir.mkdir(parents=True, exist_ok=True)
136
+ save_path = save_dir / self.WEIGHTS_FILE
137
+
155
138
  params_dict = self.neural_network.get_params()
156
- with open(save_path, 'wb') as f:
139
+ with save_path.open('wb') as f:
157
140
  pickle.dump(params_dict, f)
158
141
 
159
- def load_weights(self, load_dir: str) -> None:
142
+ def load_weights(self, load_dir: Path) -> None:
160
143
  """Restaure les paramètres depuis un fichier pickle."""
161
- load_path = os.path.join(load_dir, self.WEIGHTS_FILE)
162
- with open(load_path, 'rb') as f:
144
+ load_path = load_dir / self.WEIGHTS_FILE
145
+ with load_path.open('rb') as f:
163
146
  params_dict = pickle.load(f)
164
147
 
165
148
  self.neural_network.init_params(params_dict)
@@ -1,7 +1,7 @@
1
1
  from __future__ import annotations
2
2
 
3
3
  import dataclasses
4
- import os
4
+ from pathlib import Path
5
5
  from typing import Any
6
6
 
7
7
 
@@ -9,14 +9,10 @@ class Summary:
9
9
  def summary(self, file=None) -> None:
10
10
  print(self._summary(self, indent=0), file=file)
11
11
 
12
- def save(self, save_dir: str) -> None:
12
+ def save(self, path: Path) -> None:
13
+ path.parent.mkdir(parents=True, exist_ok=True)
13
14
 
14
- save_path = os.path.join(save_dir, "config.txt")
15
-
16
- if os.path.dirname(save_path):
17
- os.makedirs(os.path.dirname(save_path), exist_ok=True)
18
-
19
- with open(save_path, "w", encoding="utf-8") as f:
15
+ with path.open("w", encoding="utf-8") as f:
20
16
  print(self._summary(self, indent=0), file=f)
21
17
 
22
18
  @classmethod
@@ -1,7 +1,7 @@
1
1
  from __future__ import annotations
2
2
 
3
3
  import itertools
4
- import os
4
+ from pathlib import Path
5
5
  from dataclasses import dataclass, replace
6
6
  from typing import Union
7
7
 
@@ -28,7 +28,6 @@ class ExperimentConfig(Summary):
28
28
  data_creator: DataCreator
29
29
  variations: list[VariationConfig]
30
30
 
31
-
32
31
  def instantiate(self) -> Experiment:
33
32
  return Experiment(self)
34
33
 
@@ -74,31 +73,32 @@ class Experiment:
74
73
  for model in self.models:
75
74
  model.test()
76
75
 
77
- def save_df(self, save_dir: str) -> None:
76
+ def save_df(self, save_dir: Path) -> None:
78
77
  """
79
78
  saves the models parameters and their args
80
79
  """
81
- os.makedirs(save_dir, exist_ok=True)
80
+ save_dir.mkdir(parents=True, exist_ok=True)
81
+ save_path = save_dir / self.ACCURACY_FILE
82
82
  print(f" Saving results to: {save_dir}")
83
83
 
84
84
  data = [model.id | {"test_loss": model.test_loss, "test_accuracy": model.test_accuracy}
85
85
  for model in self.models]
86
86
 
87
87
  df = pd.DataFrame(data)
88
- df.to_csv(os.path.join(save_dir, self.ACCURACY_FILE), index_label="iteration")
88
+ df.to_csv(save_path)
89
89
 
90
- def save_weights(self, save_dir: str) -> None:
90
+ def save_weights(self, save_dir: Path) -> None:
91
91
  for i, model in enumerate(self.models):
92
- save_path = os.path.join(save_dir, model.name)
92
+ save_path = save_dir / model.name
93
93
  model.save_weights(save_path)
94
94
 
95
- def save_configs(self, save_dir: str) -> None:
95
+ def save_configs(self, save_dir: Path) -> None:
96
96
 
97
- config_path = os.path.join(save_dir, self.CONFIG_FILE)
97
+ config_path = save_dir / self.CONFIG_FILE
98
98
  self.config.save(config_path)
99
99
 
100
100
  for i, model in enumerate(self.models):
101
- save_path = os.path.join(save_dir, model.name)
101
+ save_path = save_dir / model.name
102
102
  model.config.save(save_path)
103
103
 
104
104
 
@@ -11,7 +11,7 @@ from matplotlib.axes import Axes
11
11
  from .types import Array
12
12
 
13
13
  if TYPE_CHECKING:
14
- from .abstract.model import ModelRecord
14
+ from .abstract.model import Model
15
15
 
16
16
 
17
17
  def set_style() -> None:
@@ -80,7 +80,7 @@ def smooth_curve(loss: Array, window_length: int) -> Array:
80
80
  return np.exp(pd.Series(np.log(loss)).ewm(span=window_length, adjust=True).mean())
81
81
 
82
82
 
83
- def plot_0d(models: list[ModelRecord], title: str, smooth_fraction: float = 50) -> plt.Figure:
83
+ def plot_0d(models: list[Model], title: str, smooth_fraction: float = 50) -> plt.Figure:
84
84
  """
85
85
  Plots a single graph overlaying multiple models that share the same hyperparameters.
86
86
  """
@@ -108,7 +108,7 @@ def plot_0d(models: list[ModelRecord], title: str, smooth_fraction: float = 50)
108
108
  return fig
109
109
 
110
110
 
111
- def plot_1d(models: list[ModelRecord], title: str, key: str, smooth_fraction: float = 50) -> plt.Figure:
111
+ def plot_1d(models: list[Model], title: str, key: str, smooth_fraction: float = 50) -> plt.Figure:
112
112
  """
113
113
  Plots a row of subplots, varying one hyperparameter (key) across columns.
114
114
  """
@@ -147,7 +147,7 @@ def plot_1d(models: list[ModelRecord], title: str, key: str, smooth_fraction: fl
147
147
  return fig
148
148
 
149
149
 
150
- def plot_2d(models: list[ModelRecord], title: str, row_key: str, col_key: str, smooth_fraction: float) -> plt.Figure:
150
+ def plot_2d(models: list[Model], title: str, row_key: str, col_key: str, smooth_fraction: float) -> plt.Figure:
151
151
  """
152
152
  Plots a grid of subplots varying two hyperparameters: one across rows, one across columns.
153
153
 
@@ -199,7 +199,7 @@ def plot_2d(models: list[ModelRecord], title: str, row_key: str, col_key: str, s
199
199
  return fig
200
200
 
201
201
 
202
- def plot_losses(title: str, dimension: int, models: list[ModelRecord], smooth_fraction: float) -> plt.Figure:
202
+ def plot_losses(title: str, dimension: int, models: list[Model], smooth_fraction: float) -> plt.Figure:
203
203
  """
204
204
  Main entry point for plotting. Automatically detects if the plot should be 0D, 1D, or 2D
205
205
  based on the number of variation parameters.
@@ -27,7 +27,7 @@ class GlobalFiniteDifferenceConfig(GradientEstimatorConfig):
27
27
  dA: float
28
28
 
29
29
  def instantiate(self, nb_params) -> GlobalFiniteDifference:
30
- return GlobalFiniteDifference(self, nb_params)
30
+ return GlobalFiniteDifference(self.dA, nb_params)
31
31
 
32
32
 
33
33
  @dataclass(frozen=True)
@@ -36,7 +36,7 @@ class PartialFiniteDifferenceConfig(GradientEstimatorConfig):
36
36
  indexes: list[int]
37
37
 
38
38
  def instantiate(self, nb_params) -> PartialFiniteDifference:
39
- return PartialFiniteDifference(self, nb_params)
39
+ return PartialFiniteDifference(self.dA, self.indexes, nb_params)
40
40
 
41
41
 
42
42
  @dataclass(frozen=True)
@@ -46,7 +46,7 @@ class SimultaneousPerturbationConfig(GradientEstimatorConfig):
46
46
  get_perturbation_matrix: PerturbationMatrix
47
47
 
48
48
  def instantiate(self, nb_params: int) -> SimultaneousPerturbation:
49
- return SimultaneousPerturbation(self, nb_params)
49
+ return SimultaneousPerturbation(self.dA, self.nb_perturbations, self.get_perturbation_matrix, nb_params)
50
50
 
51
51
 
52
52
  class GradientEstimator(ABC):
@@ -89,12 +89,12 @@ class NullGradientEstimator(GradientEstimator):
89
89
 
90
90
 
91
91
  class GlobalFiniteDifference(GradientEstimator):
92
- def __init__(self, config: GlobalFiniteDifferenceConfig, nb_params: int) -> None:
92
+ def __init__(self, dA: float, nb_params: int) -> None:
93
93
  self.nb_params: int = nb_params
94
- self.dA: float = config.dA
94
+ self.dA: float = dA
95
95
 
96
96
  self.perturbation_matrix: Array = np.vstack((np.zeros((1, self.nb_params)), np.eye(nb_params)))
97
- self.Ps: Array = config.dA * self.perturbation_matrix
97
+ self.Ps: Array = dA * self.perturbation_matrix
98
98
 
99
99
  def perturb(self, Theta: Array) -> Array:
100
100
  return Theta + self.Ps
@@ -105,16 +105,16 @@ class GlobalFiniteDifference(GradientEstimator):
105
105
 
106
106
 
107
107
  class PartialFiniteDifference(GradientEstimator):
108
- def __init__(self, config: PartialFiniteDifferenceConfig, nb_params: int) -> None:
108
+ def __init__(self, dA: float, indexes: list[int], nb_params: int) -> None:
109
109
  self.nb_params: int = nb_params
110
- self.dA: float = config.dA
111
- self.indexes = config.indexes
112
- self.nb_perturbations = len(config.indexes)
110
+ self.dA: float = dA
111
+ self.indexes: list[int] = indexes
112
+ self.nb_perturbations: int = len(indexes)
113
113
 
114
114
  self.perturbation_matrix: Array = np.zeros((self.nb_perturbations + 1, self.nb_params))
115
115
  self.perturbation_matrix[range(1, self.nb_perturbations + 1), self.indexes] = 1
116
116
 
117
- self.Ps: Array = config.dA * self.perturbation_matrix
117
+ self.Ps: Array = dA * self.perturbation_matrix
118
118
 
119
119
  def perturb(self, Theta: Array) -> Array:
120
120
  return Theta + self.Ps
@@ -125,11 +125,11 @@ class PartialFiniteDifference(GradientEstimator):
125
125
 
126
126
 
127
127
  class SimultaneousPerturbation(GradientEstimator):
128
- def __init__(self, config: SimultaneousPerturbationConfig, nb_params: int) -> None:
128
+ def __init__(self, dA: float, nb_perturbations: int, get_perturbation_matrix: PerturbationMatrix, nb_params: int) -> None:
129
129
  self.nb_params: int = nb_params
130
- self.dA: float = config.dA
131
- self.nb_perturbations: int = config.nb_perturbations
132
- self.get_perturbation_matrix: PerturbationMatrix = config.get_perturbation_matrix
130
+ self.dA: float = dA
131
+ self.nb_perturbations: int = nb_perturbations
132
+ self.get_perturbation_matrix: PerturbationMatrix = get_perturbation_matrix
133
133
 
134
134
  nb_copies = 3
135
135
  self.Ps_extended: Array = np.vstack((np.zeros((1, self.nb_params * nb_copies)),
@@ -24,7 +24,7 @@ class ZerothOrderSGDConfig(ZerothOrderOptimizerConfig):
24
24
  learning_rate: float
25
25
 
26
26
  def instantiate(self, gradient_estimator: GradientEstimator) -> ZerothOrderSGD:
27
- return ZerothOrderSGD(self, gradient_estimator)
27
+ return ZerothOrderSGD(self.learning_rate, gradient_estimator)
28
28
 
29
29
 
30
30
  @dataclass(frozen=True)
@@ -35,7 +35,7 @@ class ZerothOrderAdamConfig(ZerothOrderSGDConfig):
35
35
  epsilon: float
36
36
 
37
37
  def instantiate(self, gradient_estimator: GradientEstimator) -> ZerothOrderAdam:
38
- return ZerothOrderAdam(self, gradient_estimator)
38
+ return ZerothOrderAdam(self.learning_rate, self.beta1, self.beta2, self.epsilon, gradient_estimator)
39
39
 
40
40
 
41
41
  class ZerothOrderOptimizer(Optimizer):
@@ -60,9 +60,9 @@ class ZerothOrderSGD(ZerothOrderOptimizer):
60
60
  the gradient by evaluating the loss on perturbed versions of the parameters.
61
61
  """
62
62
 
63
- def __init__(self, config: ZerothOrderSGDConfig, gradient_estimator: GradientEstimator) -> None:
64
- self.learning_rate = config.learning_rate
65
- self.gradient_estimator = gradient_estimator
63
+ def __init__(self, learning_rate: float, gradient_estimator: GradientEstimator) -> None:
64
+ self.learning_rate: float = learning_rate
65
+ self.gradient_estimator: GradientEstimator = gradient_estimator
66
66
 
67
67
  def do_descent(self, blackbox: ZerothOrderBlackBox, loss: Loss, X: Array, Y_true: Array) -> float:
68
68
  """Performs one optimization step using zeroth_order.
@@ -104,16 +104,16 @@ class ZerothOrderAdam(ZerothOrderSGD):
104
104
  as its momentum terms (m, v) help smooth out the noise over time.
105
105
  """
106
106
  name = "Adam"
107
- def __init__(self, config: ZerothOrderAdamConfig, gradient_estimator: GradientEstimator) -> None:
108
- self.beta1: float = config.beta1
109
- self.beta2: float = config.beta2
110
- self.epsilon: float = config.epsilon
111
- self.beta1t: float = config.beta1
112
- self.beta2t: float = config.beta2
107
+ def __init__(self, learning_rate: float, beta1: float, beta2: float, epsilon: float, gradient_estimator: GradientEstimator) -> None:
108
+ self.beta1: float = beta1
109
+ self.beta2: float = beta2
110
+ self.epsilon: float = epsilon
111
+ self.beta1t: float = beta1
112
+ self.beta2t: float = beta2
113
113
  self.m: Array = np.array([0])
114
114
  self.v: Array = np.array([0])
115
115
 
116
- super().__init__(config, gradient_estimator)
116
+ super().__init__(learning_rate, gradient_estimator)
117
117
 
118
118
  def _apply_update_rule(self, grad: Array) -> Array:
119
119
  self.m = self.beta1 * self.m + (1 - self.beta1) * grad
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: zeroth-learn
3
- Version: 0.2.3
3
+ Version: 0.2.4
4
4
  Requires-Dist: numpy
5
5
  Requires-Dist: pandas
6
6
  Requires-Dist: matplotlib
File without changes
File without changes
File without changes
File without changes