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

使用Deeplearning4j库,如何将Bag of Words输入Feed Forward Neural Network?

Hey there! Let me walk you through exactly how to feed Bag of Words (BoW) inputs into a Feed Forward Neural Network using Deeplearning4j. I’ll break this down step by step with practical Java code examples so you can follow along easily.

Step 1: Prepare Text Data & Generate Bag of Words Features

First, you need to turn raw text into numerical BoW vectors—this is the core of getting text ready for a neural network. DL4J has built-in tools to handle tokenization and vectorization without extra hassle.

Here’s a hands-on example:

import org.deeplearning4j.text.tokenization.tokenizerfactory.UimaTokenizerFactory;
import org.deeplearning4j.text.tokenization.tokenizerfactory.TokenizerFactory;
import org.deeplearning4j.text.vectorization.BagOfWordsVectorizer;
import org.nd4j.linalg.api.ndarray.INDArray;
import java.util.List;

// Replace this with your actual text dataset
List<String> textData = List.of(
    "Deeplearning4j makes neural network development straightforward",
    "Bag of Words converts text into count-based numerical vectors",
    "Feed forward networks excel at structured numerical input tasks"
);

// Initialize a tokenizer to split text into individual words
TokenizerFactory tokenizerFactory = new UimaTokenizerFactory();

// Build and fit the BoW vectorizer to your corpus
BagOfWordsVectorizer vectorizer = new BagOfWordsVectorizer.Builder()
    .setTokenizerFactory(tokenizerFactory)
    .setMinWordFrequency(1) // Include all words that appear at least once
    .build();
vectorizer.fit(textData);

// Convert all text samples to BoW feature vectors
INDArray bowFeatures = vectorizer.transform(textData);
Step 2: Configure the Feed Forward Neural Network

Next, set up your network architecture. The critical thing here is that the input layer size must match the dimension of your BoW vectors (which equals the size of your generated vocabulary).

For a binary classification task, here’s a sample configuration:

import org.deeplearning4j.nn.conf.MultiLayerConfiguration;
import org.deeplearning4j.nn.conf.NeuralNetConfiguration;
import org.deeplearning4j.nn.conf.layers.DenseLayer;
import org.deeplearning4j.nn.conf.layers.OutputLayer;
import org.deeplearning4j.nn.multilayer.MultiLayerNetwork;
import org.deeplearning4j.nn.weights.WeightInit;
import org.nd4j.linalg.activations.Activation;
import org.nd4j.linalg.lossfunctions.LossFunctions;

int inputSize = vectorizer.getVocab().size(); // Match BoW vector dimension
int numClasses = 2; // Adjust based on your task (e.g., 5 for multi-class classification)

MultiLayerConfiguration config = new NeuralNetConfiguration.Builder()
    .seed(42) // Ensure reproducible results
    .updater(new org.deeplearning4j.nn.conf.Updater.Adam(0.001)) // Adam optimizer with learning rate 0.001
    .weightInit(WeightInit.XAVIER)
    .list()
    // Hidden layer with 64 neurons and ReLU activation
    .layer(0, new DenseLayer.Builder()
        .nIn(inputSize)
        .nOut(64)
        .activation(Activation.RELU)
        .build())
    // Output layer for classification (softmax + cross-entropy loss)
    .layer(1, new OutputLayer.Builder(LossFunctions.LossFunction.NEGATIVELOGLIKELIHOOD)
        .nIn(64)
        .nOut(numClasses)
        .activation(Activation.SOFTMAX)
        .build())
    .build();

MultiLayerNetwork model = new MultiLayerNetwork(config);
model.init();
Step 3: Train the Model with BoW Data

Now pair your BoW features with labels, split into training/test sets, and start training. Assume your labels are one-hot encoded (standard for classification tasks):

import org.deeplearning4j.datasets.iterator.impl.ListDataSetIterator;
import org.nd4j.linalg.dataset.DataSet;
import org.nd4j.linalg.dataset.api.iterator.DataSetIterator;
import org.nd4j.linalg.factory.Nd4j;
import java.util.ArrayList;
import java.util.Collections;

// Sample one-hot encoded labels (replace with your actual labels)
INDArray labels = Nd4j.create(new double[][]{{1,0}, {0,1}, {1,0}});

// Create a DataSet pairing features and labels
DataSet dataSet = new DataSet(bowFeatures, labels);

// Split into training (80%) and test (20%) sets
List<DataSet> dataList = dataSet.asList();
Collections.shuffle(dataList);
int trainSize = (int) (dataList.size() * 0.8);
DataSet trainData = DataSet.merge(dataList.subList(0, trainSize));
DataSet testData = DataSet.merge(dataList.subList(trainSize, dataList.size()));

// Create iterators for training
DataSetIterator trainIterator = new ListDataSetIterator<>(List.of(trainData), 1);
DataSetIterator testIterator = new ListDataSetIterator<>(List.of(testData), 1);

// Add a listener to track training progress (prints score every 10 iterations)
model.setListeners(new org.deeplearning4j.optimize.listeners.ScoreIterationListener(10));

// Train the model for 100 epochs
model.fit(trainIterator, 100);
Step 4: Make Predictions on New Text

Once trained, use the same vectorizer to convert new text into BoW vectors and get predictions:

// New text to classify
String newText = "Deeplearning4j streamlines building feed forward networks";

// Convert to BoW vector
INDArray newBowVector = vectorizer.transform(List.of(newText));

// Get model prediction
INDArray prediction = model.output(newBowVector);

// Extract the predicted class (index of the highest probability)
int predictedClass = Nd4j.argMax(prediction, 1).getInt(0);
System.out.println("Predicted class index: " + predictedClass);

Quick Tips for Better Performance

  • Control Vocabulary Size: For large corpora, use .setVocabSize(10000) in the BagOfWordsVectorizer builder to avoid overly large input vectors.
  • Use TF-IDF: Instead of raw word counts, try TfidfVectorizer for normalized, importance-weighted vectors—this often improves model accuracy.
  • Tune Hyperparameters: Adjust hidden layer size, learning rate, and epoch count based on your dataset size and task complexity.

内容的提问来源于stack exchange,提问作者Rajani

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.19 08:34:55