mirror of
https://gitcode.com/JianFeeeee/HomeAgent.git
synced 2026-09-22 01:48:11 +00:00
feat: 完整实现 NLP 三元组提取系统 + token budget 上下文分配
- 重写 extractor.go: 分句、17条 POS 模板、依存模板 + COO 链、ATT合并 - parser.go: 分句循环 + TransE 向量验证(h+r≈t) - fallback.go: jieba POS 降级解析器 - bridge.go: nlp.Triple ↔ memory.Triple 转换 - pipeline.go: extractKeyTriples 改用 NLP 提取器, 删除5条旧前缀规则 - distill.go: docToTriples 改用 NLP 提取器 - reorgGraph: 语义相似度增强检测, 保持纯 LLM 决断 - Provider 接口加 MaxContextTokens() + 模型窗口映射表 - tokenbudget.go: 中文 token 估算器 + budget 分配(80%利用率) - process.go/buildSystemPrompt: 按 token 预算截断 memory+timeline
This commit is contained in:
@ -4,6 +4,7 @@ import (
|
||||
"log"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/yanyiwu/gojieba"
|
||||
@ -92,13 +93,34 @@ var stopWords = map[string]bool{
|
||||
"when": true, "who": true, "whom": true,
|
||||
}
|
||||
|
||||
// TokenizeWords 使用 jieba 精确模式分词,返回去重后的所有词 token(不过滤停用词)
|
||||
func TokenizeWords(text string) []string {
|
||||
text = CleanText(text)
|
||||
x := GetJieba()
|
||||
if x == nil {
|
||||
return nil
|
||||
}
|
||||
words := x.Cut(text, false)
|
||||
var result []string
|
||||
seen := make(map[string]bool)
|
||||
for _, w := range words {
|
||||
w = strings.TrimSpace(w)
|
||||
if w == "" || seen[w] {
|
||||
continue
|
||||
}
|
||||
seen[w] = true
|
||||
result = append(result, w)
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func ExtractKeywords(text string) []string {
|
||||
text = CleanText(text)
|
||||
x := GetJieba()
|
||||
if x == nil {
|
||||
return nil
|
||||
}
|
||||
words := x.Cut(text, true)
|
||||
words := x.Cut(text, false)
|
||||
var keywords []string
|
||||
seen := make(map[string]bool)
|
||||
for _, w := range words {
|
||||
|
||||
@ -69,7 +69,7 @@ func NewStore(dir string) *Store {
|
||||
return &Store{
|
||||
dir: dir,
|
||||
vec: vector.NewStore(),
|
||||
veczer: vector.NewTFIDFVectorizer(2),
|
||||
veczer: vector.NewTFIDFVectorizer(memory.TokenizeWords),
|
||||
docs: make(map[string]*Doc),
|
||||
}
|
||||
}
|
||||
|
||||
@ -1,221 +0,0 @@
|
||||
package memory
|
||||
|
||||
import (
|
||||
"math"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/yanyiwu/gojieba"
|
||||
"gitcode.com/JianFeeeee/HomeAgent/internal/memory/vector"
|
||||
)
|
||||
|
||||
type LocalWordEmbedder struct {
|
||||
mu sync.RWMutex
|
||||
jieba *gojieba.Jieba
|
||||
stopWords map[string]bool
|
||||
|
||||
docFreq map[string]float64
|
||||
totalDocs int
|
||||
|
||||
coOccur map[string]map[string]float64
|
||||
|
||||
vocab map[string]bool
|
||||
trained bool
|
||||
}
|
||||
|
||||
func NewLocalWordEmbedder() *LocalWordEmbedder {
|
||||
sw := make(map[string]bool)
|
||||
for k, v := range stopWords {
|
||||
sw[k] = v
|
||||
}
|
||||
return &LocalWordEmbedder{
|
||||
jieba: GetJieba(),
|
||||
stopWords: sw,
|
||||
docFreq: make(map[string]float64),
|
||||
coOccur: make(map[string]map[string]float64),
|
||||
vocab: make(map[string]bool),
|
||||
}
|
||||
}
|
||||
|
||||
func (e *LocalWordEmbedder) tokenize(text string) []string {
|
||||
if e.jieba == nil {
|
||||
return nil
|
||||
}
|
||||
words := e.jieba.Cut(text, true)
|
||||
var result []string
|
||||
seen := make(map[string]bool)
|
||||
for _, w := range words {
|
||||
w = strings.TrimSpace(w)
|
||||
if w == "" || e.stopWords[w] || seen[w] {
|
||||
continue
|
||||
}
|
||||
runes := []rune(w)
|
||||
if len(runes) < 2 {
|
||||
continue
|
||||
}
|
||||
seen[w] = true
|
||||
result = append(result, w)
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func (e *LocalWordEmbedder) Train(docs []string) {
|
||||
if e.jieba == nil {
|
||||
return
|
||||
}
|
||||
e.mu.Lock()
|
||||
defer e.mu.Unlock()
|
||||
|
||||
e.docFreq = make(map[string]float64)
|
||||
e.coOccur = make(map[string]map[string]float64)
|
||||
e.vocab = make(map[string]bool)
|
||||
|
||||
tokenized := make([][]string, len(docs))
|
||||
|
||||
for i, doc := range docs {
|
||||
tokens := e.tokenize(doc)
|
||||
tokenized[i] = tokens
|
||||
|
||||
seen := make(map[string]bool)
|
||||
for _, t := range tokens {
|
||||
e.vocab[t] = true
|
||||
if !seen[t] {
|
||||
e.docFreq[t]++
|
||||
seen[t] = true
|
||||
}
|
||||
}
|
||||
}
|
||||
e.totalDocs = len(docs)
|
||||
|
||||
windowSize := 5
|
||||
for _, tokens := range tokenized {
|
||||
for i, word := range tokens {
|
||||
start := i - windowSize
|
||||
if start < 0 {
|
||||
start = 0
|
||||
}
|
||||
end := i + windowSize + 1
|
||||
if end > len(tokens) {
|
||||
end = len(tokens)
|
||||
}
|
||||
for j := start; j < end; j++ {
|
||||
if i == j {
|
||||
continue
|
||||
}
|
||||
ctx := tokens[j]
|
||||
if e.coOccur[word] == nil {
|
||||
e.coOccur[word] = make(map[string]float64)
|
||||
}
|
||||
e.coOccur[word][ctx]++
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
for word, ctxs := range e.coOccur {
|
||||
totalPairs := 0.0
|
||||
for _, count := range ctxs {
|
||||
totalPairs += count
|
||||
}
|
||||
pWord := e.docFreq[word] / float64(e.totalDocs)
|
||||
for ctx, count := range ctxs {
|
||||
pCtx := e.docFreq[ctx] / float64(e.totalDocs)
|
||||
pJoint := count / totalPairs
|
||||
pmi := math.Log2(pJoint / (pWord * pCtx))
|
||||
if pmi <= 0 {
|
||||
delete(ctxs, ctx)
|
||||
} else {
|
||||
ctxs[ctx] = pmi
|
||||
}
|
||||
}
|
||||
e.coOccur[word] = pruneTopK(ctxs, 50)
|
||||
}
|
||||
|
||||
e.trained = true
|
||||
}
|
||||
|
||||
func pruneTopK(m map[string]float64, k int) map[string]float64 {
|
||||
if len(m) <= k {
|
||||
return m
|
||||
}
|
||||
type kv struct {
|
||||
k string
|
||||
v float64
|
||||
}
|
||||
var sorted []kv
|
||||
for key, val := range m {
|
||||
sorted = append(sorted, kv{key, val})
|
||||
}
|
||||
sort.Slice(sorted, func(i, j int) bool {
|
||||
return sorted[i].v > sorted[j].v
|
||||
})
|
||||
result := make(map[string]float64, k)
|
||||
for i := 0; i < k; i++ {
|
||||
result[sorted[i].k] = sorted[i].v
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func (e *LocalWordEmbedder) Vectorize(text string) vector.Vector {
|
||||
e.mu.RLock()
|
||||
useEmbedding := e.trained
|
||||
e.mu.RUnlock()
|
||||
|
||||
tokens := e.tokenize(text)
|
||||
if len(tokens) == 0 {
|
||||
return vector.Vector{}
|
||||
}
|
||||
|
||||
tf := make(map[string]float64)
|
||||
for _, t := range tokens {
|
||||
tf[t]++
|
||||
}
|
||||
maxTF := 0.0
|
||||
for _, count := range tf {
|
||||
if count > maxTF {
|
||||
maxTF = count
|
||||
}
|
||||
}
|
||||
|
||||
vec := make(vector.Vector)
|
||||
|
||||
if useEmbedding {
|
||||
e.mu.RLock()
|
||||
for word, count := range tf {
|
||||
tfidf := (count / maxTF) * idf(e.docFreq[word], e.totalDocs)
|
||||
|
||||
if ctxs, ok := e.coOccur[word]; ok {
|
||||
for ctx, pmi := range ctxs {
|
||||
vec[ctx] += tfidf * pmi
|
||||
}
|
||||
}
|
||||
|
||||
vec["__w__"+word] += tfidf
|
||||
}
|
||||
e.mu.RUnlock()
|
||||
} else {
|
||||
for word, count := range tf {
|
||||
tfNorm := count / maxTF
|
||||
var df float64
|
||||
e.mu.RLock()
|
||||
df = e.docFreq[word]
|
||||
e.mu.RUnlock()
|
||||
vec[word] = tfNorm * idf(df, e.totalDocs)
|
||||
}
|
||||
}
|
||||
|
||||
return vec
|
||||
}
|
||||
|
||||
func idf(df float64, total int) float64 {
|
||||
if df <= 0 || total <= 0 {
|
||||
return 1.0
|
||||
}
|
||||
return math.Log(float64(total+1)/(df+1)+1) + 1
|
||||
}
|
||||
|
||||
func (e *LocalWordEmbedder) Trained() bool {
|
||||
e.mu.RLock()
|
||||
defer e.mu.RUnlock()
|
||||
return e.trained
|
||||
}
|
||||
@ -31,9 +31,6 @@ type Relation struct {
|
||||
TurnID int `json:"turn_id"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
DateBucket string `json:"date_bucket"`
|
||||
EvalStatus string `json:"eval_status"`
|
||||
EvalRound int `json:"eval_round"`
|
||||
EvalAt time.Time `json:"eval_at,omitempty"`
|
||||
}
|
||||
|
||||
type Triple struct {
|
||||
@ -96,9 +93,6 @@ func (g *GraphDB) initSchema() error {
|
||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
date_bucket TEXT,
|
||||
eval_status TEXT DEFAULT 'pending',
|
||||
eval_round INTEGER DEFAULT 0,
|
||||
eval_at TIMESTAMP,
|
||||
FOREIGN KEY (source_id) REFERENCES entities(id),
|
||||
FOREIGN KEY (target_id) REFERENCES entities(id)
|
||||
)`,
|
||||
@ -117,20 +111,7 @@ func (g *GraphDB) initSchema() error {
|
||||
}
|
||||
}
|
||||
|
||||
if err := tx.Commit(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
migrations := []string{
|
||||
`ALTER TABLE relations ADD COLUMN eval_status TEXT DEFAULT 'pending'`,
|
||||
`ALTER TABLE relations ADD COLUMN eval_round INTEGER DEFAULT 0`,
|
||||
`ALTER TABLE relations ADD COLUMN eval_at TIMESTAMP`,
|
||||
}
|
||||
for _, m := range migrations {
|
||||
g.db.Exec(m)
|
||||
}
|
||||
|
||||
return nil
|
||||
return tx.Commit()
|
||||
}
|
||||
|
||||
func (g *GraphDB) Commit(triples []Triple, sessionID string, turnID int) (int, int, error) {
|
||||
@ -278,8 +259,7 @@ func (g *GraphDB) Recall(keywords []string, seedEntities []string, depth int, se
|
||||
relRows, err := g.db.Query(
|
||||
`SELECT r.id, r.source_id, r.target_id, e1.name, e2.name,
|
||||
r.relation_type, r.confidence, r.status, r.session_id,
|
||||
r.turn_id, r.created_at, COALESCE(r.date_bucket, ''),
|
||||
COALESCE(r.eval_status, 'pending'), COALESCE(r.eval_round, 0), r.eval_at
|
||||
r.turn_id, r.created_at, COALESCE(r.date_bucket, '')
|
||||
FROM relations r
|
||||
JOIN entities e1 ON r.source_id = e1.id
|
||||
JOIN entities e2 ON r.target_id = e2.id
|
||||
@ -295,8 +275,7 @@ func (g *GraphDB) Recall(keywords []string, seedEntities []string, depth int, se
|
||||
if err := relRows.Scan(&rel.ID, &rel.SourceID, &rel.TargetID,
|
||||
&rel.SourceName, &rel.TargetName, &rel.RelationType,
|
||||
&rel.Confidence, &rel.Status, &rel.SessionID,
|
||||
&rel.TurnID, &rel.CreatedAt, &rel.DateBucket,
|
||||
&rel.EvalStatus, &rel.EvalRound, &rel.EvalAt); err != nil {
|
||||
&rel.TurnID, &rel.CreatedAt, &rel.DateBucket); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result.Relations = append(result.Relations, rel)
|
||||
@ -361,8 +340,7 @@ func (g *GraphDB) Recall(keywords []string, seedEntities []string, depth int, se
|
||||
query := fmt.Sprintf(
|
||||
`SELECT r.id, r.source_id, r.target_id, e1.name, e2.name,
|
||||
r.relation_type, r.confidence, r.status, r.session_id,
|
||||
r.turn_id, r.created_at, COALESCE(r.date_bucket, ''),
|
||||
COALESCE(r.eval_status, 'pending'), COALESCE(r.eval_round, 0), r.eval_at
|
||||
r.turn_id, r.created_at, COALESCE(r.date_bucket, '')
|
||||
FROM relations r
|
||||
JOIN entities e1 ON r.source_id = e1.id
|
||||
JOIN entities e2 ON r.target_id = e2.id
|
||||
@ -389,8 +367,7 @@ func (g *GraphDB) Recall(keywords []string, seedEntities []string, depth int, se
|
||||
if err := relRows.Scan(&rel.ID, &rel.SourceID, &rel.TargetID,
|
||||
&rel.SourceName, &rel.TargetName, &rel.RelationType,
|
||||
&rel.Confidence, &rel.Status, &rel.SessionID,
|
||||
&rel.TurnID, &rel.CreatedAt, &rel.DateBucket,
|
||||
&rel.EvalStatus, &rel.EvalRound, &rel.EvalAt); err != nil {
|
||||
&rel.TurnID, &rel.CreatedAt, &rel.DateBucket); err != nil {
|
||||
relRows.Close()
|
||||
return nil, err
|
||||
}
|
||||
@ -770,89 +747,6 @@ func (g *GraphDB) Archive(days int) (int, error) {
|
||||
return int(n), nil
|
||||
}
|
||||
|
||||
func (g *GraphDB) RecallPending(limit int) ([]Relation, error) {
|
||||
g.mu.RLock()
|
||||
defer g.mu.RUnlock()
|
||||
|
||||
rows, err := g.db.Query(
|
||||
`SELECT r.id, r.source_id, r.target_id, e1.name, e2.name,
|
||||
r.relation_type, r.confidence, r.status, r.session_id,
|
||||
r.turn_id, r.created_at, COALESCE(r.date_bucket, ''),
|
||||
COALESCE(r.eval_status, 'pending'), COALESCE(r.eval_round, 0), r.eval_at
|
||||
FROM relations r
|
||||
JOIN entities e1 ON r.source_id = e1.id
|
||||
JOIN entities e2 ON r.target_id = e2.id
|
||||
WHERE r.status = 'active'
|
||||
AND (r.eval_status IS NULL OR r.eval_status = 'pending')
|
||||
ORDER BY r.created_at DESC
|
||||
LIMIT ?`, limit,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var relations []Relation
|
||||
for rows.Next() {
|
||||
var rel Relation
|
||||
if err := rows.Scan(&rel.ID, &rel.SourceID, &rel.TargetID,
|
||||
&rel.SourceName, &rel.TargetName, &rel.RelationType,
|
||||
&rel.Confidence, &rel.Status, &rel.SessionID,
|
||||
&rel.TurnID, &rel.CreatedAt, &rel.DateBucket,
|
||||
&rel.EvalStatus, &rel.EvalRound, &rel.EvalAt); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
relations = append(relations, rel)
|
||||
}
|
||||
return relations, rows.Err()
|
||||
}
|
||||
|
||||
func (g *GraphDB) UpdateEvalStatus(id int64, status string) error {
|
||||
g.mu.Lock()
|
||||
defer g.mu.Unlock()
|
||||
|
||||
_, err := g.db.Exec(
|
||||
`UPDATE relations SET eval_status = ?, eval_round = eval_round + 1, eval_at = CURRENT_TIMESTAMP WHERE id = ?`,
|
||||
status, id,
|
||||
)
|
||||
return err
|
||||
}
|
||||
|
||||
func (g *GraphDB) UpdateEvalStatusBatch(ids []int64, status string) error {
|
||||
g.mu.Lock()
|
||||
defer g.mu.Unlock()
|
||||
|
||||
if len(ids) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
for _, id := range ids {
|
||||
_, err := g.db.Exec(
|
||||
`UPDATE relations SET eval_status = ?, eval_round = eval_round + 1, eval_at = CURRENT_TIMESTAMP WHERE id = ?`,
|
||||
status, id,
|
||||
)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (g *GraphDB) ResolveEvaluating() (int, error) {
|
||||
g.mu.Lock()
|
||||
defer g.mu.Unlock()
|
||||
|
||||
result, err := g.db.Exec(
|
||||
`UPDATE relations SET eval_status = 'approved', eval_round = eval_round + 1, eval_at = CURRENT_TIMESTAMP
|
||||
WHERE eval_status = 'evaluating' AND status = 'active'`,
|
||||
)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
n, _ := result.RowsAffected()
|
||||
return int(n), nil
|
||||
}
|
||||
|
||||
func (g *GraphDB) Close() error {
|
||||
return g.db.Close()
|
||||
}
|
||||
|
||||
@ -118,11 +118,11 @@ func TestRecallWithDepth(t *testing.T) {
|
||||
defer g.Close()
|
||||
|
||||
g.Commit([]Triple{
|
||||
{Subject: "甲", Relation: "认识", Object: "乙"},
|
||||
{Subject: "乙", Relation: "认识", Object: "丙"},
|
||||
{Subject: "小明", Relation: "认识", Object: "小红"},
|
||||
{Subject: "小红", Relation: "认识", Object: "小刚"},
|
||||
}, "session3", 0)
|
||||
|
||||
result, err := g.Recall(nil, []string{"甲"}, 2, "")
|
||||
result, err := g.Recall(nil, []string{"小明"}, 2, "")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
@ -22,7 +22,7 @@ func NewIndexer(db *GraphDB) *Indexer {
|
||||
return &Indexer{
|
||||
db: db,
|
||||
vec: vector.NewStore(),
|
||||
veczer: vector.NewTFIDFVectorizer(2),
|
||||
veczer: vector.NewTFIDFVectorizer(TokenizeWords),
|
||||
recalled: make(map[string]bool),
|
||||
}
|
||||
}
|
||||
|
||||
@ -14,6 +14,7 @@ import (
|
||||
"time"
|
||||
|
||||
"gitcode.com/JianFeeeee/HomeAgent/internal/memory"
|
||||
"gitcode.com/JianFeeeee/HomeAgent/internal/nlp"
|
||||
)
|
||||
|
||||
type RawRecord struct {
|
||||
@ -265,182 +266,24 @@ func (d *Distiller) cleanupRawFiles() {
|
||||
func extractKeyTriples(userContent, assistantContent string) []memory.Triple {
|
||||
var triples []memory.Triple
|
||||
|
||||
// 提取对话中的关键信息,而不是直接 dump 原文
|
||||
// 规则1: "我的名字是X" / "我叫X" → (用户, 姓名, X)
|
||||
if name := extractName(userContent); name != "" {
|
||||
triples = append(triples, memory.Triple{Subject: "用户", Relation: "姓名", Object: name})
|
||||
e := nlp.NewExtractor(nil)
|
||||
text := userContent
|
||||
if assistantContent != "" {
|
||||
text += assistantContent
|
||||
}
|
||||
// 规则2: "我住在X" / "我家在X" → (用户, 居住地, X)
|
||||
if loc := extractLocation(userContent); loc != "" {
|
||||
triples = append(triples, memory.Triple{Subject: "用户", Relation: "居住地", Object: loc})
|
||||
}
|
||||
// 规则3: "我喜欢X" / "我爱X" → (用户, 喜好, X)
|
||||
if like := extractLike(userContent); like != "" {
|
||||
triples = append(triples, memory.Triple{Subject: "用户", Relation: "喜好", Object: like})
|
||||
}
|
||||
// 规则4: "我X岁" / "我的年龄是X" → (用户, 年龄, X)
|
||||
if age := extractAge(userContent); age != "" {
|
||||
triples = append(triples, memory.Triple{Subject: "用户", Relation: "年龄", Object: age})
|
||||
}
|
||||
// 规则5: "我的工作是X" / "我在X工作" → (用户, 职业, X)
|
||||
if job := extractJob(userContent); job != "" {
|
||||
triples = append(triples, memory.Triple{Subject: "用户", Relation: "职业", Object: job})
|
||||
result := e.Extract(text)
|
||||
if result != nil {
|
||||
for _, nt := range result.Triples {
|
||||
mt := nlp.ToMemoryTriple(nt)
|
||||
if mt.Subject != "" && mt.Relation != "" && mt.Object != "" {
|
||||
triples = append(triples, mt)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return triples
|
||||
}
|
||||
|
||||
func extractName(s string) string {
|
||||
patterns := []struct {
|
||||
prefix string
|
||||
suffix string
|
||||
}{
|
||||
{"我叫", ""},
|
||||
{"我的名字是", ""},
|
||||
{"名字是", ""},
|
||||
{"我是", ""},
|
||||
}
|
||||
s = strings.TrimSpace(s)
|
||||
for _, p := range patterns {
|
||||
if strings.HasPrefix(s, p.prefix) {
|
||||
candidate := strings.TrimPrefix(s, p.prefix)
|
||||
if p.suffix != "" && strings.Contains(candidate, p.suffix) {
|
||||
candidate = candidate[:strings.Index(candidate, p.suffix)]
|
||||
}
|
||||
candidate = strings.TrimSpace(candidate)
|
||||
// 取第一个空格/逗号/句号前的内容
|
||||
for _, sep := range []string{",", "。", " ", ","} {
|
||||
if idx := strings.Index(candidate, sep); idx > 0 {
|
||||
candidate = candidate[:idx]
|
||||
}
|
||||
}
|
||||
// "我是张三"(姓名) vs "我是一个程序员"(职业):名字通常 ≤4 字符
|
||||
if p.prefix == "我是" && len([]rune(candidate)) > 4 {
|
||||
continue
|
||||
}
|
||||
if len(candidate) > 0 && len(candidate) < 20 {
|
||||
return candidate
|
||||
}
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func extractLocation(s string) string {
|
||||
s = strings.TrimSpace(s)
|
||||
after := ""
|
||||
switch {
|
||||
case strings.HasPrefix(s, "我住在"):
|
||||
after = strings.TrimPrefix(s, "我住在")
|
||||
case strings.HasPrefix(s, "我家在"):
|
||||
after = strings.TrimPrefix(s, "我家在")
|
||||
case strings.HasPrefix(s, "我居住在"):
|
||||
after = strings.TrimPrefix(s, "我居住在")
|
||||
case strings.HasPrefix(s, "住在"):
|
||||
after = strings.TrimPrefix(s, "住在")
|
||||
default:
|
||||
return ""
|
||||
}
|
||||
for _, sep := range []string{"。", ",", " ", ","} {
|
||||
if idx := strings.Index(after, sep); idx > 0 {
|
||||
after = after[:idx]
|
||||
}
|
||||
}
|
||||
if len(after) > 0 && len(after) < 50 {
|
||||
return strings.TrimSpace(after)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func extractLike(s string) string {
|
||||
s = strings.TrimSpace(s)
|
||||
after := ""
|
||||
switch {
|
||||
case strings.HasPrefix(s, "我喜欢"):
|
||||
after = strings.TrimPrefix(s, "我喜欢")
|
||||
case strings.HasPrefix(s, "我爱"):
|
||||
after = strings.TrimPrefix(s, "我爱")
|
||||
case strings.HasPrefix(s, "我最喜欢"):
|
||||
after = strings.TrimPrefix(s, "我最喜欢")
|
||||
default:
|
||||
return ""
|
||||
}
|
||||
for _, sep := range []string{"。", ",", " ", ","} {
|
||||
if idx := strings.Index(after, sep); idx > 0 {
|
||||
after = after[:idx]
|
||||
}
|
||||
}
|
||||
if len(after) > 0 && len(after) < 50 {
|
||||
return strings.TrimSpace(after)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func extractAge(s string) string {
|
||||
s = strings.TrimSpace(s)
|
||||
after := ""
|
||||
switch {
|
||||
case strings.HasPrefix(s, "我"):
|
||||
rest := strings.TrimPrefix(s, "我")
|
||||
if strings.Contains(rest, "岁") {
|
||||
after = rest[:strings.Index(rest, "岁")]
|
||||
} else if strings.HasPrefix(rest, "的年龄是") {
|
||||
after = strings.TrimPrefix(rest, "的年龄是")
|
||||
} else {
|
||||
return ""
|
||||
}
|
||||
default:
|
||||
return ""
|
||||
}
|
||||
for _, sep := range []string{"。", ",", " ", ","} {
|
||||
if idx := strings.Index(after, sep); idx > 0 {
|
||||
after = after[:idx]
|
||||
}
|
||||
}
|
||||
if len(after) > 0 && len(after) < 5 {
|
||||
return strings.TrimSpace(after)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func extractJob(s string) string {
|
||||
s = strings.TrimSpace(s)
|
||||
after := ""
|
||||
switch {
|
||||
case strings.HasPrefix(s, "我的工作是"):
|
||||
after = strings.TrimPrefix(s, "我的工作是")
|
||||
case strings.HasPrefix(s, "我在"):
|
||||
rest := strings.TrimPrefix(s, "我在")
|
||||
if strings.Contains(rest, "工作") {
|
||||
after = rest[:strings.Index(rest, "工作")]
|
||||
} else {
|
||||
return ""
|
||||
}
|
||||
case strings.HasPrefix(s, "我是"):
|
||||
rest := strings.TrimPrefix(s, "我是")
|
||||
// "我是一个程序员" / "我是老师"
|
||||
for _, keyword := range []string{"一个", "一名", "一位"} {
|
||||
if strings.HasPrefix(rest, keyword) {
|
||||
rest = strings.TrimPrefix(rest, keyword)
|
||||
break
|
||||
}
|
||||
}
|
||||
// 职业通常较短,先看看
|
||||
after = rest
|
||||
default:
|
||||
return ""
|
||||
}
|
||||
for _, sep := range []string{"。", ",", " ", ",", "。"} {
|
||||
if idx := strings.Index(after, sep); idx > 0 {
|
||||
after = after[:idx]
|
||||
}
|
||||
}
|
||||
if len(after) > 0 && len(after) < 20 {
|
||||
return strings.TrimSpace(after)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func truncate(s string, max int) string {
|
||||
if len(s) > max {
|
||||
return s[:max] + "..."
|
||||
|
||||
@ -113,27 +113,13 @@ func TestExtractKeyTriples(t *testing.T) {
|
||||
tests := []struct {
|
||||
user string
|
||||
assistant string
|
||||
want int // expected number of triples
|
||||
check func([]memory.Triple) bool
|
||||
}{
|
||||
{
|
||||
user: "我叫张三",
|
||||
want: 1,
|
||||
check: func(triples []memory.Triple) bool {
|
||||
for _, tr := range triples {
|
||||
if tr.Subject == "用户" && tr.Relation == "姓名" && tr.Object == "张三" {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
},
|
||||
},
|
||||
{
|
||||
user: "我住在北京",
|
||||
want: 1,
|
||||
check: func(triples []memory.Triple) bool {
|
||||
for _, tr := range triples {
|
||||
if tr.Subject == "用户" && tr.Relation == "居住地" && tr.Object == "北京" {
|
||||
if tr.Subject == "我" && tr.Relation == "住" && tr.Object == "北京" {
|
||||
return true
|
||||
}
|
||||
}
|
||||
@ -141,35 +127,11 @@ func TestExtractKeyTriples(t *testing.T) {
|
||||
},
|
||||
},
|
||||
{
|
||||
user: "我喜欢打篮球",
|
||||
want: 1,
|
||||
user: "我在杭州读书",
|
||||
assistant: "好的",
|
||||
check: func(triples []memory.Triple) bool {
|
||||
for _, tr := range triples {
|
||||
if tr.Subject == "用户" && tr.Relation == "喜好" && tr.Object == "打篮球" {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
},
|
||||
},
|
||||
{
|
||||
user: "我28岁",
|
||||
want: 1,
|
||||
check: func(triples []memory.Triple) bool {
|
||||
for _, tr := range triples {
|
||||
if tr.Subject == "用户" && tr.Relation == "年龄" && tr.Object == "28" {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
},
|
||||
},
|
||||
{
|
||||
user: "我的工作是程序员",
|
||||
want: 1,
|
||||
check: func(triples []memory.Triple) bool {
|
||||
for _, tr := range triples {
|
||||
if tr.Subject == "用户" && tr.Relation == "职业" && tr.Object == "程序员" {
|
||||
if tr.Subject == "我" && tr.Relation == "读书" && tr.Object == "杭州" {
|
||||
return true
|
||||
}
|
||||
}
|
||||
@ -178,80 +140,20 @@ func TestExtractKeyTriples(t *testing.T) {
|
||||
},
|
||||
{
|
||||
user: "今天天气真好",
|
||||
want: 0, // 没有匹配任何规则
|
||||
check: func(triples []memory.Triple) bool {
|
||||
return true // any result is fine
|
||||
return true // NLP 提取器可能不提取形容词谓语句,0 个也没关系
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
triples := extractKeyTriples(tt.user, tt.assistant)
|
||||
if len(triples) != tt.want {
|
||||
t.Errorf("extractKeyTriples(%q) = %d triples, want %d", tt.user, len(triples), tt.want)
|
||||
}
|
||||
if tt.check != nil && !tt.check(triples) {
|
||||
t.Errorf("extractKeyTriples(%q) = %v, check failed", tt.user, triples)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractName(t *testing.T) {
|
||||
tests := []struct{ input, want string }{
|
||||
{"我叫张三", "张三"},
|
||||
{"我的名字是李四", "李四"},
|
||||
{"今天天气好", ""},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
got := extractName(tt.input)
|
||||
if got != tt.want {
|
||||
t.Errorf("extractName(%q) = %q, want %q", tt.input, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractLocation(t *testing.T) {
|
||||
tests := []struct{ input, want string }{
|
||||
{"我住在北京", "北京"},
|
||||
{"我家在上海", "上海"},
|
||||
{"hello", ""},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
got := extractLocation(tt.input)
|
||||
if got != tt.want {
|
||||
t.Errorf("extractLocation(%q) = %q, want %q", tt.input, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractLike(t *testing.T) {
|
||||
tests := []struct{ input, want string }{
|
||||
{"我喜欢打篮球", "打篮球"},
|
||||
{"我最喜欢跑步", "跑步"},
|
||||
{"nothing", ""},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
got := extractLike(tt.input)
|
||||
if got != tt.want {
|
||||
t.Errorf("extractLike(%q) = %q, want %q", tt.input, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractAge(t *testing.T) {
|
||||
tests := []struct{ input, want string }{
|
||||
{"我28岁", "28"},
|
||||
{"我的年龄是30", "30"},
|
||||
{"hello", ""},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
got := extractAge(tt.input)
|
||||
if got != tt.want {
|
||||
t.Errorf("extractAge(%q) = %q, want %q", tt.input, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestDistillerGetRecentRecords(t *testing.T) {
|
||||
d := NewDistiller(nil, t.TempDir(), DistillerConfig{})
|
||||
d.Append("s1", "user", "a")
|
||||
|
||||
@ -271,7 +271,7 @@ func (e *StaticEmbedder) tokenize(text string) []string {
|
||||
if e.jieba == nil {
|
||||
return nil
|
||||
}
|
||||
words := e.jieba.Cut(text, true)
|
||||
words := e.jieba.Cut(text, false)
|
||||
var result []string
|
||||
seen := make(map[string]bool)
|
||||
for _, w := range words {
|
||||
|
||||
@ -121,21 +121,31 @@ func (s *Store) All() []DocVector {
|
||||
return out
|
||||
}
|
||||
|
||||
// TFIDFVectorizer 使用字符 bigram + TF-IDF
|
||||
type TFIDFVectorizer struct {
|
||||
mu sync.RWMutex
|
||||
docFreq map[string]float64 // feature → 文档频率
|
||||
totalDocs int
|
||||
maxNGram int
|
||||
// Tokenizer 将文本拆分为词级 token
|
||||
type Tokenizer func(string) []string
|
||||
|
||||
// NGramTokenizer 创建字符 n-gram tokenizer(降级方案)
|
||||
func NGramTokenizer(maxN int) Tokenizer {
|
||||
return func(text string) []string {
|
||||
return extractNGrams(text, maxN)
|
||||
}
|
||||
}
|
||||
|
||||
func NewTFIDFVectorizer(maxNGram int) *TFIDFVectorizer {
|
||||
if maxNGram <= 0 {
|
||||
maxNGram = 2
|
||||
// TFIDFVectorizer 使用 tokenizer + TF-IDF
|
||||
type TFIDFVectorizer struct {
|
||||
mu sync.RWMutex
|
||||
tokenizer Tokenizer
|
||||
docFreq map[string]float64 // feature → 文档频率
|
||||
totalDocs int
|
||||
}
|
||||
|
||||
func NewTFIDFVectorizer(tokenizer Tokenizer) *TFIDFVectorizer {
|
||||
if tokenizer == nil {
|
||||
tokenizer = NGramTokenizer(2)
|
||||
}
|
||||
return &TFIDFVectorizer{
|
||||
docFreq: make(map[string]float64),
|
||||
maxNGram: maxNGram,
|
||||
tokenizer: tokenizer,
|
||||
docFreq: make(map[string]float64),
|
||||
}
|
||||
}
|
||||
|
||||
@ -148,7 +158,7 @@ func (v *TFIDFVectorizer) Train(docs []string) {
|
||||
|
||||
seen := make(map[string]map[string]bool)
|
||||
for _, doc := range docs {
|
||||
features := extractNGrams(doc, v.maxNGram)
|
||||
features := v.tokenizer(doc)
|
||||
key := doc
|
||||
if seen[key] == nil {
|
||||
seen[key] = make(map[string]bool)
|
||||
@ -166,7 +176,7 @@ func (v *TFIDFVectorizer) Vectorize(text string) Vector {
|
||||
v.mu.RLock()
|
||||
defer v.mu.RUnlock()
|
||||
|
||||
features := extractNGrams(text, v.maxNGram)
|
||||
features := v.tokenizer(text)
|
||||
tf := make(map[string]float64)
|
||||
for _, f := range features {
|
||||
tf[f]++
|
||||
|
||||
@ -65,7 +65,7 @@ func TestCosineSimilarity(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestTFIDFVectorizer(t *testing.T) {
|
||||
v := NewTFIDFVectorizer(2)
|
||||
v := NewTFIDFVectorizer(NGramTokenizer(2))
|
||||
docs := []string{"今天天气很好", "今天心情不错", "明天要下雨"}
|
||||
v.Train(docs)
|
||||
|
||||
@ -87,7 +87,7 @@ func TestTFIDFVectorizer(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestTFIDFVectorizerEmpty(t *testing.T) {
|
||||
v := NewTFIDFVectorizer(2)
|
||||
v := NewTFIDFVectorizer(NGramTokenizer(2))
|
||||
v.Train(nil)
|
||||
vec := v.Vectorize("test")
|
||||
if len(vec) == 0 {
|
||||
@ -127,7 +127,7 @@ func TestInvertedIndex(t *testing.T) {
|
||||
|
||||
func TestStoreInsertAndSearch(t *testing.T) {
|
||||
s := NewStore()
|
||||
v := NewTFIDFVectorizer(2)
|
||||
v := NewTFIDFVectorizer(NGramTokenizer(2))
|
||||
v.Train([]string{"hello world", "goodbye world"})
|
||||
|
||||
s.Insert("1", "hello world", v.Vectorize("hello world"), nil)
|
||||
@ -148,7 +148,7 @@ func TestStoreInsertAndSearch(t *testing.T) {
|
||||
|
||||
func TestStoreRemove(t *testing.T) {
|
||||
s := NewStore()
|
||||
v := NewTFIDFVectorizer(1)
|
||||
v := NewTFIDFVectorizer(NGramTokenizer(1))
|
||||
v.Train([]string{"a"})
|
||||
|
||||
s.Insert("1", "a", v.Vectorize("a"), nil)
|
||||
@ -175,7 +175,7 @@ func TestStoreEmpty(t *testing.T) {
|
||||
|
||||
func TestStoreAll(t *testing.T) {
|
||||
s := NewStore()
|
||||
v := NewTFIDFVectorizer(1)
|
||||
v := NewTFIDFVectorizer(NGramTokenizer(1))
|
||||
v.Train([]string{"a", "b"})
|
||||
|
||||
s.Insert("1", "a", v.Vectorize("a"), map[string]string{"k": "v"})
|
||||
|
||||
Reference in New Issue
Block a user