基于Deep Java Library构建逻辑AND分类模型的问题求助
基于DJL实现逻辑AND运算神经网络的问题排查与解决
问题描述
尝试使用Deep Java Library(DJL)构建实现逻辑AND运算的简单神经网络,完成官方入门教程后修改代码实现以下功能:
- 从CSV文件读取训练数据
- 使用CSV数据训练模型
- 对二维浮点输入向量进行分类
但分类结果不符合预期:
输入:
float one [] = {1f,1f}; classify(one);
输出:
0: 0.5816633701324463 1: 0.4183366000652313
输入:
float zero [] = {1f,0f}; classify(zero);
输出:
0: 0.5276625156402588 1: 0.47233742475509644
提供的代码与训练数据
完整Java代码
import ai.djl.*; import ai.djl.training.*; import java.io.IOException; import java.nio.file.*; import ai.djl.ndarray.types.*; import ai.djl.training.loss.*; import ai.djl.training.listener.*; import ai.djl.training.evaluator.*; import ai.djl.basicmodelzoo.basic.*; import java.util.*; import java.util.stream.*; import org.apache.commons.csv.CSVFormat; import ai.djl.ndarray.*; import ai.djl.modality.*; import ai.djl.translate.*; import ai.djl.ndarray.NDList; import ai.djl.translate.TranslateException; import ai.djl.basicdataset.tabular.CsvDataset; import ai.djl.basicdataset.tabular.utils.Feature; public class App { public static void main( String[] args ) { boolean train = false; if(train) { try { training(); } catch (Exception e) { System.out.println("[ERROR] Could not train"); e.printStackTrace(); } } else { try { // define some input vectors for the neural network float zero [] = {1f,0f}; float zero2 [] = {0f,0f}; float one [] = {1f,1f}; classify(zero); } catch (Exception e) { System.out.println("[ERROR] Could not classify"); e.printStackTrace(); } } } /** * Classify with the trained neural network * @throws MalformedModelException * @throws IOException * @throws TranslateException */ static void classify(float [] input) throws MalformedModelException, IOException, TranslateException { Path modelDir = Paths.get("build/mlp"); Model model = Model.newInstance("mlpBlock"); model.setBlock(new Mlp(2, 2, new int[] {2})); model.load(modelDir); Translator<float[], Classifications> translator = new Translator<float[], Classifications>() { @Override public NDList processInput(TranslatorContext ctx, float[] input) { NDArray array = ctx.getNDManager().create(input); NDList ndList = new NDList(); ndList.add(array); return ndList; } @Override public Classifications processOutput(TranslatorContext ctx, NDList list) { NDArray probabilities = list.singletonOrThrow().softmax(0); List<String> classNames = IntStream.range(0, 2).mapToObj(String::valueOf).collect(Collectors.toList()); return new Classifications(classNames, probabilities); } @Override public Batchifier getBatchifier() { return Batchifier.STACK; } }; var predictor = model.newPredictor(translator); var classifications = predictor.predict(input); for(int i = 0; i < classifications.getProbabilities().size(); i++) { System.out.println(classifications.getClassNames().get(i) + ": " + classifications.getProbabilities().get(i)); } } /** * Train the neural network * @throws IOException * @throws TranslateException */ static void training() throws IOException, TranslateException { Path csvPath = Paths.get("TrainingDataAND.csv"); CSVFormat csvFormat = CSVFormat.DEFAULT.withHeader(); CsvDataset dataset = CsvDataset.builder() .optCsvFile(csvPath) .addFeature(new Feature("in1", true)) .addFeature(new Feature("in2", true)) .addLabel(new Feature("result", true)) .setSampling(2, true) .setCsvFormat(csvFormat) .build(); Model model = Model.newInstance("mlpBlock"); model.setBlock(new Mlp(2, 2, new int[] {2})); DefaultTrainingConfig config = new DefaultTrainingConfig(Loss.softmaxCrossEntropyLoss()) .addEvaluator(new Accuracy()) .addTrainingListeners(TrainingListener.Defaults.logging()); Trainer trainer = model.newTrainer(config); trainer.initialize(new Shape(1, 2)); int epoch = 2; EasyTrain.fit(trainer, epoch, dataset, null); Path modelDir = Paths.get("build/mlp"); Files.createDirectories(modelDir); model.setProperty("Epoch", String.valueOf(epoch)); model.save(modelDir, "mlpBlock"); } }
训练数据(TrainingDataAND.csv)
in1,in2,result 1,1,1 1,0,0 0,1,0 0,0,0
问题排查
1. 训练轮次严重不足
当前训练仅设置了2个epoch,即使是简单的AND逻辑任务,模型也没有足够的迭代次数来学习数据中的模式,导致权重几乎未收敛,输出接近随机猜测。
2. 数据集采样配置不合理
数据集仅4条样本,但设置setSampling(2, true)意味着每次取2条样本训练,小批量在极小数据集上会降低学习效率,无法充分利用全部数据信息。
3. 缺少适配的优化器配置
默认训练配置可能使用学习率较高的SGD优化器,在小数据集上容易导致模型震荡,难以稳定收敛。
解决方案
1. 大幅增加训练轮次
将训练轮次从2调整为500,确保模型有足够迭代次数收敛到正确权重。
2. 调整数据集采样设置
将批量大小设置为4(即全量样本),每次训练使用全部数据,提升学习效率:
.setSampling(4, true)
3. 配置合适的优化器
使用Adam优化器并设置0.1的学习率,替代默认优化器,加快收敛速度:
DefaultTrainingConfig config = new DefaultTrainingConfig(Loss.softmaxCrossEntropyLoss()) .optOptimizer(Optimizer.adam().setLearningRate(0.1f).build()) .addEvaluator(new Accuracy()) .addTrainingListeners(TrainingListener.Defaults.logging());
4. 确保训练开关开启
训练时将main方法中的train变量设为true,保证模型被重新训练并保存。
修正后的完整代码
import ai.djl.*; import ai.djl.training.*; import ai.djl.training.optimizer.Optimizer; import java.io.IOException; import java.nio.file.*; import ai.djl.ndarray.types.*; import ai.djl.training.loss.*; import ai.djl.training.listener.*; import ai.djl.training.evaluator.*; import ai.djl.basicmodelzoo.basic.*; import java.util.*; import java.util.stream.*; import org.apache.commons.csv.CSVFormat; import ai.djl.ndarray.*; import ai.djl.modality.*; import ai.djl.translate.*; import ai.djl.ndarray.NDList; import ai.djl.translate.TranslateException; import ai.djl.basicdataset.tabular.CsvDataset; import ai.djl.basicdataset.tabular.utils.Feature; public class App { public static void main( String[] args ) { boolean train = true; // 设置为true进行训练 if(train) { try { training(); } catch (Exception e) { System.out.println("[ERROR] Could not train"); e.printStackTrace(); } } else { try { // 测试各个输入 float zero [] = {1f,0f}; float zero2 [] = {0f,0f}; float one [] = {1f,1f}; float zero3 [] = {0f,1f}; System.out.println("输入[1,0]的分类结果:"); classify(zero); System.out.println("\n输入[0,0]的分类结果:"); classify(zero2); System.out.println("\n输入[1,1]的分类结果:"); classify(one); System.out.println("\n输入[0,1]的分类结果:"); classify(zero3); } catch (Exception e) { System.out.println("[ERROR] Could not classify"); e.printStackTrace(); } } } static void classify(float [] input) throws MalformedModelException, IOException, TranslateException { Path modelDir = Paths.get("build/mlp"); Model model = Model.newInstance("mlpBlock"); model.setBlock(new Mlp(2, 2, new int[] {2})); model.load(modelDir); Translator<float[], Classifications> translator = new Translator<float[], Classifications>() { @Override public NDList processInput(TranslatorContext ctx, float[] input) { NDArray array = ctx.getNDManager().create(input).reshape(new Shape(1, 2)); return new NDList(array); } @Override public Classifications processOutput(TranslatorContext ctx, NDList list) { NDArray probabilities = list.singletonOrThrow().softmax(1).squeeze(); List<String> classNames = Arrays.asList("0", "1"); return new Classifications(classNames, probabilities); } @Override public Batchifier getBatchifier() { return Batchifier.STACK; } }; var predictor = model.newPredictor(translator); var classifications = predictor.predict(input); for(int i = 0; i < classifications.getProbabilities().size(); i++) { System.out.println(classifications.getClassNames().get(i) + ": " + classifications.getProbabilities().get(i)); } } static void training() throws IOException, TranslateException { Path csvPath = Paths.get("TrainingDataAND.csv"); CSVFormat csvFormat = CSVFormat.DEFAULT.withHeader(); CsvDataset dataset = CsvDataset.builder() .optCsvFile(csvPath) .addFeature(new Feature("in1", true)) .addFeature(new Feature("in2", true)) .addLabel(new Feature("result", true)) .setSampling(4, true) // 使用全量样本作为批量 .setCsvFormat(csvFormat) .build(); Model model = Model.newInstance("mlpBlock"); model.setBlock(new Mlp(2, 2, new int[] {2})); DefaultTrainingConfig config = new DefaultTrainingConfig(Loss.softmaxCrossEntropyLoss()) .optOptimizer(Optimizer.adam().setLearningRate(0.1f).build()) // 使用Adam优化器 .addEvaluator(new Accuracy()) .addTrainingListeners(TrainingListener.Defaults.logging()); Trainer trainer = model.newTrainer(config); trainer.initialize(new Shape(1, 2)); int epoch = 500; // 增加训练轮次 EasyTrain.fit(trainer, epoch, dataset, null); Path modelDir = Paths.get("build/mlp"); Files.createDirectories(modelDir); model.setProperty("Epoch", String.valueOf(epoch)); model.save(modelDir, "mlpBlock"); System.out.println("模型训练完成并保存到:" + modelDir); } }
预期结果
训练完成后,测试各个输入会得到符合AND逻辑的结果:
- 输入
[1,1]:类别1的概率接近1.0 - 输入
[1,0]、[0,1]、[0,0]:类别0的概率接近1.0
内容的提问来源于stack exchange,提问作者Philipp S
相关产品推荐
相关产品推荐

