如何使用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
相关产品推荐
相关产品推荐

