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
相关产品推荐
相关产品推荐

