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

DJL中PyTorch加载GloVe词嵌入时类型转换异常排查

异常原因
  • ModelZooTextEmbedding的构造逻辑要求传入的ZooModel内部模型必须是DJL原生的ai.djl.nn.core.Embedding类实例,但通过PyTorch导出的TorchScript模型在DJL中会被封装为PtSymbolBlock,二者属于完全不同的类型,无法强制转换,因此触发ClassCastException,进而抛出“模型不是embedding”的参数异常。
解决方法

方法一:自定义TextEmbedding实现(推荐)

绕过ModelZooTextEmbedding的类型限制,自己实现TextEmbedding接口,直接调用加载好的PyTorch模型获取词嵌入。

1. 加载PyTorch模型(原有加载逻辑不变)

Criteria<NDList, NDList> criteria = Criteria.builder()
    .setTypes(NDList.class, NDList.class)
    .optEngine("PyTorch")
    .optModelPath(Paths.get("/home/user01/Downloads/RNN_Files/DJLdependencies"))
    .optModelName("GloVe.6B.50d.embedding.torchscript.pt")
    .build();

ZooModel<NDList, NDList> embeddingModel = criteria.loadModel();
Predictor<NDList, NDList> predictor = embeddingModel.newPredictor();

2. 自定义TextEmbedding类

import ai.djl.modality.nlp.embedding.TextEmbedding;
import ai.djl.ndarray.NDArray;
import ai.djl.ndarray.NDList;
import ai.djl.ndarray.types.DataType;
import ai.djl.ndarray.types.Shape;
import java.util.List;

public class CustomTorchEmbedding implements TextEmbedding {

    private Predictor<NDList, NDList> predictor;
    // 维护与Python中GloVe一致的词表映射(词转索引)
    private Vocabulary vocabulary;

    public CustomTorchEmbedding(Predictor<NDList, NDList> predictor, Vocabulary vocabulary) {
        this.predictor = predictor;
        this.vocabulary = vocabulary;
    }

    @Override
    public NDArray embedText(String text) {
        // 将文本拆分为词并转换为索引序列
        List<Integer> indices = vocabulary.getIndices(text.split(" "));
        // 转换为DJL可识别的NDArray(shape: [1, 序列长度])
        NDArray input = predictor.getManager().create(indices, new Shape(1, indices.size()), DataType.INT32);
        try {
            NDList output = predictor.predict(new NDList(input));
            return output.get(0);
        } catch (Exception e) {
            throw new RuntimeException("获取词嵌入失败", e);
        }
    }

    @Override
    public NDArray embedBatch(List<String> texts) {
        // 批量处理:将所有文本转换为索引序列,统一长度后生成batch输入
        List<List<Integer>> batchIndices = texts.stream()
                .map(text -> vocabulary.getIndices(text.split(" ")))
                .toList();
        
        // 此处需根据实际需求处理不同长度的文本(如padding补全)
        int batchSize = batchIndices.size();
        int seqLength = batchIndices.get(0).size();
        NDArray input = predictor.getManager().create(batchIndices, new Shape(batchSize, seqLength), DataType.INT32);
        
        try {
            NDList output = predictor.predict(new NDList(input));
            return output.get(0);
        } catch (Exception e) {
            throw new RuntimeException("批量获取词嵌入失败", e);
        }
    }

    @Override
    public long getEmbeddingSize() {
        // 返回GloVe的嵌入维度,此处为50
        return 50;
    }
}

3. 使用自定义词嵌入

// 初始化与Python GloVe词表一致的Vocabulary实例
Vocabulary vocab = ...; // 自行实现或加载词表映射逻辑
CustomTorchEmbedding customEmbedding = new CustomTorchEmbedding(predictor, vocab);

// 调用示例
NDArray embedding = customEmbedding.embedText("hello world");

方法二:适配DJL原生Embedding结构(不推荐)

ModelZooTextEmbedding依赖DJL内部的Embedding结构,PyTorch导出的模型无法直接适配,需要额外的模型转换与封装逻辑,复杂度高且灵活性差,因此不推荐使用。

额外注意事项
  • 输入维度一致性:Python中追踪模型时使用的输入形状(如[1,4])需与DJL调用时的输入形状匹配,避免维度不兼容。
  • 词表一致性:DJL中使用的词表必须与Python中GloVe的词表完全对应,否则索引映射错误会导致嵌入结果失效。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.24 00:25:04