Files
SuperBizAgent-java/src/main/java/org/example/service/VectorEmbeddingService.java
T
2026-05-31 21:45:14 +08:00

126 lines
4.0 KiB
Java

package org.example.service;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.stereotype.Service;
import java.util.ArrayList;
import java.util.Collections;
import java.util.List;
/**
* 向量嵌入服务
* 使用阿里云 DashScope Text Embedding API
*/
@Service
public class VectorEmbeddingService {
private static final Logger logger = LoggerFactory.getLogger(VectorEmbeddingService.class);
@Autowired
private EmbeddingModel embeddingModel;
/**
* 生成向量嵌入
* 调用阿里云 DashScope Text Embedding API
*
* @param content 文本内容
* @return 向量嵌入(浮点数列表)
*/
public List<Float> generateEmbedding(String content) {
try {
if (content == null || content.trim().isEmpty()) {
logger.warn("内容为空,无法生成向量");
throw new IllegalArgumentException("内容不能为空");
}
logger.debug("开始生成向量嵌入, 内容长度: {} 字符", content.length());
float[] embedding = embeddingModel.embed(content);
List<Float> floatEmbedding = new ArrayList<>(embedding.length);
for (float v : embedding) {
floatEmbedding.add(v);
}
logger.info("成功生成向量嵌入, 内容长度: {} 字符, 向量维度: {}",
content.length(), floatEmbedding.size());
return floatEmbedding;
} catch (Exception e) {
logger.error("生成向量嵌入失败, 内容长度: {}", content != null ? content.length() : 0, e);
throw new RuntimeException("生成向量嵌入失败: " + e.getMessage(), e);
}
}
public List<List<Float>> generateEmbeddings(List<String> contents) {
try {
if (contents == null || contents.isEmpty()) {
logger.warn("内容列表为空,无法生成向量");
return Collections.emptyList();
}
logger.info("开始批量生成向量嵌入, 数量: {}", contents.size());
List<float[]> embeddings = embeddingModel.embed(contents);
List<List<Float>> result = new ArrayList<>();
for (float[] embedding : embeddings) {
List<Float> floatEmbedding = new ArrayList<>(embedding.length);
for (float v : embedding) {
floatEmbedding.add(v);
}
result.add(floatEmbedding);
}
logger.info("成功批量生成向量嵌入, 数量: {}, 维度: {}",
result.size(),
result.isEmpty() ? 0 : result.get(0).size());
return result;
} catch (Exception e) {
logger.error("批量生成向量嵌入失败", e);
throw new RuntimeException("批量生成向量嵌入失败: " + e.getMessage(), e);
}
}
/**
* 生成查询向量
*
* @param query 查询文本
* @return 向量嵌入
*/
public List<Float> generateQueryVector(String query) {
return generateEmbedding(query);
}
/**
* 计算两个向量的余弦相似度
*
* @param vector1 向量1
* @param vector2 向量2
* @return 余弦相似度 [-1, 1]
*/
public float calculateCosineSimilarity(List<Float> vector1, List<Float> vector2) {
if (vector1.size() != vector2.size()) {
throw new IllegalArgumentException("向量维度不匹配");
}
float dotProduct = 0.0f;
float norm1 = 0.0f;
float norm2 = 0.0f;
for (int i = 0; i < vector1.size(); i++) {
dotProduct += vector1.get(i) * vector2.get(i);
norm1 += vector1.get(i) * vector1.get(i);
norm2 += vector2.get(i) * vector2.get(i);
}
return dotProduct / (float) (Math.sqrt(norm1) * Math.sqrt(norm2));
}
}