You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

基于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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.06 06:05:50