mirror of
https://gitcode.com/JianFeeeee/HomeAgent.git
synced 2026-10-04 00:03:59 +00:00
蒸馏嵌入接线:Distiller/Agent 注入共享 embedder,修复配置缺失时蒸馏零产出
This commit is contained in:
@ -136,6 +136,7 @@ type AgentConfig struct {
|
||||
MaxContextSize int // 活跃上下文最大条数,超出按相关性裁剪
|
||||
ContextSavePath string // 上下文持久化路径,空则不持久化
|
||||
EmbeddingModelPath string // 预训练词嵌入模型路径(word2vec 文本格式),空则不使用
|
||||
Embedder *memory.StaticEmbedder // 共享词嵌入实例;nil 时按 EmbeddingModelPath 自建
|
||||
StageHost *StageHost
|
||||
EventBus *events.Bus
|
||||
ThinkingEnabled bool
|
||||
@ -161,7 +162,10 @@ func New(cfg AgentConfig) *Agent {
|
||||
cfg.MaxContextSize = 30
|
||||
}
|
||||
|
||||
embedder := memory.NewStaticEmbedder(strings.Split(cfg.EmbeddingModelPath, ",")...)
|
||||
embedder := cfg.Embedder
|
||||
if embedder == nil {
|
||||
embedder = memory.NewStaticEmbedder(strings.Split(cfg.EmbeddingModelPath, ",")...)
|
||||
}
|
||||
if cfg.DocStore != nil {
|
||||
cfg.DocStore.SetVectorizer(embedder)
|
||||
cfg.DocStore.ReindexWithVectorizer(embedder)
|
||||
|
||||
@ -32,7 +32,7 @@ func TestDocToTriples(t *testing.T) {
|
||||
Source: "context",
|
||||
}
|
||||
|
||||
triples := docToTriples(doc)
|
||||
triples := docToTriples(doc, nil)
|
||||
|
||||
foundSummary := false
|
||||
foundSource := false
|
||||
@ -54,7 +54,7 @@ func TestDocToTriples(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestDocToTriplesNil(t *testing.T) {
|
||||
triples := docToTriples(nil)
|
||||
triples := docToTriples(nil, nil)
|
||||
if len(triples) != 0 {
|
||||
t.Errorf("expected empty for nil doc, got %d", len(triples))
|
||||
}
|
||||
@ -65,7 +65,7 @@ func TestDocToTriplesNoSource(t *testing.T) {
|
||||
Summary: "无来源文档",
|
||||
Content: "content",
|
||||
}
|
||||
triples := docToTriples(doc)
|
||||
triples := docToTriples(doc, nil)
|
||||
for _, tr := range triples {
|
||||
if tr.Relation == "来源" {
|
||||
t.Error("should not have source triple when Source is empty")
|
||||
@ -80,7 +80,7 @@ func TestDocToTriplesTypes(t *testing.T) {
|
||||
Source: "test",
|
||||
}
|
||||
|
||||
triples := docToTriples(doc)
|
||||
triples := docToTriples(doc, nil)
|
||||
|
||||
for _, tr := range triples {
|
||||
if tr.Subject == "文档" {
|
||||
|
||||
@ -17,7 +17,7 @@ func TestDocToTriplesEmpty(t *testing.T) {
|
||||
Content: "",
|
||||
Source: "test",
|
||||
}
|
||||
triples := docToTriples(doc)
|
||||
triples := docToTriples(doc, nil)
|
||||
if len(triples) < 2 {
|
||||
t.Fatalf("expected at least 2 triples (主题+来源), got %d", len(triples))
|
||||
}
|
||||
@ -38,7 +38,7 @@ func TestDocToTriplesConversation(t *testing.T) {
|
||||
Content: "[15:04] qq: 今天天气怎么样\n[15:05] agent: 今天天气很好",
|
||||
Source: "qq",
|
||||
}
|
||||
triples := docToTriples(doc)
|
||||
triples := docToTriples(doc, nil)
|
||||
|
||||
if len(triples) < 2 {
|
||||
t.Errorf("expected at least 2 triples (主题+来源), got %d", len(triples))
|
||||
@ -60,7 +60,7 @@ func TestDocToTriplesMultiLine(t *testing.T) {
|
||||
Content: "[10:00] user: 你好\n[10:01] agent: 你好,有什么可以帮助你的\n[10:02] user: 今天天气如何\n[10:03] agent: 今天天气很好",
|
||||
Source: "qq",
|
||||
}
|
||||
triples := docToTriples(doc)
|
||||
triples := docToTriples(doc, nil)
|
||||
if len(triples) < 2 {
|
||||
t.Fatalf("expected at least 2 triples, got %d", len(triples))
|
||||
}
|
||||
@ -80,7 +80,7 @@ func TestDocToTriplesEmptyContent(t *testing.T) {
|
||||
Content: "",
|
||||
Source: "test",
|
||||
}
|
||||
triples := docToTriples(doc)
|
||||
triples := docToTriples(doc, nil)
|
||||
if len(triples) != 2 {
|
||||
t.Fatalf("expected exactly 2 triples (主题+来源) for empty content, got %d", len(triples))
|
||||
}
|
||||
|
||||
@ -178,7 +178,7 @@ func (a *Agent) archiveColdDocs() {
|
||||
if a.docStore != nil {
|
||||
coldDocs := a.docStore.FindColdDocs(72*time.Hour, 2)
|
||||
for _, doc := range coldDocs {
|
||||
triples := docToTriples(doc)
|
||||
triples := docToTriples(doc, a.embedder)
|
||||
if len(triples) > 0 {
|
||||
ec, rc, err := a.memory.Commit(triples, string(a.id)+"_doc_archival", 0)
|
||||
if err != nil {
|
||||
@ -387,7 +387,7 @@ func entitySimilarity(a, b string) float64 {
|
||||
return float64(intersect) / float64(union)
|
||||
}
|
||||
|
||||
func docToTriples(doc *document.Doc) []memory.Triple {
|
||||
func docToTriples(doc *document.Doc, embedder nlp.Vectorizer) []memory.Triple {
|
||||
var triples []memory.Triple
|
||||
if doc == nil {
|
||||
return triples
|
||||
@ -409,6 +409,9 @@ func docToTriples(doc *document.Doc) []memory.Triple {
|
||||
|
||||
// NLP 通用提取
|
||||
e := nlp.NewExtractor(nil)
|
||||
if embedder != nil {
|
||||
e.SetEmbedder(embedder)
|
||||
}
|
||||
result := e.Extract(doc.Content)
|
||||
if result != nil {
|
||||
for _, nt := range result.Triples {
|
||||
|
||||
@ -42,8 +42,11 @@ type Distiller struct {
|
||||
ctx context.Context
|
||||
cancel context.CancelFunc
|
||||
onMemory func(input, response string)
|
||||
embedder nlp.Vectorizer
|
||||
}
|
||||
|
||||
func (d *Distiller) SetEmbedder(ev nlp.Vectorizer) { d.embedder = ev }
|
||||
|
||||
func NewDistiller(db *memory.GraphDB, dataDir string, cfg DistillerConfig) *Distiller {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
return &Distiller{
|
||||
@ -233,7 +236,7 @@ func (d *Distiller) distillBatch(batch []RawRecord) {
|
||||
assistantContent += r.Content + " "
|
||||
}
|
||||
}
|
||||
triples := extractKeyTriples(userContent, assistantContent)
|
||||
triples := extractKeyTriples(userContent, assistantContent, d.embedder)
|
||||
if len(triples) > 0 {
|
||||
sessionID := ""
|
||||
for sid := range sessionIDs {
|
||||
@ -263,10 +266,13 @@ func (d *Distiller) cleanupRawFiles() {
|
||||
}
|
||||
}
|
||||
|
||||
func extractKeyTriples(userContent, assistantContent string) []memory.Triple {
|
||||
func extractKeyTriples(userContent, assistantContent string, embedder nlp.Vectorizer) []memory.Triple {
|
||||
var triples []memory.Triple
|
||||
|
||||
e := nlp.NewExtractor(nil)
|
||||
if embedder != nil {
|
||||
e.SetEmbedder(embedder)
|
||||
}
|
||||
text := userContent
|
||||
if assistantContent != "" {
|
||||
text += assistantContent
|
||||
|
||||
@ -7,8 +7,16 @@ import (
|
||||
"time"
|
||||
|
||||
"gitcode.com/JianFeeeee/HomeAgent/internal/memory"
|
||||
"gitcode.com/JianFeeeee/HomeAgent/internal/memory/vector"
|
||||
)
|
||||
|
||||
// constVectorizer 恒等向量器:所有候选统一过融合阈值,验证嵌入接线是否生效。
|
||||
type constVectorizer struct{}
|
||||
|
||||
func (constVectorizer) Vectorize(text string) vector.Vector {
|
||||
return vector.Vector{"0": 1.0, "1": 0.5}
|
||||
}
|
||||
|
||||
func TestNewDistiller(t *testing.T) {
|
||||
d := NewDistiller(nil, t.TempDir(), DistillerConfig{
|
||||
Interval: 10 * time.Minute,
|
||||
@ -147,7 +155,7 @@ func TestExtractKeyTriples(t *testing.T) {
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
triples := extractKeyTriples(tt.user, tt.assistant)
|
||||
triples := extractKeyTriples(tt.user, tt.assistant, constVectorizer{})
|
||||
if tt.check != nil && !tt.check(triples) {
|
||||
t.Errorf("extractKeyTriples(%q) = %v, check failed", tt.user, triples)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user