This repository has been archived on 2026-08-12. You can view files and clone it. You cannot open issues or pull requests or push a commit.
Files
Cyrene/backend/ai-core/internal/memory/retriever.go
T
AskaEth 02852347d9 feat: n-gram文本相似度作为embedding降级方案
SimpleEmbedder重写为基于字符bigram+trigram的FNV哈希向量:
- 相似中文短语共享n-gram → 余弦相似度有意义
- 不需要任何外部API,纯本地计算
- 作为EMBEDDING_API_URL不可用时的自动降级

keywordSearch升级为滑动窗口分词匹配:
- 2-4字窗口切分查询词,分别匹配记忆内容/摘要/标签
- 按匹配分数降序排列,分数越高越相关
- 替代原来的完整字符串包含匹配,中文召回率大幅提升

Co-Authored-By: Claude <noreply@anthropic.com>
2026-07-06 20:46:40 +08:00

304 lines
7.4 KiB
Go

package memory
import (
"context"
"fmt"
"strings"
"git.yeij.top/AskaEth/Cyrene/ai-core/internal/model"
)
// MemoryEntry 记忆条目别名(避免与model包冲突)
type MemoryEntry = model.MemoryEntry
// Retriever 记忆检索器
type Retriever struct {
store *Store
embedder Embedder // 文本转向量的接口
}
// Embedder 文本嵌入接口
type Embedder interface {
Embed(ctx context.Context, text string) ([]float64, error)
IsAvailable() bool
}
// SimpleEmbedder 基于 n-gram 哈希的本地嵌入(无需外部API)
// 用于嵌入API不可用时的降级方案,中文效果显著优于字符频率哈希。
type SimpleEmbedder struct{}
const embedDim = 1536
// Embed 使用字符 bigram + trigram 哈希生成稀疏向量。
// 相似的中文短语会共享 n-gram → 哈希碰撞产生有意义的余弦相似度。
func (e *SimpleEmbedder) Embed(ctx context.Context, text string) ([]float64, error) {
vec := make([]float64, embedDim)
runes := []rune(strings.TrimSpace(text))
if len(runes) == 0 {
return vec, nil
}
// 统计 n-gram 频率(bigram + trigram),用 TF 加权
grams := make(map[uint64]float64)
addGram := func(start, n int) {
if start+n > len(runes) {
return
}
h := hashRunes(runes[start : start+n])
grams[h] += 1.0
}
for i := 0; i < len(runes); i++ {
addGram(i, 2) // bigram
addGram(i, 3) // trigram
}
// 单字也加入,捕获关键词
for i := 0; i < len(runes); i++ {
h := hashRunes(runes[i : i+1])
grams[h] += 0.3
}
// 归一化后写入向量
var total float64
for _, v := range grams {
total += v * v
}
if total == 0 {
return vec, nil
}
norm := 1.0 / total // approximate L2 norm
for h, v := range grams {
idx := int(h % embedDim)
vec[idx] += v * norm
}
return vec, nil
}
// hashRunes computes a simple FNV-like hash of rune slice.
func hashRunes(r []rune) uint64 {
var h uint64 = 14695981039346656037
for _, c := range r {
h ^= uint64(c)
h *= 1099511628211
}
return h
}
// IsAvailable returns true (SimpleEmbedder is always available as fallback).
func (e *SimpleEmbedder) IsAvailable() bool { return true }
// NewRetriever 创建记忆检索器
func NewRetriever(store *Store, embedder Embedder) *Retriever {
if embedder == nil {
embedder = &SimpleEmbedder{}
}
return &Retriever{
store: store,
embedder: embedder,
}
}
// Retrieve 检索与查询相关的记忆
// 策略: 向量相似度 + 关键词匹配混合 → 按重要性降序返回
func (r *Retriever) Retrieve(ctx context.Context, userID string, query string) ([]MemoryEntry, error) {
var allEntries []MemoryEntry
seen := make(map[string]bool)
// 1. 向量相似度检索
embedding, err := r.embedder.Embed(ctx, query)
if err == nil {
vecEntries, err := r.store.SearchByVector(ctx, userID, embedding, 8)
if err == nil {
for _, e := range vecEntries {
if !seen[e.ID] {
seen[e.ID] = true
allEntries = append(allEntries, e)
}
}
}
}
// 2. 关键词匹配检索(包含关键词标签匹配)
keywordEntries, err := r.keywordSearch(ctx, userID, query)
if err == nil {
for _, e := range keywordEntries {
if !seen[e.ID] {
seen[e.ID] = true
allEntries = append(allEntries, e)
}
}
}
// 3. 如果没有匹配,返回最近的重要记忆
if len(allEntries) == 0 {
recentEntries, err := r.store.Query(ctx, model.MemoryQuery{
UserID: userID,
Priority: model.MemoryImportant,
Limit: 5,
})
if err == nil {
allEntries = recentEntries
}
}
// 4. 去重合并:对高度相似的记忆只保留Importance更高的
allEntries = r.deduplicate(allEntries)
// 5. 按重要性降序排列
sortByImportance(allEntries)
// 限制返回数量
if len(allEntries) > 10 {
allEntries = allEntries[:10]
}
return allEntries, nil
}
// RetrieveByCategory 按分类检索记忆
func (r *Retriever) RetrieveByCategory(ctx context.Context, userID string, category model.MemoryCategory, limit int) ([]MemoryEntry, error) {
if limit <= 0 {
limit = 20
}
return r.store.Query(ctx, model.MemoryQuery{
UserID: userID,
Category: category,
Limit: limit,
})
}
// keywordSearch 关键词匹配检索(包含关键词标签和n-gram分词匹配)
func (r *Retriever) keywordSearch(ctx context.Context, userID string, query string) ([]MemoryEntry, error) {
// 将查询切分为中文分词 tokens(2-4字的滑动窗口)
queryRunes := []rune(query)
var tokens []string
for size := 2; size <= 4; size++ {
for i := 0; i+size <= len(queryRunes); i++ {
tokens = append(tokens, string(queryRunes[i:i+size]))
}
}
// 也保留完整查询
tokens = append(tokens, query)
// 查询记忆
entries, err := r.store.Query(ctx, model.MemoryQuery{
UserID: userID,
Limit: 100,
})
if err != nil {
return nil, err
}
type scoredEntry struct {
entry MemoryEntry
score int
}
var scored []scoredEntry
seen := make(map[string]bool)
for _, entry := range entries {
if seen[entry.ID] {
continue
}
s := 0
contentLower := strings.ToLower(entry.Content)
summaryLower := strings.ToLower(entry.Summary)
for _, tok := range tokens {
tokLower := strings.ToLower(tok)
if len([]rune(tok)) < 2 {
continue // skip single-char matches (too noisy)
}
if strings.Contains(contentLower, tokLower) {
s += 2 // content match is stronger
}
if strings.Contains(summaryLower, tokLower) {
s += 3 // summary match is even stronger (distilled info)
}
}
// 关键词标签直接匹配
for _, kw := range entry.Keywords {
kwLower := strings.ToLower(kw)
for _, tok := range tokens {
if strings.Contains(strings.ToLower(tok), kwLower) || strings.Contains(kwLower, strings.ToLower(tok)) {
s += 2
break
}
}
}
if s > 0 {
seen[entry.ID] = true
scored = append(scored, scoredEntry{entry, s})
}
}
// 按分数降序
for i := 0; i < len(scored); i++ {
for j := i + 1; j < len(scored); j++ {
if scored[j].score > scored[i].score {
scored[i], scored[j] = scored[j], scored[i]
}
}
}
result := make([]MemoryEntry, 0, len(scored))
for _, s := range scored {
result = append(result, s.entry)
}
return result, nil
}
// deduplicate 去重合并:对高度相似的记忆只保留 Importance 更高的
func (r *Retriever) deduplicate(entries []MemoryEntry) []MemoryEntry {
if len(entries) < 2 {
return entries
}
result := make([]MemoryEntry, 0, len(entries))
discarded := make(map[int]bool)
for i := 0; i < len(entries); i++ {
if discarded[i] {
continue
}
for j := i + 1; j < len(entries); j++ {
if discarded[j] {
continue
}
score := entries[i].SimilarityScore(&entries[j])
if score >= deDupThreshold {
// 保留更重要的那条
if entries[j].Importance > entries[i].Importance ||
(entries[j].Importance == entries[i].Importance && entries[j].Priority > entries[i].Priority) {
discarded[i] = true
break
} else {
discarded[j] = true
}
}
}
if !discarded[i] {
result = append(result, entries[i])
}
}
return result
}
// sortByImportance 按 Importance 降序, Priority 降序排列
func sortByImportance(entries []MemoryEntry) {
for i := 0; i < len(entries); i++ {
for j := i + 1; j < len(entries); j++ {
if entries[j].Importance > entries[i].Importance ||
(entries[j].Importance == entries[i].Importance && entries[j].Priority > entries[i].Priority) {
entries[i], entries[j] = entries[j], entries[i]
}
}
}
}
// Ensure fmt is used
var _ = fmt.Sprintf