|
|
@@ -0,0 +1,299 @@
|
|
|
+package com.zsjz.ai.module.embedding.core;
|
|
|
+
|
|
|
+import ai.djl.huggingface.tokenizers.Encoding;
|
|
|
+import ai.djl.huggingface.tokenizers.HuggingFaceTokenizer;
|
|
|
+import ai.onnxruntime.OnnxTensor;
|
|
|
+import ai.onnxruntime.OrtEnvironment;
|
|
|
+import ai.onnxruntime.OrtSession;
|
|
|
+import ai.onnxruntime.OrtSession.SessionOptions;
|
|
|
+import lombok.extern.slf4j.Slf4j;
|
|
|
+
|
|
|
+import java.nio.file.Files;
|
|
|
+import java.nio.file.Path;
|
|
|
+import java.util.ArrayList;
|
|
|
+import java.util.HashSet;
|
|
|
+import java.util.LinkedHashMap;
|
|
|
+import java.util.List;
|
|
|
+import java.util.Locale;
|
|
|
+import java.util.Map;
|
|
|
+import java.util.Set;
|
|
|
+
|
|
|
+/**
|
|
|
+ * ONNX 本地嵌入推理引擎(面向 bge-m3 导出物,兼容同类 XLM-RoBERTa/BERT 结构模型)。
|
|
|
+ *
|
|
|
+ * <p>加载约定:{@code modelDir} 下含 {@code tokenizer.json} 与 {@code onnx/model.onnx}
|
|
|
+ * (外部权重 {@code model.onnx_data} 必须与主文件同目录;不存在 onnx 子目录时回退
|
|
|
+ * {@code modelDir/model.onnx})。
|
|
|
+ *
|
|
|
+ * <p>推理流程:DJL HuggingFaceTokenizer 原生批内 padding/截断 → ORT 会话推理 →
|
|
|
+ * 取 {@code last_hidden_state[:, 0]}(CLS 位置)→ L2 归一化(bge-m3 官方 dense 用法)。
|
|
|
+ * 输入/输出名在加载时从会话动态探测,兼容不同导出命名(token_type_ids / position_ids
|
|
|
+ * 按需补零或按 RoBERTa 规则生成)。
|
|
|
+ *
|
|
|
+ * <p>实例由 {@code EmbeddingService} 持有为单例:权重常驻堆外内存(约 2.5~3GB),
|
|
|
+ * 懒加载由服务层控制,引擎本身保证 {@link #ensureLoaded()} 幂等线程安全。
|
|
|
+ * ORT {@code session.run} 线程安全;分词调用内部串行化以规避原生句柄并发风险。
|
|
|
+ */
|
|
|
+@Slf4j
|
|
|
+public class OnnxEmbeddingEngine implements AutoCloseable {
|
|
|
+
|
|
|
+ /** 模型可能接受、引擎能够供给的输入名(其余输入出现即报错,便于定位异常导出物) */
|
|
|
+ private static final Set<String> SUPPORTED_INPUTS = Set.of(
|
|
|
+ "input_ids", "attention_mask", "token_type_ids", "position_ids");
|
|
|
+
|
|
|
+ private final OrtEnvironment ortEnv = OrtEnvironment.getEnvironment();
|
|
|
+ private final Path modelDir;
|
|
|
+ private final int maxSeqLen;
|
|
|
+ private final int intraOpThreads;
|
|
|
+ private final String optLevel;
|
|
|
+
|
|
|
+ private volatile HuggingFaceTokenizer tokenizer;
|
|
|
+ private volatile OrtSession session;
|
|
|
+ private volatile boolean closed;
|
|
|
+
|
|
|
+ /** 会话实际声明的输入名集合 */
|
|
|
+ private volatile Set<String> inputNames;
|
|
|
+ /** 池化输出名(优先 last_hidden_state,否则首个输出) */
|
|
|
+ private volatile String embeddingOutputName;
|
|
|
+ /** 输出向量维度,首次推理后填充 */
|
|
|
+ private volatile int dimension = -1;
|
|
|
+
|
|
|
+ public OnnxEmbeddingEngine(Path modelDir, int maxSeqLen, int intraOpThreads, String optLevel) {
|
|
|
+ this.modelDir = modelDir;
|
|
|
+ this.maxSeqLen = maxSeqLen;
|
|
|
+ this.intraOpThreads = intraOpThreads;
|
|
|
+ this.optLevel = optLevel;
|
|
|
+ }
|
|
|
+
|
|
|
+ /**
|
|
|
+ * 懒加载分词器与 ORT 会话(幂等,首次调用承担 2.3GB 权重的加载耗时)
|
|
|
+ */
|
|
|
+ public synchronized void ensureLoaded() {
|
|
|
+ if (session != null || closed) {
|
|
|
+ return;
|
|
|
+ }
|
|
|
+ if (!Files.isDirectory(modelDir)) {
|
|
|
+ throw new EmbeddingException("嵌入模型目录不存在: " + modelDir.toAbsolutePath()
|
|
|
+ + "(请检查 zsjz.embedding.model-path 配置,目录下应含 tokenizer.json 与 onnx/model.onnx)");
|
|
|
+ }
|
|
|
+ long start = System.currentTimeMillis();
|
|
|
+ try {
|
|
|
+ // 显式传 modelMaxLength:目录方式加载不读 tokenizer_config.json,DJL 缺省会按 512 截断
|
|
|
+ tokenizer = HuggingFaceTokenizer.newInstance(modelDir,
|
|
|
+ Map.of("modelMaxLength", String.valueOf(maxSeqLen)));
|
|
|
+ } catch (Exception e) {
|
|
|
+ throw new EmbeddingException("加载分词器失败: " + modelDir.resolve("tokenizer.json"), e);
|
|
|
+ }
|
|
|
+ Path modelFile = modelDir.resolve("onnx").resolve("model.onnx");
|
|
|
+ if (!Files.isRegularFile(modelFile)) {
|
|
|
+ modelFile = modelDir.resolve("model.onnx");
|
|
|
+ }
|
|
|
+ if (!Files.isRegularFile(modelFile)) {
|
|
|
+ throw new EmbeddingException("模型目录下未找到 ONNX 模型文件(期望 onnx/model.onnx 或 model.onnx): "
|
|
|
+ + modelDir.toAbsolutePath());
|
|
|
+ }
|
|
|
+ try (SessionOptions options = new SessionOptions()) {
|
|
|
+ options.setOptimizationLevel(parseOptLevel(optLevel));
|
|
|
+ if (intraOpThreads > 0) {
|
|
|
+ options.setIntraOpNumThreads(intraOpThreads);
|
|
|
+ }
|
|
|
+ // 从文件路径创建会话:ONNX Runtime 依据模型文件所在目录解析外部权重 model.onnx_data
|
|
|
+ session = ortEnv.createSession(modelFile.toAbsolutePath().toString(), options);
|
|
|
+ } catch (EmbeddingException e) {
|
|
|
+ throw e;
|
|
|
+ } catch (Exception e) {
|
|
|
+ throw new EmbeddingException("创建 ONNX 会话失败: " + modelFile, e);
|
|
|
+ }
|
|
|
+
|
|
|
+ inputNames = session.getInputNames();
|
|
|
+ Set<String> unsupported = new HashSet<>(inputNames);
|
|
|
+ unsupported.removeAll(SUPPORTED_INPUTS);
|
|
|
+ if (!unsupported.isEmpty()) {
|
|
|
+ throw new EmbeddingException("模型存在引擎无法供给的输入: " + unsupported
|
|
|
+ + "(模型输入: " + inputNames + ")");
|
|
|
+ }
|
|
|
+ if (!inputNames.contains("input_ids")) {
|
|
|
+ throw new EmbeddingException("模型缺少必需输入 input_ids(实际输入: " + inputNames + ")");
|
|
|
+ }
|
|
|
+ embeddingOutputName = session.getOutputNames().stream()
|
|
|
+ .filter("last_hidden_state"::equals)
|
|
|
+ .findFirst()
|
|
|
+ .orElse(session.getOutputNames().iterator().next());
|
|
|
+ log.info("本地嵌入模型加载完成: modelDir={}, 耗时 {} ms, inputs={}, output={}, maxSeqLen={}",
|
|
|
+ modelDir.toAbsolutePath(), System.currentTimeMillis() - start, inputNames,
|
|
|
+ embeddingOutputName, maxSeqLen);
|
|
|
+ }
|
|
|
+
|
|
|
+ public boolean isReady() {
|
|
|
+ return session != null && !closed;
|
|
|
+ }
|
|
|
+
|
|
|
+ /**
|
|
|
+ * 输出向量维度(首次调用会触发模型加载与一次最小推理)
|
|
|
+ */
|
|
|
+ public int getDimension() {
|
|
|
+ ensureLoaded();
|
|
|
+ if (dimension <= 0) {
|
|
|
+ embed(List.of("维度探测"));
|
|
|
+ }
|
|
|
+ return dimension;
|
|
|
+ }
|
|
|
+
|
|
|
+ /**
|
|
|
+ * 批量嵌入一批文本(一次会话调用,批内由分词器原生按最长 padding)
|
|
|
+ *
|
|
|
+ * <p>调用方负责按配置的 batch-size 分批;本方法不做条数限制。
|
|
|
+ *
|
|
|
+ * @return 与输入顺序一致的 L2 归一化向量列表
|
|
|
+ */
|
|
|
+ public List<float[]> embed(List<String> texts) {
|
|
|
+ ensureLoaded();
|
|
|
+ if (texts == null || texts.isEmpty()) {
|
|
|
+ return List.of();
|
|
|
+ }
|
|
|
+ for (int i = 0; i < texts.size(); i++) {
|
|
|
+ if (texts.get(i) == null || texts.get(i).isBlank()) {
|
|
|
+ throw new EmbeddingException("嵌入文本不能为空(第 " + i + " 条)");
|
|
|
+ }
|
|
|
+ }
|
|
|
+
|
|
|
+ Encoding[] encodings;
|
|
|
+ synchronized (tokenizer) {
|
|
|
+ encodings = tokenizer.batchEncode(texts);
|
|
|
+ }
|
|
|
+ int batch = encodings.length;
|
|
|
+ int seqLen = encodings[0].getIds().length;
|
|
|
+ long[][] inputIds = new long[batch][];
|
|
|
+ long[][] attentionMask = new long[batch][];
|
|
|
+ for (int i = 0; i < batch; i++) {
|
|
|
+ inputIds[i] = encodings[i].getIds();
|
|
|
+ attentionMask[i] = encodings[i].getAttentionMask();
|
|
|
+ if (encodings[i].getIds().length != seqLen) {
|
|
|
+ throw new EmbeddingException("分词器批内 padding 结果长度不一致(期望 " + seqLen
|
|
|
+ + ",第 " + i + " 条 " + encodings[i].getIds().length + ")");
|
|
|
+ }
|
|
|
+ }
|
|
|
+
|
|
|
+ Map<String, OnnxTensor> inputs = new LinkedHashMap<>();
|
|
|
+ try {
|
|
|
+ inputs.put("input_ids", OnnxTensor.createTensor(ortEnv, inputIds));
|
|
|
+ inputs.put("attention_mask", OnnxTensor.createTensor(ortEnv, attentionMask));
|
|
|
+ if (inputNames.contains("token_type_ids")) {
|
|
|
+ inputs.put("token_type_ids", OnnxTensor.createTensor(ortEnv, new long[batch][seqLen]));
|
|
|
+ }
|
|
|
+ if (inputNames.contains("position_ids")) {
|
|
|
+ inputs.put("position_ids", OnnxTensor.createTensor(ortEnv, robertaPositionIds(attentionMask)));
|
|
|
+ }
|
|
|
+ } catch (Exception e) {
|
|
|
+ inputs.values().forEach(OnnxTensor::close);
|
|
|
+ throw new EmbeddingException("构建模型输入张量失败", e);
|
|
|
+ }
|
|
|
+
|
|
|
+ long start = System.currentTimeMillis();
|
|
|
+ try (OrtSession.Result result = session.run(inputs)) {
|
|
|
+ Object value = result.get(embeddingOutputName)
|
|
|
+ .orElseThrow(() -> new EmbeddingException("模型输出缺失: " + embeddingOutputName))
|
|
|
+ .getValue();
|
|
|
+ List<float[]> vectors = poolAndNormalize(value, batch);
|
|
|
+ log.debug("本地嵌入推理完成: batch={}, seqLen={}, 耗时 {} ms",
|
|
|
+ batch, seqLen, System.currentTimeMillis() - start);
|
|
|
+ return vectors;
|
|
|
+ } catch (EmbeddingException e) {
|
|
|
+ throw e;
|
|
|
+ } catch (Exception e) {
|
|
|
+ throw new EmbeddingException("ONNX 推理失败: " + e.getMessage(), e);
|
|
|
+ } finally {
|
|
|
+ inputs.values().forEach(OnnxTensor::close);
|
|
|
+ }
|
|
|
+ }
|
|
|
+
|
|
|
+ /**
|
|
|
+ * CLS 池化 + L2 归一化:兼容 [batch, seq, hidden] 与 [batch, hidden] 两种输出形态
|
|
|
+ */
|
|
|
+ private List<float[]> poolAndNormalize(Object value, int batch) {
|
|
|
+ if (value instanceof float[][][] hidden) {
|
|
|
+ if (hidden.length != batch || hidden[0].length == 0) {
|
|
|
+ throw new EmbeddingException("模型输出 batch 维度异常(期望 " + batch
|
|
|
+ + ",实际 " + hidden.length + ")");
|
|
|
+ }
|
|
|
+ List<float[]> vectors = new ArrayList<>(batch);
|
|
|
+ for (int i = 0; i < batch; i++) {
|
|
|
+ vectors.add(l2Normalize(hidden[i][0]));
|
|
|
+ }
|
|
|
+ dimension = vectors.get(0).length;
|
|
|
+ return vectors;
|
|
|
+ }
|
|
|
+ if (value instanceof float[][] pooled) {
|
|
|
+ if (pooled.length != batch) {
|
|
|
+ throw new EmbeddingException("模型输出 batch 维度异常(期望 " + batch
|
|
|
+ + ",实际 " + pooled.length + ")");
|
|
|
+ }
|
|
|
+ List<float[]> vectors = new ArrayList<>(batch);
|
|
|
+ for (int i = 0; i < batch; i++) {
|
|
|
+ vectors.add(l2Normalize(pooled[i]));
|
|
|
+ }
|
|
|
+ dimension = vectors.get(0).length;
|
|
|
+ return vectors;
|
|
|
+ }
|
|
|
+ throw new EmbeddingException("不支持的模型输出形态: "
|
|
|
+ + (value == null ? "null" : value.getClass().getName()));
|
|
|
+ }
|
|
|
+
|
|
|
+ private static float[] l2Normalize(float[] vector) {
|
|
|
+ double sum = 0;
|
|
|
+ for (float v : vector) {
|
|
|
+ sum += (double) v * v;
|
|
|
+ }
|
|
|
+ float norm = (float) Math.sqrt(sum);
|
|
|
+ if (norm > 0) {
|
|
|
+ for (int i = 0; i < vector.length; i++) {
|
|
|
+ vector[i] /= norm;
|
|
|
+ }
|
|
|
+ }
|
|
|
+ return vector;
|
|
|
+ }
|
|
|
+
|
|
|
+ /**
|
|
|
+ * RoBERTa/XLM-R 位置编码:真实 token 位置 = attention_mask 前缀和 + padding_idx + 1,
|
|
|
+ * 右填充布局下与 HF create_position_ids_from_input_ids 结果一致(填充位被注意力掩码屏蔽)
|
|
|
+ */
|
|
|
+ private static long[][] robertaPositionIds(long[][] attentionMask) {
|
|
|
+ long[][] positions = new long[attentionMask.length][];
|
|
|
+ for (int i = 0; i < attentionMask.length; i++) {
|
|
|
+ long[] row = new long[attentionMask[i].length];
|
|
|
+ long cumsum = 0;
|
|
|
+ for (int j = 0; j < row.length; j++) {
|
|
|
+ cumsum += attentionMask[i][j];
|
|
|
+ row[j] = cumsum + 1;
|
|
|
+ }
|
|
|
+ positions[i] = row;
|
|
|
+ }
|
|
|
+ return positions;
|
|
|
+ }
|
|
|
+
|
|
|
+ private static SessionOptions.OptLevel parseOptLevel(String optLevel) {
|
|
|
+ try {
|
|
|
+ return SessionOptions.OptLevel.valueOf(optLevel.trim().toUpperCase(Locale.ROOT));
|
|
|
+ } catch (Exception e) {
|
|
|
+ log.warn("无效的图优化级别 {},回退 BASIC_OPT", optLevel);
|
|
|
+ return SessionOptions.OptLevel.BASIC_OPT;
|
|
|
+ }
|
|
|
+ }
|
|
|
+
|
|
|
+ @Override
|
|
|
+ public synchronized void close() {
|
|
|
+ closed = true;
|
|
|
+ if (session != null) {
|
|
|
+ try {
|
|
|
+ session.close();
|
|
|
+ } catch (Exception e) {
|
|
|
+ log.warn("关闭 ONNX 会话失败: {}", e.getMessage());
|
|
|
+ }
|
|
|
+ session = null;
|
|
|
+ }
|
|
|
+ if (tokenizer != null) {
|
|
|
+ tokenizer.close();
|
|
|
+ tokenizer = null;
|
|
|
+ }
|
|
|
+ }
|
|
|
+}
|