如何将Keras Tokenizer导入Java深度学习框架Deeplearning4j(DL4J)
将Keras Tokenizer导入Deeplearning4j的解决方案
我了解你已经用Keras(TensorFlow 2.1.4版)完成了20 News Group文本分类任务,准确率达到0.87,并且能保存模型和Tokenizer用于跨程序预测。现在需要把Keras的Tokenizer导入Deeplearning4j(DL4J),由于两者的Tokenizer体系结构并不完全兼容,我们需要通过导出核心词汇信息+在DL4J中重建Tokenizer的方式来实现,具体步骤如下:
步骤1:从Keras Tokenizer导出核心信息(Python端)
你已经用pickle保存了Tokenizer,现在先在Python里加载它,然后把关键数据导出为Java能读取的JSON格式:
import pickle import json # 加载保存的Keras Tokenizer with open('tokenizer.pickle', 'rb') as handle: tokenizer = pickle.load(handle) # 提取Keras Tokenizer的核心配置与词汇映射 tokenizer_info = { "word_index": tokenizer.word_index, "index_word": tokenizer.index_word, "oov_token": tokenizer.oov_token, "vocab_size": len(tokenizer.word_index) + 1, # Keras默认索引从1开始,0留作padding "lowercase": tokenizer.lowercase, "filters": tokenizer.filters } # 保存为JSON文件供Java读取 with open('tokenizer_config.json', 'w') as f: json.dump(tokenizer_info, f, indent=4)
步骤2:在DL4J中重建匹配的Tokenizer(Java端)
DL4J中我们可以通过自定义VocabCache和TokenPreProcessor,来还原Keras Tokenizer的词汇映射与分词规则:
import org.deeplearning4j.text.tokenization.tokenizer.DefaultTokenizer; import org.deeplearning4j.text.tokenization.tokenizer.Tokenizer; import org.deeplearning4j.text.tokenization.tokenizerfactory.DefaultTokenizerFactory; import org.nd4j.linalg.primitives.Pair; import com.google.gson.Gson; import java.io.FileReader; import java.util.Map; // 映射JSON中的Tokenizer配置信息 class KerasTokenizerInfo { public Map<String, Integer> word_index; public Map<Integer, String> index_word; public String oov_token; public int vocab_size; public boolean lowercase; public String filters; } public class KerasTokenizerImporter { public static void main(String[] args) throws Exception { // 读取JSON配置文件 Gson gson = new Gson(); KerasTokenizerInfo tokenizerInfo = gson.fromJson( new FileReader("tokenizer_config.json"), KerasTokenizerInfo.class ); // 创建DL4J TokenizerFactory,匹配Keras的基础配置 DefaultTokenizerFactory tokenizerFactory = new DefaultTokenizerFactory(); tokenizerFactory.setLowercase(tokenizerInfo.lowercase); // 设置与Keras一致的字符过滤规则 tokenizerFactory.setTokenPreProcessor(token -> { StringBuilder sb = new StringBuilder(); for (char c : token.toCharArray()) { if (tokenizerInfo.filters.indexOf(c) == -1) { sb.append(c); } } return sb.toString(); }); // 自定义VocabCache,导入Keras的词汇索引映射 CustomVocabCache vocabCache = new CustomVocabCache( tokenizerInfo.word_index, tokenizerInfo.index_word, tokenizerInfo.oov_token ); tokenizerFactory.setVocabCache(vocabCache); // 测试Tokenizer功能 Tokenizer tokenizer = tokenizerFactory.create("Sample test text from 20 News Group"); while (tokenizer.hasMoreTokens()) { String token = tokenizer.nextToken(); int index = vocabCache.indexOf(token); System.out.printf("Token: %s | Index: %d%n", token, index); } } // 自定义VocabCache,还原Keras的词汇索引逻辑 static class CustomVocabCache implements org.deeplearning4j.text.vocab.VocabCache { private final Map<String, Integer> wordToIndex; private final Map<Integer, String> indexToWord; private final String oovToken; public CustomVocabCache(Map<String, Integer> wordToIndex, Map<Integer, String> indexToWord, String oovToken) { this.wordToIndex = wordToIndex; this.indexToWord = indexToWord; this.oovToken = oovToken; } @Override public int indexOf(String word) { // 匹配Keras的OOV处理逻辑:未知词返回OOV索引(默认1) return wordToIndex.getOrDefault(word, oovToken != null ? 1 : 0); } @Override public String wordAtIndex(int index) { return indexToWord.getOrDefault(index, ""); } // 实现其他必要接口方法 @Override public int size() { return wordToIndex.size(); } @Override public void addWord(String word) {} @Override public boolean containsWord(String word) { return wordToIndex.containsKey(word); } @Override public int countWord(String word) { return 1; // 若需要词频,可在Python导出时添加tokenizer.word_counts字段 } @Override public Pair<String, Integer> wordAtIndex(int index, boolean add) { return Pair.of(wordAtIndex(index), index); } } }
关键注意事项
- 索引一致性:Keras Tokenizer默认词汇索引从1开始,0用于padding;DL4J默认索引从0开始,必须确保映射完全匹配,否则会导致模型预测出错。
- OOV处理对齐:如果你的Keras Tokenizer设置了
oov_token,要保证DL4J中未知词映射到相同的索引(Keras默认OOV索引为1)。 - 分词规则匹配:必须同步Keras的
lowercase和filters参数,否则分词结果不一致会直接影响预测准确性。 - 模型导入配合:导入Tokenizer的同时,可使用DL4J的
KerasModelImport.importKerasModelAndWeights()方法导入你的my_model.h5模型,确保端到端流程一致。
内容的提问来源于stack exchange,提问作者Pavan_opencodez
相关产品推荐
相关产品推荐

