Skip to content

Commit 2fa8de7

Browse files
authored
Merge pull request #87 from INGEOTEC/develop
Version - 1.5.6
2 parents 71fb982 + 8faa023 commit 2fa8de7

3 files changed

Lines changed: 111 additions & 39 deletions

File tree

EvoMSA/__init__.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -11,7 +11,7 @@
1111
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
1212
# See the License for the specific language governing permissions and
1313
# limitations under the License.
14-
__version__ = '1.5.5'
14+
__version__ = '1.5.6'
1515

1616
try:
1717
from EvoMSA.evodag import BoW, TextRepresentations, StackGeneralization

EvoMSA/evodag.py

Lines changed: 93 additions & 37 deletions
Original file line numberDiff line numberDiff line change
@@ -54,18 +54,42 @@ def __init__(self, lang: str='es',
5454
assert lang is None or lang in MODEL_LANG
5555
self._n_jobs = n_jobs
5656
self._lang = lang
57-
self._key = key
58-
self._label_key = label_key
57+
self.key = key
58+
self.label_key = label_key
5959
self._mixer_func = mixer_func
60-
self._decision_function = decision_function
61-
self._estimator_class = estimator_class
62-
self._estimator_kwargs = estimator_kwargs
60+
self.decision_function_name = decision_function
61+
self.estimator_class = estimator_class
62+
self.estimator_kwargs = estimator_kwargs
6363
self._b4msa_kwargs = b4msa_kwargs
6464
self._pretrain = pretrain
65-
self._kfold_instance = kfold_instance
66-
self._kfold_kwargs = kfold_kwargs
65+
self.kfold_instance = kfold_instance
66+
self.kfold_kwargs = kfold_kwargs
6767
self._b4msa_estimated = False
6868

69+
@property
70+
def label_key(self):
71+
return self._label_key
72+
73+
@label_key.setter
74+
def label_key(self, value):
75+
self._label_key = value
76+
77+
@property
78+
def key(self):
79+
return self._key
80+
81+
@key.setter
82+
def key(self, value):
83+
self._key = value
84+
85+
@property
86+
def decision_function_name(self):
87+
return self._decision_function
88+
89+
@decision_function_name.setter
90+
def decision_function_name(self, value):
91+
self._decision_function = value
92+
6993
@property
7094
def names(self):
7195
_names = [None] * len(self.bow.id2token)
@@ -81,31 +105,21 @@ def pretrain(self):
81105
def lang(self):
82106
return self._lang
83107

84-
def b4msa_fit(self, D):
85-
assert len(D)
86-
self._b4msa_estimated = True
87-
if self._key == 'text' or isinstance(D[0], str):
88-
return self.bow.fit(D)
89-
assert isinstance(D[0], dict)
90-
if isinstance(self._key, str):
91-
key = self._key
92-
return self.bow.fit([x[key] for x in D])
93-
_ = [[x[key] for key in self._key] for x in D]
94-
return self.bow.fit(_)
108+
@property
109+
def kfold_instance(self):
110+
return self._kfold_instance
95111

96-
def transform(self, D: List[Union[List, dict]], y=None) -> csr_matrix:
97-
assert len(D)
98-
if not self.pretrain:
99-
assert self._b4msa_estimated
100-
if self._key == 'text' or isinstance(D[0], str):
101-
return self.bow.transform(D)
102-
assert isinstance(D[0], dict)
103-
if isinstance(self._key, str):
104-
key = self._key
105-
return self.bow.transform([x[key] for x in D])
106-
Xs = [self.bow.transform([x[key] for x in D])
107-
for key in self._key]
108-
return self._mixer_func(Xs)
112+
@kfold_instance.setter
113+
def kfold_instance(self, value):
114+
self._kfold_instance = value
115+
116+
@property
117+
def kfold_kwargs(self):
118+
return self._kfold_kwargs
119+
120+
@kfold_kwargs.setter
121+
def kfold_kwargs(self, value):
122+
self._kfold_kwargs = value
109123

110124
@property
111125
def bow(self):
@@ -124,10 +138,52 @@ def bow(self):
124138
def bow(self, value):
125139
self._bow = value
126140

141+
@property
142+
def estimator_class(self):
143+
return self._estimator_class
144+
145+
@estimator_class.setter
146+
def estimator_class(self, value):
147+
self._estimator_class = value
148+
149+
@property
150+
def estimator_kwargs(self):
151+
return self._estimator_kwargs
152+
153+
@estimator_kwargs.setter
154+
def estimator_kwargs(self, value):
155+
self._estimator_kwargs = value
156+
157+
def b4msa_fit(self, D):
158+
assert len(D)
159+
self._b4msa_estimated = True
160+
if self.key == 'text' or isinstance(D[0], str):
161+
return self.bow.fit(D)
162+
assert isinstance(D[0], dict)
163+
if isinstance(self.key, str):
164+
key = self.key
165+
return self.bow.fit([x[key] for x in D])
166+
_ = [[x[key] for key in self.key] for x in D]
167+
return self.bow.fit(_)
168+
169+
def transform(self, D: List[Union[List, dict]], y=None) -> csr_matrix:
170+
assert len(D)
171+
if not self.pretrain:
172+
assert self._b4msa_estimated
173+
if self.key == 'text' or isinstance(D[0], str):
174+
return self.bow.transform(D)
175+
assert isinstance(D[0], dict)
176+
if isinstance(self.key, str):
177+
key = self.key
178+
return self.bow.transform([x[key] for x in D])
179+
Xs = [self.bow.transform([x[key] for x in D])
180+
for key in self.key]
181+
return self._mixer_func(Xs)
182+
127183
def dependent_variable(self, D: List[Union[dict, list]],
128184
y: Union[np.ndarray, None]=None) -> np.ndarray:
129185
assert isinstance(D, list) and len(D)
130-
label_key = self._label_key
186+
label_key = self.label_key
131187
if y is None:
132188
assert isinstance(D[0], dict)
133189
y = np.array([x[label_key] for x in D])
@@ -149,10 +205,10 @@ def train_predict_decision_function(self, D: List[Union[dict, list]],
149205
y: Union[np.ndarray, None]=None) -> np.ndarray:
150206
def train_predict(tr, vs):
151207
m = self.estimator().fit(X[tr], y[tr])
152-
return getattr(m, self._decision_function)(X[vs])
208+
return getattr(m, self.decision_function_name)(X[vs])
153209

154210
y = self.dependent_variable(D, y=y)
155-
kf = self._kfold_instance(**self._kfold_kwargs)
211+
kf = self.kfold_instance(**self.kfold_kwargs)
156212
kfolds = [x for x in kf.split(D, y)]
157213
X = self.transform(D, y=y)
158214
hys = Parallel(n_jobs=self._n_jobs)(delayed(train_predict)(tr, vs)
@@ -182,7 +238,7 @@ def predict(self, D: List[Union[dict, list]]) -> np.ndarray:
182238

183239
def decision_function(self, D: List[Union[dict, list]]) -> Union[list, np.ndarray]:
184240
_ = self.transform(D)
185-
hy = getattr(self.estimator_instance, self._decision_function)(_)
241+
hy = getattr(self.estimator_instance, self.decision_function_name)(_)
186242
if hy.ndim == 1:
187243
return np.atleast_2d(hy).T
188244
return hy
@@ -273,7 +329,7 @@ def load_dataset(self) -> None:
273329
for k, name in zip(_, names)]
274330

275331
def transform(self, D: List[Union[List, dict]], y=None) -> np.ndarray:
276-
if isinstance(self._key, str):
332+
if isinstance(self.key, str):
277333
X = super(TextRepresentations, self).transform(D, y=y)
278334
models = Parallel(n_jobs=self._n_jobs)(delayed(m.decision_function)(X)
279335
for m in self.text_representations)
@@ -284,7 +340,7 @@ def transform(self, D: List[Union[List, dict]], y=None) -> np.ndarray:
284340
return _
285341
assert len(D) and isinstance(D[0], dict)
286342
Xs = [super(TextRepresentations, self).transform([x[key] for x in D], y=y)
287-
for key in self._key]
343+
for key in self.key]
288344
with Parallel(n_jobs=self._n_jobs) as parallel:
289345
models = []
290346
for X in Xs:

EvoMSA/tests/test_evodag.py

Lines changed: 17 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -250,4 +250,20 @@ def test_TextRepresentations_unit():
250250
unit_vector=True)
251251
X = text_repr.transform([dict(text='buenos días')])
252252
_ = np.sqrt((X ** 2).sum(axis=1))
253-
np.testing.assert_almost_equal(_, 1)
253+
np.testing.assert_almost_equal(_, 1)
254+
255+
256+
def test_BoW_property():
257+
from EvoMSA.evodag import BoW
258+
bow = BoW()
259+
bow.kfold_instance = '!'
260+
bow.kfold_kwargs = '*'
261+
assert bow._kfold_instance == '!' and bow._kfold_kwargs == '*'
262+
bow.estimator_class = '1'
263+
bow.estimator_kwargs = '2'
264+
assert bow._estimator_class == '1' and bow._estimator_kwargs == '2'
265+
bow.decision_function_name = '3'
266+
assert bow._decision_function == '3'
267+
bow.key = '4'
268+
bow.label_key = '5'
269+
assert bow._key == '4' and bow._label_key == '5'

0 commit comments

Comments
 (0)