77import org .springframework .beans .factory .annotation .Autowired ;
88import org .springframework .stereotype .Component ;
99import sentiment .data .Dataset ;
10+ import sentiment .data .DatasetUtils ;
1011import sentiment .data .SimpleDatasetLoader ;
1112import sentiment .data .DatasetLoadResult ;
1213import sentiment .data .SplitManifest ;
2122import sentiment .config .FeatureExtractionProperties ;
2223import weka .core .Instances ;
2324
24- import java .io .FileWriter ;
2525import java .io .IOException ;
2626import java .nio .file .Files ;
2727import java .nio .file .Path ;
2828import java .nio .file .Paths ;
2929import java .util .*;
30- import java .util .stream .Collectors ;
3130
3231/**
3332 * Service for training and persisting sentiment analysis models.
@@ -327,59 +326,8 @@ private List<Dataset> loadTrainingData(String dataPath, int maxSamples) throws E
327326
328327 // Balance classes by undersampling majority class to match minority class
329328 // This ensures fair comparison across datasets (IMDB 50/50, Amazon 50/50, Yelp 75/25 -> 50/50)
330- List <Dataset > balanced = balanceClasses (shuffled );
331-
332- return balanced ;
333- }
334-
335- /**
336- * Balances classes by undersampling the majority class to match the minority class.
337- * This ensures fair model comparison across datasets with different class distributions.
338- *
339- * @param data shuffled dataset (may be imbalanced)
340- * @return balanced dataset with equal positive/negative samples
341- */
342- private List <Dataset > balanceClasses (List <Dataset > data ) {
343- // Separate by class
344- List <Dataset > positive = data .stream ()
345- .filter (d -> d .getSentiment () == Dataset .SentimentLabel .POSITIVE )
346- .collect (Collectors .toList ());
347- List <Dataset > negative = data .stream ()
348- .filter (d -> d .getSentiment () == Dataset .SentimentLabel .NEGATIVE )
349- .collect (Collectors .toList ());
350-
351- // Find minority class size
352- int minoritySize = Math .min (positive .size (), negative .size ());
353-
354- if (minoritySize == 0 ) {
355- logger .warn ("One class has zero samples - cannot balance. Returning original data." );
356- return data ;
357- }
358-
359- // Check if already balanced (within 5% tolerance)
360- double ratio = (double ) Math .min (positive .size (), negative .size ()) /
361- Math .max (positive .size (), negative .size ());
362- if (ratio >= 0.95 ) {
363- logger .info ("Classes already balanced (ratio: {}) - no undersampling needed" ,
364- String .format ("%.2f" , ratio ));
365- return data ;
366- }
367-
368- // Undersample majority class
369- List <Dataset > balancedPositive = positive .subList (0 , minoritySize );
370- List <Dataset > balancedNegative = negative .subList (0 , minoritySize );
371-
372- // Combine and shuffle again to mix classes
373- List <Dataset > balanced = new ArrayList <>(minoritySize * 2 );
374- balanced .addAll (balancedPositive );
375- balanced .addAll (balancedNegative );
376- Collections .shuffle (balanced , new Random (42 ));
377-
378- logger .info ("Balanced classes: {} -> {} samples (undersampled {} from majority class)" ,
379- data .size (), balanced .size (), data .size () - balanced .size ());
380- logger .info ("Final distribution: positive={}, negative={}" , minoritySize , minoritySize );
381-
382- return balanced ;
329+ // Uses shared utility to ensure consistent balancing across training and data preparation
330+ return DatasetUtils .balanceByUndersampling (shuffled , 42L );
383331 }
384332
385333 /**
@@ -675,28 +623,19 @@ private void saveDataSplits(List<Dataset> trainData, List<Dataset> testData, Str
675623 }
676624
677625 // Save both splits (backward compatibility path - no manifest)
626+ // Uses shared CSV writer to ensure format consistency with data preparation pipeline
678627 logger .warn ("Saving splits without manifest (legacy mode). Run ./scripts/prepare_data.sh for proper setup." );
679- saveSplitToFile (trainData , processedDir .resolve ("train.csv" ), "train" );
680- saveSplitToFile (testData , processedDir .resolve ("test.csv" ), "test" );
681- }
682-
683- /**
684- * Save a data split to CSV file.
685- */
686- private void saveSplitToFile (List <Dataset > data , Path filePath , String splitName ) {
687- try (FileWriter writer = new FileWriter (filePath .toFile ())) {
688- writer .write ("review,sentiment\n " );
689-
690- for (Dataset sample : data ) {
691- String sentiment = sample .getSentiment ().name ().toLowerCase ();
692- String text = escapeCsv (sample .getText ());
693- writer .write ("\" " + text + "\" ," + sentiment + "\n " );
694- }
695-
696- logger .info ("Saved {} split ({} samples) to: {}" , splitName , data .size (), filePath );
697-
698- } catch (Exception e ) {
699- logger .error ("ARTIFACT_SAVE_FAILED: {} split - {}" , splitName , e .getMessage ());
628+ try {
629+ SimpleDatasetLoader .writeCsv (trainData , processedDir .resolve ("train.csv" ));
630+ logger .info ("Saved train split ({} samples) to: {}" , trainData .size (), processedDir .resolve ("train.csv" ));
631+ } catch (IOException e ) {
632+ logger .error ("ARTIFACT_SAVE_FAILED: train split - {}" , e .getMessage ());
633+ }
634+ try {
635+ SimpleDatasetLoader .writeCsv (testData , processedDir .resolve ("test.csv" ));
636+ logger .info ("Saved test split ({} samples) to: {}" , testData .size (), processedDir .resolve ("test.csv" ));
637+ } catch (IOException e ) {
638+ logger .error ("ARTIFACT_SAVE_FAILED: test split - {}" , e .getMessage ());
700639 }
701640 }
702641
@@ -722,15 +661,6 @@ private String inferDatasetName(String dataPath) {
722661 return dotIndex > 0 ? filename .substring (0 , dotIndex ) : filename ;
723662 }
724663
725- /**
726- * Escape CSV special characters in text.
727- */
728- private String escapeCsv (String text ) {
729- if (text == null ) return "" ;
730- // Escape quotes by doubling them
731- return text .replace ("\" " , "\" \" " );
732- }
733-
734664 private void analyzeAndPrintFeatureImportance (SentimentClassifier classifier ,
735665 StratifiedDataSplitter .DataSplit split ,
736666 int topFeaturesCount ,
0 commit comments