@@ -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 :
0 commit comments