memoryweave/go/internal/storage/reranker.go

222 lines
5.8 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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 rerankAPI 不可用时的后备)。
func (r *Reranker) SetEmbedder(emb *Embedder) {
r.embedder = emb
}
// Rerank 对候选文档重排,返回 top_n 条最相关结果。
// 优先级1. 模力方舟 APIcross-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
}