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

如何使用Deeplearning4Java更新训练后的Word2Vec模型词汇表

用Deeplearning4j实现Word2Vec的增量训练(添加新词汇)

要给已训练的Word2Vec模型添加新词汇并更新权重,你需要通过增量训练实现——新词汇必须出现在带上下文的语料中,才能学到合理的向量。以下是具体步骤和代码修改方案:

核心逻辑

Deeplearning4j支持加载已有Word2Vec模型后,继续在包含新词汇的新语料上训练,自动将新词汇加入词汇表并基于上下文更新权重。关键是保持原模型核心参数一致,用较小学习率避免破坏原有训练成果。

具体实现步骤

1. 加载已训练模型

用WordVectorSerializer.readWord2VecModel读取之前保存的模型,确保复用原模型的层大小、窗口大小等核心配置。

2. 准备新语料

创建新的SentenceIterator指向包含新词汇的语料文件,新词汇必须搭配足够的上下文语句,否则向量无实际意义。

3. 配置增量训练参数

  • 复用原模型的layerSize、windowSize等参数
  • 设置较小的学习率(建议0.01左右,原训练通常用0.025)
  • 训练轮次无需太多(1-3轮即可,避免过度更新原有权重)

4. 执行增量训练并保存模型

调用模型的fit方法传入新语料迭代器,训练完成后重新保存更新后的模型。

修改后的完整代码示例

import org.deeplearning4j.models.word2vec.Word2Vec;
import org.deeplearning4j.text.sentenceiterator.FileSentenceIterator;
import org.deeplearning4j.text.sentenceiterator.SentenceIterator;
import org.deeplearning4j.text.tokenization.tokenizerfactory.DefaultTokenizerFactory;
import org.deeplearning4j.text.tokenization.tokenizerfactory.TokenizerFactory;
import org.deeplearning4j.text.tokenization.tokenizer.preprocessor.CommonPreprocessor;
import org.deeplearning4j.models.word2vec.WordVectorSerializer;

import java.io.File;
import java.io.IOException;

public class Test {
    private String inputFilePath = "原语料文件路径";
    private String modelFilePath = "已训练模型保存路径";
    private String newCorpusPath = "包含新词汇的新语料路径";

    public static void main(String[] args) throws IOException {
        Test test = new Test();
        // 初次训练(已完成训练可注释此行)
        // test.train();

        // 加载模型并执行增量训练
        Word2Vec word2VecModel = WordVectorSerializer.readWord2VecModel(new File(test.modelFilePath));
        test.incrementalTrain(word2VecModel);

        // 测试新词汇向量(假设新增词汇为"httpdbpediaorgresourceNewTerm")
        double similarity = word2VecModel.similarity("httpdbpediaorgresourcethe_terminator", "httpdbpediaorgresourceNewTerm");
        System.out.println(similarity);
    }

    // 原初次训练方法
    public void train() throws IOException {
        SentenceIterator sentenceIterator = new FileSentenceIterator(new File(inputFilePath));
        TokenizerFactory tokenizerFactory = new DefaultTokenizerFactory();
        tokenizerFactory.setTokenPreProcessor(new CommonPreprocessor());

        Word2Vec vec = new Word2Vec.Builder()
                .layerSize(100)
                .windowSize(5)
                .epochs(5)
                .elementsLearningAlgorithm(new org.deeplearning4j.models.word2vec.learning.SkipGram<>())
                .iterate(sentenceIterator)
                .tokenizerFactory(tokenizerFactory)
                .build();
        vec.fit();
        WordVectorSerializer.writeWordVectors(vec, modelFilePath);
    }

    // 新增增量训练方法
    public void incrementalTrain(Word2Vec existingModel) throws IOException {
        // 初始化新语料迭代器
        SentenceIterator newSentenceIterator = new FileSentenceIterator(new File(newCorpusPath));
        TokenizerFactory tokenizerFactory = new DefaultTokenizerFactory();
        tokenizerFactory.setTokenPreProcessor(new CommonPreprocessor());

        // 配置增量训练参数
        existingModel.setTokenizerFactory(tokenizerFactory);
        existingModel.setIterate(newSentenceIterator);
        existingModel.setLearningRate(0.01); // 小学习率保护原有权重
        existingModel.setEpochs(2); // 增量训练轮次无需过多

        // 执行训练:自动添加新词汇到词汇表并更新权重
        existingModel.fit();

        // 保存更新后的模型
        WordVectorSerializer.writeWordVectors(existingModel, modelFilePath);
    }
}

注意事项

  • 新语料必须包含新词汇的上下文语句,否则新词汇的向量仅为随机初始化后的微调,不具备语义价值。
  • 增量训练的学习率不可过大,否则会覆盖原模型已学到的有效权重。
  • 若新增词汇数量较多,可适当增加训练轮次,但建议不超过5轮。

内容的提问来源于stack exchange,提问作者Ahmed M Mahdi

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.19 02:15:41