Skip to content

Commit f38047b

Browse files
committed
refactor(data): extract shared utilities for class balancing and CSV writing
1 parent cc00ee0 commit f38047b

6 files changed

Lines changed: 353 additions & 157 deletions

File tree

src/main/java/sentiment/data/DataPreparer.java

Lines changed: 6 additions & 72 deletions
Original file line numberDiff line numberDiff line change
@@ -4,14 +4,12 @@
44
import org.slf4j.LoggerFactory;
55
import sentiment.evaluation.StratifiedDataSplitter;
66

7-
import java.io.FileWriter;
87
import java.io.IOException;
98
import java.nio.file.Files;
109
import java.nio.file.Path;
1110
import java.nio.file.Paths;
1211
import java.time.Instant;
1312
import java.util.*;
14-
import java.util.stream.Collectors;
1513

1614
/**
1715
* Prepares and manages immutable data splits for training.
@@ -165,8 +163,10 @@ private SplitManifest createSplits(String domain, String rawDataPath, int maxSam
165163
Path trainPath = processedDir.resolve("train.csv");
166164
Path testPath = processedDir.resolve("test.csv");
167165

168-
saveSplitToCsv(split.train, trainPath);
169-
saveSplitToCsv(split.test, testPath);
166+
SimpleDatasetLoader.writeCsv(split.train, trainPath);
167+
logger.info("Saved {} samples to {}", split.train.size(), trainPath);
168+
SimpleDatasetLoader.writeCsv(split.test, testPath);
169+
logger.info("Saved {} samples to {}", split.test.size(), testPath);
170170

171171
// Step 4: Create and save manifest
172172
logger.info("Step 4/4: Creating manifest with checksums...");
@@ -221,74 +221,8 @@ private List<Dataset> loadAndBalance(String rawDataPath, int maxSamples, int see
221221
List<Dataset> shuffled = new ArrayList<>(allData);
222222
Collections.shuffle(shuffled, new Random(seed));
223223

224-
// Balance classes by undersampling majority
225-
return balanceClasses(shuffled, seed);
226-
}
227-
228-
/**
229-
* Balance classes by undersampling the majority class.
230-
*/
231-
private List<Dataset> balanceClasses(List<Dataset> data, int seed) {
232-
List<Dataset> positive = data.stream()
233-
.filter(d -> d.getSentiment() == Dataset.SentimentLabel.POSITIVE)
234-
.collect(Collectors.toList());
235-
List<Dataset> negative = data.stream()
236-
.filter(d -> d.getSentiment() == Dataset.SentimentLabel.NEGATIVE)
237-
.collect(Collectors.toList());
238-
239-
int minoritySize = Math.min(positive.size(), negative.size());
240-
241-
if (minoritySize == 0) {
242-
logger.warn("One class has zero samples - cannot balance");
243-
return data;
244-
}
245-
246-
// Check if already balanced (within 5% tolerance)
247-
double ratio = (double) Math.min(positive.size(), negative.size()) /
248-
Math.max(positive.size(), negative.size());
249-
if (ratio >= 0.95) {
250-
logger.info("Classes already balanced (ratio: {}) - no undersampling needed",
251-
String.format("%.2f", ratio));
252-
return data;
253-
}
254-
255-
// Undersample majority class
256-
List<Dataset> balancedPositive = positive.subList(0, minoritySize);
257-
List<Dataset> balancedNegative = negative.subList(0, minoritySize);
258-
259-
List<Dataset> balanced = new ArrayList<>(minoritySize * 2);
260-
balanced.addAll(balancedPositive);
261-
balanced.addAll(balancedNegative);
262-
Collections.shuffle(balanced, new Random(seed));
263-
264-
logger.info("Balanced classes: {} -> {} samples (undersampled {} from majority)",
265-
data.size(), balanced.size(), data.size() - balanced.size());
266-
267-
return balanced;
268-
}
269-
270-
/**
271-
* Save a data split to CSV file.
272-
*/
273-
private void saveSplitToCsv(List<Dataset> data, Path filePath) throws IOException {
274-
try (FileWriter writer = new FileWriter(filePath.toFile())) {
275-
writer.write("review,sentiment\n");
276-
277-
for (Dataset sample : data) {
278-
String sentiment = sample.getSentiment().name().toLowerCase();
279-
String text = escapeCsv(sample.getText());
280-
writer.write("\"" + text + "\"," + sentiment + "\n");
281-
}
282-
}
283-
logger.info("Saved {} samples to {}", data.size(), filePath);
284-
}
285-
286-
/**
287-
* Escape CSV special characters in text.
288-
*/
289-
private String escapeCsv(String text) {
290-
if (text == null) return "";
291-
return text.replace("\"", "\"\"");
224+
// Balance classes by undersampling majority (uses shared utility)
225+
return DatasetUtils.balanceByUndersampling(shuffled, seed);
292226
}
293227

