Skip to content

Commit 5b38163

Browse files
committed
fix SMOTE n_neighbors issues
1 parent 6436dcc commit 5b38163

2 files changed

Lines changed: 358 additions & 7 deletions

File tree

‎superstyl/svm.py‎

Lines changed: 35 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -104,15 +104,40 @@ def train_svm(
104104
# Ensures that the resampling method does not attempt to use more neighbors than available samples
105105
# in the minority class, which produced the error.
106106
min_class_size = min(Counter(classes).values())
107-
n_neighbors = min(5, min_class_size - 1) # Default n_neighbors in SMOTE is 5
108-
# In case we have to temper with the n_neighbors, we print a warning message to the user
109-
# (might be written more clearly, but we want a short message, right?)
110-
if 0 < n_neighbors >= min_class_size:
111-
print(f"Warning: Adjusting n_neighbors for SMOTE to {n_neighbors} due to small class size.")
112107

113-
if n_neighbors == 0:
114-
print("Warning: at least one class only has a single individual; cannot apply SMOTE(Tomek).")
108+
# For cross-validation, we need to be more conservative since each fold
109+
# will have fewer samples. We use a safety margin to account for this.
110+
if cross_validate is not None:
111+
# In CV, each fold typically has fewer samples
112+
# Use a more conservative estimate
113+
if cross_validate == 'leave-one-out':
114+
# LOO: each fold removes 1 sample
115+
effective_min = max(1, min_class_size - 1)
116+
elif cross_validate == 'k-fold':
117+
k_folds = k if k > 0 else 10
118+
effective_min = max(1, int(min_class_size * (k_folds - 1) / k_folds))
119+
elif cross_validate == 'group-k-fold':
120+
# use conservative estimate
121+
k_folds = k if k > 0 else 10
122+
effective_min = max(2, min_class_size // k_folds)
123+
else:
124+
effective_min = min_class_size
115125
else:
126+
effective_min = min_class_size
127+
128+
# Calculate n_neighbors with safety margin
129+
# We need at least 2 samples to have 1 neighbor
130+
n_neighbors = min(5, effective_min - 1) if effective_min > 1 else 0
131+
132+
# Ensure n_neighbors is at least 1 for SMOTE to work
133+
if n_neighbors < 1:
134+
print(f"Warning: Smallest class has only {min_class_size} sample(s). Cannot apply SMOTE.")
135+
print(" Skipping SMOTE resampling. Consider using 'upsampling' or 'downsampling' instead.")
136+
elif n_neighbors < 5:
137+
print(f"Warning: Adjusting n_neighbors for SMOTE to {n_neighbors} due to small class size.")
138+
print(f" (Minimum class size: {min_class_size}, effective for CV: {effective_min})")
139+
140+
if n_neighbors >= 1:
116141
if balance == 'SMOTE':
117142
estimators.append(('sampling', over.SMOTE(k_neighbors=n_neighbors, random_state=42)))
118143
elif balance == 'SMOTETomek':
@@ -183,6 +208,9 @@ def train_svm(
183208
pipe.fit(train, classes)
184209

185210
if final_pred:
211+
if test is None:
212+
raise ValueError("final_pred=True requires a test set!")
213+
preds = pipe.predict(test)
186214
preds = pipe.predict(test)
187215

188216
# And now the simple case where there is only one svm to train

‎tests/test_SMOTE.py‎

Lines changed: 323 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,323 @@
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

Comments
 (0)