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

如何将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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.26 11:01:22