222 lines
5.8 KiB
Go
222 lines
5.8 KiB
Go
package storage
|
||
|
||
import (
|
||
"bytes"
|
||
"encoding/json"
|
||
"fmt"
|
||
"io"
|
||
"math"
|
||
"net/http"
|
||
"os"
|
||
"sort"
|
||
"strings"
|
||
"time"
|
||
|
||
"github.com/xiaoxue/memoryweave/internal/models"
|
||
)
|
||
|
||
// Reranker 重排客户端。优先模力方舟 cross-encoder API,本地为后备。
|
||
type Reranker struct {
|
||
endpoint string
|
||
apiKey string
|
||
httpClient *http.Client
|
||
embedder *Embedder // 可选;API 不可用时启用本地 cosine rerank
|
||
}
|
||
|
||
// NewReranker 创建重排客户端。embedder 可为 nil(只用 API)。
|
||
func NewReranker(endpoint string, emb *Embedder) *Reranker {
|
||
if endpoint == "" {
|
||
endpoint = os.Getenv("RERANK_ENDPOINT")
|
||
if endpoint == "" {
|
||
endpoint = "https://ai.gitee.com/v1"
|
||
}
|
||
}
|
||
// 确保路径完整
|
||
if !strings.Contains(endpoint, "/rerank") {
|
||
endpoint = strings.TrimSuffix(endpoint, "/") + "/rerank"
|
||
}
|
||
return &Reranker{
|
||
endpoint: endpoint,
|
||
apiKey: os.Getenv("MOLIFANG_API_KEY"),
|
||
httpClient: &http.Client{Timeout: 10 * time.Second},
|
||
embedder: emb, // 内置 Embedder(环境已有 VLLM_ENDPOINT)
|
||
}
|
||
}
|
||
|
||
// SetEmbedder 注入 Embedder 以启用本地 cosine rerank(API 不可用时的后备)。
|
||
func (r *Reranker) SetEmbedder(emb *Embedder) {
|
||
r.embedder = emb
|
||
}
|
||
|
||
// Rerank 对候选文档重排,返回 top_n 条最相关结果。
|
||
// 优先级:1. 模力方舟 API(cross-encoder,高精度)2. 本地 cosine(零延迟)3. 降级 score=0.5
|
||
func (r *Reranker) Rerank(query string, documents []string, topN int) ([]models.RerankResult, error) {
|
||
if len(documents) == 0 {
|
||
return nil, nil
|
||
}
|
||
|
||
// 1. 优先:模力方舟 API rerank
|
||
if r.apiKey != "" {
|
||
results, err := r.apiRerank(query, documents, topN)
|
||
if err == nil {
|
||
return results, nil
|
||
}
|
||
}
|
||
|
||
// 2. 后备:本地 cosine rerank
|
||
if r.embedder != nil {
|
||
results, err := r.localCosineRerank(query, documents, topN)
|
||
if err == nil {
|
||
return results, nil
|
||
}
|
||
}
|
||
|
||
// 3. 降级:全部 0.5(保持 recall 可用)
|
||
results := make([]models.RerankResult, len(documents))
|
||
for i := range results {
|
||
results[i] = models.RerankResult{Index: i, Score: 0.5, Text: documents[i]}
|
||
}
|
||
return results, nil
|
||
}
|
||
|
||
// localCosineRerank 使用 Embedder 编码 query → 对每个 doc 的向量做 cosine
|
||
func (r *Reranker) localCosineRerank(query string, documents []string, topN int) ([]models.RerankResult, error) {
|
||
qVec, err := r.embedder.EncodeSingle(query)
|
||
if err != nil {
|
||
return nil, fmt.Errorf("local rerank encode query: %w", err)
|
||
}
|
||
|
||
// 分批编码(本地 vLLM batch 上限 32)
|
||
const chunkSize = 32
|
||
docVecs := make([][]float32, len(documents))
|
||
for start := 0; start < len(documents); start += chunkSize {
|
||
end := start + chunkSize
|
||
if end > len(documents) {
|
||
end = len(documents)
|
||
}
|
||
chunk, err := r.embedder.Encode(documents[start:end])
|
||
if err != nil {
|
||
return nil, fmt.Errorf("local rerank encode docs[%d:%d]: %w", start, end, err)
|
||
}
|
||
copy(docVecs[start:end], chunk)
|
||
}
|
||
|
||
type scored struct {
|
||
idx int
|
||
score float64
|
||
}
|
||
scores := make([]scored, len(documents))
|
||
for i, dv := range docVecs {
|
||
scores[i] = scored{idx: i, score: cosineScore(qVec, dv)}
|
||
}
|
||
sort.Slice(scores, func(i, j int) bool {
|
||
return scores[i].score > scores[j].score
|
||
})
|
||
|
||
n := topN
|
||
if n > len(scores) {
|
||
n = len(scores)
|
||
}
|
||
out := make([]models.RerankResult, n)
|
||
for i := 0; i < n; i++ {
|
||
out[i] = models.RerankResult{
|
||
Index: scores[i].idx,
|
||
Score: scores[i].score,
|
||
Text: documents[scores[i].idx],
|
||
}
|
||
}
|
||
return out, nil
|
||
}
|
||
|
||
// cosineScore 计算两个向量的余弦相似度
|
||
func cosineScore(a, b []float32) float64 {
|
||
var dot, normA, normB float64
|
||
n := len(a)
|
||
if len(b) < n {
|
||
n = len(b)
|
||
}
|
||
for i := 0; i < n; i++ {
|
||
dot += float64(a[i]) * float64(b[i])
|
||
normA += float64(a[i]) * float64(a[i])
|
||
normB += float64(b[i]) * float64(b[i])
|
||
}
|
||
if normA == 0 || normB == 0 {
|
||
return 0
|
||
}
|
||
return dot / math.Sqrt(normA*normB)
|
||
}
|
||
|
||
// apiRerank 调用模力方舟 bge-reranker-v2-m3 进行 cross-encoder 重排
|
||
// 限制:molifang API 每批最多 25 个文档
|
||
func (r *Reranker) apiRerank(query string, documents []string, topN int) ([]models.RerankResult, error) {
|
||
// 分批(每批最多 25)
|
||
const batchSize = 25
|
||
allResults := make([]models.RerankResult, 0, len(documents))
|
||
for start := 0; start < len(documents); start += batchSize {
|
||
end := start + batchSize
|
||
if end > len(documents) {
|
||
end = len(documents)
|
||
}
|
||
batch := documents[start:end]
|
||
n := topN
|
||
if n > len(batch) {
|
||
n = len(batch)
|
||
}
|
||
reqBody := map[string]any{
|
||
"model": "bge-reranker-v2-m3",
|
||
"query": query,
|
||
"documents": batch,
|
||
"top_n": n,
|
||
}
|
||
body, _ := json.Marshal(reqBody)
|
||
req, err := http.NewRequest("POST", r.endpoint, bytes.NewReader(body))
|
||
if err != nil {
|
||
return nil, fmt.Errorf("new request: %w", err)
|
||
}
|
||
req.Header.Set("Content-Type", "application/json")
|
||
if r.apiKey != "" {
|
||
req.Header.Set("Authorization", "Bearer "+r.apiKey)
|
||
}
|
||
|
||
resp, err := r.httpClient.Do(req)
|
||
if err != nil {
|
||
return nil, fmt.Errorf("do request: %w", err)
|
||
}
|
||
defer resp.Body.Close()
|
||
|
||
if resp.StatusCode != http.StatusOK {
|
||
respBody, _ := io.ReadAll(resp.Body)
|
||
return nil, fmt.Errorf("status %d: %s", resp.StatusCode, string(respBody))
|
||
}
|
||
|
||
var result struct {
|
||
Results []struct {
|
||
Index int `json:"index"`
|
||
RelevanceScore float64 `json:"relevance_score"`
|
||
Document struct {
|
||
Text string `json:"text"`
|
||
} `json:"document"`
|
||
} `json:"results"`
|
||
}
|
||
if err := json.NewDecoder(resp.Body).Decode(&result); err != nil {
|
||
return nil, fmt.Errorf("decode: %w", err)
|
||
}
|
||
|
||
for _, rr := range result.Results {
|
||
allResults = append(allResults, models.RerankResult{
|
||
Index: start + rr.Index,
|
||
Score: rr.RelevanceScore,
|
||
Text: rr.Document.Text,
|
||
})
|
||
}
|
||
}
|
||
|
||
// 全局排序取 topN
|
||
sort.Slice(allResults, func(i, j int) bool {
|
||
return allResults[i].Score > allResults[j].Score
|
||
})
|
||
if topN > len(allResults) {
|
||
topN = len(allResults)
|
||
}
|
||
return allResults[:topN], nil
|
||
} |