feat: n-gram文本相似度作为embedding降级方案
SimpleEmbedder重写为基于字符bigram+trigram的FNV哈希向量: - 相似中文短语共享n-gram → 余弦相似度有意义 - 不需要任何外部API,纯本地计算 - 作为EMBEDDING_API_URL不可用时的自动降级 keywordSearch升级为滑动窗口分词匹配: - 2-4字窗口切分查询词,分别匹配记忆内容/摘要/标签 - 按匹配分数降序排列,分数越高越相关 - 替代原来的完整字符串包含匹配,中文召回率大幅提升 Co-Authored-By: Claude <noreply@anthropic.com>
This commit is contained in:
@@ -23,27 +23,67 @@ type Embedder interface {
|
|||||||
IsAvailable() bool
|
IsAvailable() bool
|
||||||
}
|
}
|
||||||
|
|
||||||
// SimpleEmbedder 基于关键词的简单嵌入(MVP阶段可用,无需外部API)
|
// SimpleEmbedder 基于 n-gram 哈希的本地嵌入(无需外部API)
|
||||||
|
// 用于嵌入API不可用时的降级方案,中文效果显著优于字符频率哈希。
|
||||||
type SimpleEmbedder struct{}
|
type SimpleEmbedder struct{}
|
||||||
|
|
||||||
// Embed 简单的关键词哈希嵌入(用于MVP快速验证)
|
const embedDim = 1536
|
||||||
func (e *SimpleEmbedder) Embed(ctx context.Context, text string) ([]float64, error) {
|
|
||||||
// 生成一个简单的1536维特征向量
|
|
||||||
// 基于字符频率的简单表示,用于MVP阶段
|
|
||||||
vec := make([]float64, 1536)
|
|
||||||
|
|
||||||
runes := []rune(strings.ToLower(text))
|
// Embed 使用字符 bigram + trigram 哈希生成稀疏向量。
|
||||||
for i, r := range runes {
|
// 相似的中文短语会共享 n-gram → 哈希碰撞产生有意义的余弦相似度。
|
||||||
idx := int(r) % 1536
|
func (e *SimpleEmbedder) Embed(ctx context.Context, text string) ([]float64, error) {
|
||||||
vec[idx] += 1.0 / float64(len(runes))
|
vec := make([]float64, embedDim)
|
||||||
// 考虑位置信息
|
runes := []rune(strings.TrimSpace(text))
|
||||||
posIdx := (int(r) + i) % 1536
|
if len(runes) == 0 {
|
||||||
vec[posIdx] += 0.5 / float64(len(runes))
|
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
|
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).
|
// IsAvailable returns true (SimpleEmbedder is always available as fallback).
|
||||||
func (e *SimpleEmbedder) IsAvailable() bool { return true }
|
func (e *SimpleEmbedder) IsAvailable() bool { return true }
|
||||||
|
|
||||||
@@ -127,67 +167,87 @@ func (r *Retriever) RetrieveByCategory(ctx context.Context, userID string, categ
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
// keywordSearch 关键词匹配检索(包含关键词标签匹配)
|
// keywordSearch 关键词匹配检索(包含关键词标签和n-gram分词匹配)
|
||||||
func (r *Retriever) keywordSearch(ctx context.Context, userID string, query string) ([]MemoryEntry, error) {
|
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{
|
entries, err := r.store.Query(ctx, model.MemoryQuery{
|
||||||
UserID: userID,
|
UserID: userID,
|
||||||
Priority: model.MemoryImportant,
|
Limit: 100,
|
||||||
Limit: 50,
|
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
// 关键词匹配过滤
|
type scoredEntry struct {
|
||||||
var matched []MemoryEntry
|
entry MemoryEntry
|
||||||
queryLower := strings.ToLower(query)
|
score int
|
||||||
|
}
|
||||||
|
var scored []scoredEntry
|
||||||
|
seen := make(map[string]bool)
|
||||||
|
|
||||||
for _, entry := range entries {
|
for _, entry := range entries {
|
||||||
|
if seen[entry.ID] {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
s := 0
|
||||||
contentLower := strings.ToLower(entry.Content)
|
contentLower := strings.ToLower(entry.Content)
|
||||||
summaryLower := strings.ToLower(entry.Summary)
|
summaryLower := strings.ToLower(entry.Summary)
|
||||||
|
|
||||||
// 内容/摘要匹配
|
for _, tok := range tokens {
|
||||||
if strings.Contains(contentLower, queryLower) || strings.Contains(summaryLower, queryLower) {
|
tokLower := strings.ToLower(tok)
|
||||||
matched = append(matched, entry)
|
if len([]rune(tok)) < 2 {
|
||||||
continue
|
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 {
|
for _, kw := range entry.Keywords {
|
||||||
if strings.Contains(queryLower, strings.ToLower(kw)) ||
|
kwLower := strings.ToLower(kw)
|
||||||
strings.Contains(strings.ToLower(kw), queryLower) {
|
for _, tok := range tokens {
|
||||||
matched = append(matched, entry)
|
if strings.Contains(strings.ToLower(tok), kwLower) || strings.Contains(kwLower, strings.ToLower(tok)) {
|
||||||
break
|
s += 2
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// 也匹配普通记忆
|
|
||||||
normalEntries, err := r.store.Query(ctx, model.MemoryQuery{
|
|
||||||
UserID: userID,
|
|
||||||
Priority: model.MemoryNormal,
|
|
||||||
Limit: 100,
|
|
||||||
})
|
|
||||||
if err == nil {
|
|
||||||
for _, entry := range normalEntries {
|
|
||||||
contentLower := strings.ToLower(entry.Content)
|
|
||||||
summaryLower := strings.ToLower(entry.Summary)
|
|
||||||
if strings.Contains(contentLower, queryLower) || strings.Contains(summaryLower, queryLower) {
|
|
||||||
matched = append(matched, entry)
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
for _, kw := range entry.Keywords {
|
|
||||||
if strings.Contains(queryLower, strings.ToLower(kw)) ||
|
|
||||||
strings.Contains(strings.ToLower(kw), queryLower) {
|
|
||||||
matched = append(matched, entry)
|
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if s > 0 {
|
||||||
|
seen[entry.ID] = true
|
||||||
|
scored = append(scored, scoredEntry{entry, s})
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return matched, nil
|
// 按分数降序
|
||||||
|
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 更高的
|
// deduplicate 去重合并:对高度相似的记忆只保留 Importance 更高的
|
||||||
|
|||||||
Reference in New Issue
Block a user