294228
// ===== CLI Entry Point =====
Lines changed: 89 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,89 @@
1+
package sentiment.data;
2+
3+
import org.slf4j.Logger;
4+
import org.slf4j.LoggerFactory;
5+
6+
import java.util.ArrayList;
7+
import java.util.Collections;
8+
import java.util.List;
9+
import java.util.Random;
10+
import java.util.stream.Collectors;
11+
12+
/**
13+
* Utility methods for dataset manipulation.
14+
* Provides shared operations used across data preparation and training pipelines.
15+
*/
16+
public final class DatasetUtils {
17+
18+
private static final Logger logger = LoggerFactory.getLogger(DatasetUtils.class);
19+
20+
/**
21+
* Tolerance threshold for considering classes "already balanced".
22+
* If minority/majority ratio >= this value, no undersampling is performed.
23+
*/
24+
private static final double BALANCE_TOLERANCE = 0.95;
25+
26+
private DatasetUtils() {
27+
// Utility class - prevent instantiation
28+
}
29+
30+
/**
31+
* Balance classes by undersampling the majority class to match the minority class.
32+
* This ensures fair model comparison across datasets with different class distributions.
33+
*
34+
* The method:
35+
* 1. Separates data by sentiment class
36+
* 2. If already balanced (within 5% tolerance), returns original data
37+
* 3. Otherwise, undersamples majority class to match minority class size
38+
* 4. Shuffles the result to mix classes
39+
*
40+
* @param data input dataset (may be imbalanced)
41+
* @param seed random seed for reproducible shuffling
42+
* @return balanced dataset with equal positive/negative samples
43+
*/
44+
public static List<Dataset> balanceByUndersampling(List<Dataset> data, long seed) {
45+
if (data == null || data.isEmpty()) {
46+
return data;
47+
}
48+
49+
// Separate by class
50+
List<Dataset> positive = data.stream()
51+
.filter(d -> d.getSentiment() == Dataset.SentimentLabel.POSITIVE)
52+
.collect(Collectors.toList());
53+
List<Dataset> negative = data.stream()
54+
.filter(d -> d.getSentiment() == Dataset.SentimentLabel.NEGATIVE)
55+
.collect(Collectors.toList());
56+
57+
// Find minority class size
58+
int minoritySize = Math.min(positive.size(), negative.size());
59+
60+
if (minoritySize == 0) {
61+
logger.warn("One class has zero samples - cannot balance. Returning original data.");
62+
return data;
63+
}
64+
65+
// Check if already balanced (within tolerance)
66+
double ratio = (double) minoritySize / Math.max(positive.size(), negative.size());
67+
if (ratio >= BALANCE_TOLERANCE) {
68+
logger.info("Classes already balanced (ratio: {}) - no undersampling needed",
69+
String.format("%.2f", ratio));
70+
return data;
71+
}
72+
73+
// Undersample majority class
74+
List<Dataset> balancedPositive = positive.subList(0, minoritySize);
75+
List<Dataset> balancedNegative = negative.subList(0, minoritySize);
76+
77+
// Combine and shuffle to mix classes
78+
List<Dataset> balanced = new ArrayList<>(minoritySize * 2);
79+
balanced.addAll(balancedPositive);
80+
balanced.addAll(balancedNegative);
81+
Collections.shuffle(balanced, new Random(seed));
82+
83+
logger.info("Balanced classes: {} -> {} samples (undersampled {} from majority class)",
84+
data.size(), balanced.size(), data.size() - balanced.size());
85+
logger.info("Final distribution: positive={}, negative={}", minoritySize, minoritySize);
86+
87+
return balanced;
88+
}
89+
}

src/main/java/sentiment/data/SimpleDatasetLoader.java

Lines changed: 39 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,7 @@
99
import sentiment.util.ValidationUtils;
1010

1111
import java.io.BufferedReader;
12+
import java.io.BufferedWriter;
1213
import java.io.FileReader;
1314
import java.io.IOException;
1415
import java.nio.charset.StandardCharsets;
@@ -397,6 +398,44 @@ private List<Dataset> loadPlaintext(String filePath) throws DataLoadingException
397398
);
398399
}
399400

401+
// CSV WRITING
402+
403+
/**
404+
* Write datasets to CSV file in the standard format used by this loader.
405+
* Format: "review,sentiment" with quoted text field and lowercase sentiment.
406+
*
407+
* This is the inverse of loadCsv() - keeping read/write together ensures format consistency.
408+
*
409+
* @param data list of datasets to write
410+
* @param path destination file path
411+
* @throws IOException if file cannot be written
412+
*/
413+
public static void writeCsv(List<Dataset> data, Path path) throws IOException {
414+
try (BufferedWriter writer = Files.newBufferedWriter(path, StandardCharsets.UTF_8)) {
415+
// Use explicit \n for cross-platform consistency (matches original format)
416+
writer.write("review,sentiment\n");
417+
418+
for (Dataset sample : data) {
419+
writer.write('"');
420+
writer.write(escapeCsvField(sample.getText()));
421+
writer.write("\",");
422+
writer.write(sample.getSentiment().name().toLowerCase());
423+
writer.write('\n');
424+
}
425+
}
426+
}
427+
428+
/**
429+
* Escape special characters in CSV field (quotes and newlines).
430+
*/
431+
private static String escapeCsvField(String text) {
432+
if (text == null) {
433+
return "";
434+
}
435+
// Double quotes to escape them in CSV
436+
return text.replace("\"", "\"\"");
437+
}
438+
400439
// UTILITIES
401440

402441
/**

src/main/java/sentiment/training/ModelTrainer.java

Lines changed: 15 additions & 85 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@
77
import org.springframework.beans.factory.annotation.Autowired;
88
import org.springframework.stereotype.Component;
99
import sentiment.data.Dataset;
10+
import sentiment.data.DatasetUtils;
1011
import sentiment.data.SimpleDatasetLoader;
1112
import sentiment.data.DatasetLoadResult;
1213
import sentiment.data.SplitManifest;
@@ -21,13 +22,11 @@
2122
import sentiment.config.FeatureExtractionProperties;
2223
import weka.core.Instances;
2324

24-
import java.io.FileWriter;
2525
import java.io.IOException;
2626
import java.nio.file.Files;
2727
import java.nio.file.Path;
2828
import java.nio.file.Paths;
2929
import 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

Comments
 (0)