1+ import unittest
2+ import superstyl
3+ import pandas
4+ import numpy as np
5+ from collections import Counter
6+
7+
8+ class TestSMOTEFixes (unittest .TestCase ):
9+ """Tests pour les corrections SMOTE avec validation croisée"""
10+
11+ def setUp (self ):
12+ """Créer des datasets de test"""
13+ # Dataset équilibré basique
14+ self .balanced_data = pandas .DataFrame ({
15+ 'author' : ['A' , 'A' , 'A' , 'A' , 'B' , 'B' , 'B' , 'B' ],
16+ 'lang' : ['en' ] * 8 ,
17+ 'feat1' : [1 , 2 , 3 , 4 , 5 , 6 , 7 , 8 ],
18+ 'feat2' : [8 , 7 , 6 , 5 , 4 , 3 , 2 , 1 ]
19+ })
20+
21+ # Dataset déséquilibré avec petites classes
22+ self .imbalanced_small = pandas .DataFrame ({
23+ 'author' : ['Majority' ] * 10 + ['Minority' ] * 3 ,
24+ 'lang' : ['en' ] * 13 ,
25+ 'feat1' : np .random .rand (13 ),
26+ 'feat2' : np .random .rand (13 ),
27+ 'feat3' : np .random .rand (13 )
28+ })
29+
30+ # Dataset pour group-k-fold avec structure de groupes
31+ # Format: Author_Work_Sample
32+ indices = (
33+ ['Smith_Work1_0-100' , 'Smith_Work1_100-200' , 'Smith_Work1_200-300' ] +
34+ ['Smith_Work2_0-100' , 'Smith_Work2_100-200' ] +
35+ ['Dupont_Work1_0-100' , 'Dupont_Work1_100-200' ] +
36+ ['Dupont_Work2_0-100' , 'Dupont_Work2_100-200' , 'Dupont_Work2_200-300' ]
37+ )
38+ self .grouped_data = pandas .DataFrame ({
39+ 'author' : ['Smith' ] * 5 + ['Dupont' ] * 5 ,
40+ 'lang' : ['en' ] * 10 ,
41+ 'feat1' : np .random .rand (10 ),
42+ 'feat2' : np .random .rand (10 )
43+ }, index = indices )
44+
45+ def test_smote_with_kfold (self ):
46+ """Test SMOTE avec k-fold classique"""
47+ # GIVEN: Un dataset déséquilibré
48+ # WHEN: On applique SMOTE avec k-fold
49+ try :
50+ results = superstyl .train_svm (
51+ self .imbalanced_small ,
52+ cross_validate = 'k-fold' ,
53+ k = 3 ,
54+ balance = 'SMOTE' ,
55+ norms = True
56+ )
57+ # THEN: Ça devrait fonctionner
58+ self .assertIn ('confusion_matrix' , results )
59+ self .assertIn ('pipeline' , results )
60+ except ValueError as e :
61+ if 'n_neighbors' in str (e ):
62+ self .fail (f"SMOTE k-fold failed with: { e } " )
63+
64+ def test_smote_with_leave_one_out (self ):
65+ """Test SMOTE avec leave-one-out"""
66+ # GIVEN: Un petit dataset
67+ small_data = pandas .DataFrame ({
68+ 'author' : ['A' ] * 4 + ['B' ] * 4 ,
69+ 'lang' : ['en' ] * 8 ,
70+ 'feat1' : [1 , 2 , 3 , 4 , 5 , 6 , 7 , 8 ],
71+ 'feat2' : [8 , 7 , 6 , 5 , 4 , 3 , 2 , 1 ]
72+ })
73+
74+ # WHEN: On applique SMOTE avec LOO
75+ try :
76+ results = superstyl .train_svm (
77+ small_data ,
78+ cross_validate = 'leave-one-out' ,
79+ balance = 'SMOTE' ,
80+ norms = True
81+ )
82+ # THEN: Ça devrait fonctionner
83+ self .assertIn ('confusion_matrix' , results )
84+ except ValueError as e :
85+ if 'n_neighbors' in str (e ):
86+ self .fail (f"SMOTE LOO failed with: { e } " )
87+
88+ def test_smote_with_group_kfold (self ):
89+ """Test SMOTE avec group-k-fold (le cas problématique)"""
90+ # GIVEN: Un dataset avec structure de groupes
91+ # WHEN: On applique SMOTE avec group-k-fold
92+ try :
93+ results = superstyl .train_svm (
94+ self .grouped_data ,
95+ cross_validate = 'group-k-fold' ,
96+ k = 2 ,
97+ balance = 'SMOTE' ,
98+ norms = True
99+ )
100+ # THEN: Ça devrait fonctionner ou skip SMOTE proprement
101+ self .assertIn ('pipeline' , results )
102+ # Si SMOTE a été skippé, le pipeline ne devrait pas contenir de sampling
103+ pipeline_steps = [step [0 ] for step in results ['pipeline' ].steps ]
104+ # C'est OK si 'sampling' est présent ou absent (selon la taille des classes)
105+ except ValueError as e :
106+ if 'n_neighbors' in str (e ):
107+ self .fail (f"SMOTE group-k-fold failed with: { e } " )
108+
109+ def test_smote_very_small_classes (self ):
110+ """Test SMOTE avec des classes extrêmement petites (devrait skip SMOTE)"""
111+ # GIVEN: Dataset avec classe de taille 2
112+ tiny_data = pandas .DataFrame ({
113+ 'author' : ['A' , 'A' , 'B' ] * 2 ,
114+ 'lang' : ['en' ] * 6 ,
115+ 'feat1' : [1 , 2 , 3 , 4 , 5 , 6 ],
116+ 'feat2' : [6 , 5 , 4 , 3 , 2 , 1 ]
117+ })
118+
119+ # WHEN: On essaie SMOTE avec k-fold
120+ try :
121+ results = superstyl .train_svm (
122+ tiny_data ,
123+ cross_validate = 'k-fold' ,
124+ k = 3 ,
125+ balance = 'SMOTE' ,
126+ norms = True
127+ )
128+ # THEN: Devrait fonctionner (SMOTE skippé ou n_neighbors=1)
129+ self .assertIn ('pipeline' , results )
130+ except ValueError as e :
131+ if 'n_neighbors' in str (e ):
132+ self .fail (f"Should have skipped SMOTE or used n_neighbors=1, but got: { e } " )
133+
134+ def test_smotetomek_with_cv (self ):
135+ """Test SMOTETomek avec validation croisée"""
136+ # WHEN: On applique SMOTETomek
137+ try :
138+ results = superstyl .train_svm (
139+ self .imbalanced_small ,
140+ cross_validate = 'k-fold' ,
141+ k = 3 ,
142+ balance = 'SMOTETomek' ,
143+ norms = True
144+ )
145+ # THEN: Ça devrait fonctionner
146+ self .assertIn ('confusion_matrix' , results )
147+ except ValueError as e :
148+ if 'n_neighbors' in str (e ):
149+ self .fail (f"SMOTETomek failed with: { e } " )
150+
151+ def test_alternative_balance_methods (self ):
152+ """Test que les méthodes alternatives fonctionnent"""
153+ # Test upsampling
154+ results1 = superstyl .train_svm (
155+ self .imbalanced_small ,
156+ cross_validate = 'k-fold' ,
157+ k = 3 ,
158+ balance = 'upsampling' ,
159+ norms = True
160+ )
161+ self .assertIn ('confusion_matrix' , results1 )
162+
163+ # Test downsampling
164+ results2 = superstyl .train_svm (
165+ self .imbalanced_small ,
166+ cross_validate = 'k-fold' ,
167+ k = 3 ,
168+ balance = 'downsampling' ,
169+ norms = True
170+ )
171+ self .assertIn ('confusion_matrix' , results2 )
172+
173+ # Test class_weights sans balance
174+ results3 = superstyl .train_svm (
175+ self .imbalanced_small ,
176+ cross_validate = 'k-fold' ,
177+ k = 3 ,
178+ balance = None ,
179+ class_weights = True ,
180+ norms = True
181+ )
182+ self .assertIn ('confusion_matrix' , results3 )
183+
184+ def test_smote_with_group_kfold_bad_distribution (self ):
185+ """Test que group-k-fold échoue proprement avec une mauvaise distribution de classes"""
186+ # GIVEN: Dataset où les groupes sont mal distribués (une classe par groupe)
187+ indices = (
188+ ['Author1_Work1_0' , 'Author1_Work1_1' , 'Author1_Work1_2' ] +
189+ ['Author1_Work2_0' , 'Author1_Work2_1' ] +
190+ ['Author2_Work1_0' , 'Author2_Work1_1' , 'Author2_Work1_2' ]
191+ )
192+ bad_grouped = pandas .DataFrame ({
193+ 'author' : ['Author1' ] * 5 + ['Author2' ] * 3 ,
194+ 'lang' : ['en' ] * 8 ,
195+ 'feat1' : np .random .rand (8 ),
196+ 'feat2' : np .random .rand (8 )
197+ }, index = indices )
198+
199+ # WHEN: On essaie SMOTE avec group-k-fold
200+ # THEN: Devrait soit skip SMOTE, soit échouer avec un message clair
201+ try :
202+ results = superstyl .train_svm (
203+ bad_grouped ,
204+ cross_validate = 'group-k-fold' ,
205+ k = 2 ,
206+ balance = 'SMOTE' ,
207+ norms = True
208+ )
209+ # Si ça marche, c'est que SMOTE a été skippé intelligemment
210+ self .assertIn ('pipeline' , results )
211+ except ValueError as e :
212+ # L'erreur "Got 1 class instead" est acceptable ici
213+ # C'est un problème inhérent aux données, pas au code
214+ if "Got 1 class instead" not in str (e ):
215+ self .fail (f"Expected 'Got 1 class' error or success, got: { e } " )
216+
217+
218+
219+
220+ class TestFinalPredFixes (unittest .TestCase ):
221+ """Tests pour la correction final_pred avec test set"""
222+
223+ def setUp (self ):
224+ """Créer des datasets train et test"""
225+ self .train = pandas .DataFrame ({
226+ 'author' : ['A' ] * 5 + ['B' ] * 5 ,
227+ 'lang' : ['en' ] * 10 ,
228+ 'feat1' : np .random .rand (10 ),
229+ 'feat2' : np .random .rand (10 ),
230+ 'feat3' : np .random .rand (10 )
231+ })
232+
233+ self .test = pandas .DataFrame ({
234+ 'author' : ['A' ] * 2 + ['B' ] * 2 ,
235+ 'lang' : ['en' ] * 4 ,
236+ 'feat1' : np .random .rand (4 ),
237+ 'feat2' : np .random .rand (4 ),
238+ 'feat3' : np .random .rand (4 )
239+ })
240+
241+ def test_final_pred_with_test_set (self ):
242+ """Test final_pred=True AVEC test set (devrait fonctionner)"""
243+ # WHEN: On fait une prédiction finale avec test set
244+ results = superstyl .train_svm (
245+ self .train ,
246+ test = self .test ,
247+ cross_validate = 'k-fold' ,
248+ k = 3 ,
249+ final_pred = True ,
250+ norms = True
251+ )
252+
253+ # THEN: Devrait contenir les prédictions finales
254+ self .assertIn ('final_predictions' , results )
255+ self .assertIn ('pipeline' , results )
256+ self .assertEqual (len (results ['final_predictions' ]), len (self .test ))
257+
258+ def test_final_pred_without_test_set (self ):
259+ """Test final_pred=True SANS test set (devrait lever une erreur)"""
260+ # WHEN: On essaie final_pred sans test set
261+ with self .assertRaises (ValueError ) as context :
262+ superstyl .train_svm (
263+ self .train ,
264+ test = None , # ← Pas de test set
265+ cross_validate = 'k-fold' ,
266+ k = 3 ,
267+ final_pred = True ,
268+ norms = True
269+ )
270+
271+ # THEN: Devrait avoir un message d'erreur clair
272+ self .assertIn ('test' , str (context .exception ).lower ())
273+
274+ def test_cv_without_final_pred (self ):
275+ """Test validation croisée sans final_pred (devrait fonctionner)"""
276+ # WHEN: On fait juste de la CV sans prédiction finale
277+ results = superstyl .train_svm (
278+ self .train ,
279+ cross_validate = 'k-fold' ,
280+ k = 3 ,
281+ final_pred = False ,
282+ norms = True
283+ )
284+
285+ # THEN: Devrait avoir confusion matrix mais pas final_predictions
286+ self .assertIn ('confusion_matrix' , results )
287+ self .assertNotIn ('final_predictions' , results )
288+
289+ def test_train_test_split_with_final_pred (self ):
290+ """Test train/test split classique avec final_pred"""
291+ # WHEN: On fait un train/test split (pas de CV)
292+ results = superstyl .train_svm (
293+ self .train ,
294+ test = self .test ,
295+ cross_validate = None , # ← Pas de CV
296+ final_pred = True ,
297+ norms = True
298+ )
299+
300+ # THEN: Devrait avoir les prédictions finales
301+ self .assertIn ('final_predictions' , results )
302+ self .assertEqual (len (results ['final_predictions' ]), len (self .test ))
303+
304+ def test_cv_then_final_on_test (self ):
305+ """Test CV sur train puis prédiction finale sur test"""
306+ # WHEN: On fait CV + prédiction finale (cas le plus complet)
307+ results = superstyl .train_svm (
308+ self .train ,
309+ test = self .test ,
310+ cross_validate = 'k-fold' ,
311+ k = 3 ,
312+ final_pred = True ,
313+ norms = True
314+ )
315+
316+ # THEN: Devrait avoir confusion matrix (CV) ET final_predictions
317+ self .assertIn ('confusion_matrix' , results )
318+ self .assertIn ('final_predictions' , results )
319+ self .assertEqual (len (results ['final_predictions' ]), len (self .test ))
320+
321+
322+ if __name__ == '__main__' :
323+ unittest .main ()
0 commit comments