mirror of
https://gitcode.com/JianFeeeee/HomeAgent.git
synced 2026-09-21 17:38:10 +00:00
- 重写 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
122 lines
2.9 KiB
Go
122 lines
2.9 KiB
Go
package nlp
|
||
|
||
import "gitcode.com/JianFeeeee/HomeAgent/internal/memory/vector"
|
||
|
||
// Parser 依存句法分析器接口
|
||
type Parser interface {
|
||
Parse(text string) (*ParseResult, error)
|
||
}
|
||
|
||
// Vectorizer 向量化接口,复用 memory/vector 或 memory/static_embedder
|
||
type Vectorizer interface {
|
||
Vectorize(text string) vector.Vector
|
||
}
|
||
|
||
// Extractor 三元组提取器
|
||
type Extractor struct {
|
||
parser Parser
|
||
fallack Parser // 降级用 POS 模板解析器
|
||
embedder Vectorizer // 可选:用于 TransE 语义验证
|
||
}
|
||
|
||
// NewExtractor 创建提取器,parser 为 nil 时纯用 fallback
|
||
func NewExtractor(parser Parser) *Extractor {
|
||
return &Extractor{
|
||
parser: parser,
|
||
fallack: newFallbackParser(),
|
||
}
|
||
}
|
||
|
||
// SetEmbedder 设置词嵌入向量化器,用于候选三元组的语义验证
|
||
func (e *Extractor) SetEmbedder(ev Vectorizer) {
|
||
e.embedder = ev
|
||
}
|
||
|
||
// Extract 从文本中提取三元组
|
||
// 优先使用 parser,失败/无结果时自动降级到 fallback
|
||
// 如果设置了 embedder,还会做 h+r≈t 向量验证过滤
|
||
func (e *Extractor) Extract(text string) *TripleSet {
|
||
if text == "" {
|
||
return &TripleSet{Src: "", Err: nil}
|
||
}
|
||
|
||
var allTriples []Triple
|
||
src := ""
|
||
|
||
sentences := splitSentences(text)
|
||
for _, sentence := range sentences {
|
||
if sentence == "" {
|
||
continue
|
||
}
|
||
var triples []Triple
|
||
|
||
// 主线:依存解析 + 模板匹配
|
||
if e.parser != nil {
|
||
result, err := e.parser.Parse(sentence)
|
||
if err == nil && result != nil && len(result.Tokens) > 1 {
|
||
triples = extractFromDep(result)
|
||
if len(triples) > 0 {
|
||
src = "dep_parser"
|
||
}
|
||
}
|
||
}
|
||
|
||
// 降级:POS 模板匹配
|
||
if len(triples) == 0 && e.fallack != nil {
|
||
result, err := e.fallack.Parse(sentence)
|
||
if err == nil && result != nil && len(result.Tokens) > 1 {
|
||
triples = extractFromPOS(result)
|
||
if len(triples) > 0 {
|
||
src = "fallback"
|
||
}
|
||
}
|
||
}
|
||
|
||
// 向量验证(可选):用 h+r≈t 过滤不合理三元组
|
||
if len(triples) > 0 && e.embedder != nil {
|
||
triples = verifyTriples(triples, e.embedder)
|
||
}
|
||
|
||
allTriples = append(allTriples, triples...)
|
||
}
|
||
|
||
if len(allTriples) > 0 {
|
||
return &TripleSet{Triples: allTriples, Src: src}
|
||
}
|
||
return &TripleSet{Src: src}
|
||
}
|
||
|
||
// verifyTriples 使用 TransE 打分 (h+r≈t) 验证三元组,过滤低分项
|
||
func verifyTriples(triples []Triple, embedder Vectorizer) []Triple {
|
||
var kept []Triple
|
||
for _, t := range triples {
|
||
h := embedder.Vectorize(t.Subject)
|
||
r := embedder.Vectorize(t.Relation)
|
||
tv := embedder.Vectorize(t.Object)
|
||
|
||
hr := addVectors(h, r)
|
||
sim := vector.CosineSimilarity(hr, tv)
|
||
|
||
// 语义一致性过低 → 过滤(除非 fallback 无其他候选)
|
||
if sim >= 0.25 {
|
||
t.Score *= (0.5 + 0.5*sim)
|
||
kept = append(kept, t)
|
||
}
|
||
}
|
||
if len(kept) == 0 {
|
||
return triples
|
||
}
|
||
return kept
|
||
}
|
||
|
||
func addVectors(a, b vector.Vector) vector.Vector {
|
||
out := make(vector.Vector)
|
||
for k, v := range a {
|
||
out[k] = v
|
||
}
|
||
for k, v := range b {
|
||
out[k] += v
|
||
}
|
||
return out
|
||
}
|