Files
HomeAgent/internal/nlp/parser.go
root 51f190e0b6 core: pass ReasoningContent to assistant messages in process.go
- process.go: attach resp.ReasoningContent when building assistant messages
- deepseek.lua v2.1.0: remove last_reasoning closure hack; rely on
  core-provided reasoning_content in messages
- openai.lua: strip reasoning_content from messages in transform_request
  (not supported by OpenAI API)
2026-07-28 17:07:17 +08:00

217 lines
5.8 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

package nlp
import (
"sort"
"gitcode.com/JianFeeeee/HomeAgent/internal/memory/vector"
)
// Parser 依存句法分析器接口
type Parser interface {
Parse(text string) (*ParseResult, error)
}
// defaultParser 包级默认解析器,由 SetDefaultParser 设置
var defaultParser Parser
// SetDefaultParser 设置包级默认解析器。
// 设置后NewExtractor(nil) 将使用此解析器而非纯降级模式。
func SetDefaultParser(p Parser) {
defaultParser = p
}
// GetDefaultParser 返回当前包级默认解析器
func GetDefaultParser() Parser {
return defaultParser
}
// Vectorizer 向量化接口,复用 memory/vector 或 memory/static_embedder
type Vectorizer interface {
Vectorize(text string) vector.Vector
}
// ExtractorConfig 三元组提取器配置
type ExtractorConfig struct {
FusionAlpha float64 // syntax_conf 权重,默认 0.4
FusionBeta float64 // vector_conf 权重,默认 0.6
FusionThreshold float64 // 最终阈值,默认 0.3
}
// DefaultExtractorConfig 返回默认的提取器配置
func DefaultExtractorConfig() ExtractorConfig {
return ExtractorConfig{
FusionAlpha: 0.4,
FusionBeta: 0.6,
FusionThreshold: 0.3,
}
}
// Extractor 三元组提取器
type Extractor struct {
parser Parser
fallack Parser // 降级用 POS 模板解析器
embedder Vectorizer // 可选:用于 TransE 语义验证
fusionCfg ExtractorConfig
}
// NewExtractor 创建提取器。
// parser 为 nil 时尝试使用包级默认解析器 (SetDefaultParser)
// 若仍未设置则纯用 fallback (POS 模板匹配)。
func NewExtractor(parser Parser) *Extractor {
if parser == nil {
parser = defaultParser
}
return &Extractor{
parser: parser,
fallack: newFallbackParser(),
fusionCfg: DefaultExtractorConfig(),
}
}
// SetEmbedder 设置词嵌入向量化器,用于候选三元组的语义验证
func (e *Extractor) SetEmbedder(ev Vectorizer) {
e.embedder = ev
}
// SetFusionWeights 设置三元组融合裁决的权重参数。
// - alpha: syntax_conf 权重 (默认 0.4)
// - beta: vector_conf 权重 (默认 0.6)
// - threshold: 最终阈值 (默认 0.3)
func (e *Extractor) SetFusionWeights(alpha, beta, threshold float64) {
e.fusionCfg.FusionAlpha = alpha
e.fusionCfg.FusionBeta = beta
e.fusionCfg.FusionThreshold = threshold
}
// FusionConfig 返回当前融合裁诀配置
func (e *Extractor) FusionConfig() ExtractorConfig {
return e.fusionCfg
}
// Extract 从文本中提取三元组(完整四阶段流水线)
// Phase 1: 句法解析LTP 分词 → POS 标注 → 依存句法树)
// Phase 2: 结构初筛(依存模板 / POS 模板 → 候选三元组 + syntax_conf
// Phase 3: 语义验证TransE h+r≈t → vector_conf
// Phase 4: 融合裁决(线性加权 → 阈值截断 → 降序输出)
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
// ——— Phase 1 & 2: 句法解析 + 结构初筛 ———
if e.parser != nil {
result, err := e.parser.Parse(sentence)
if err == nil && result != nil && len(result.Tokens) > 1 {
triples = extractFromDep(result, sentence)
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, sentence)
if len(triples) > 0 {
src = "fallback"
}
}
}
// ——— Phase 3: 语义验证 (TransE h+r≈t) ———
if len(triples) > 0 && e.embedder != nil {
triples = verifyTriples(triples, e.embedder)
}
// ——— Phase 4: 融合裁决 ———
if len(triples) > 0 {
triples = fuseTriples(triples, e.fusionCfg)
}
allTriples = append(allTriples, triples...)
}
if len(allTriples) > 0 {
return &TripleSet{Triples: allTriples, Src: src}
}
return &TripleSet{Src: src}
}
// ——— Phase 3: 语义验证 ———
// verifyTriples 使用 TransE 打分 (h+r≈t) 计算 vector_conf
// 输入:候选三元组(带 syntax_conf
// 处理cos(h+r, t) → vector_conf
// 输出:带 vector_conf 的候选三元组
func verifyTriples(triples []Triple, embedder Vectorizer) []Triple {
for i := range triples {
t := &triples[i]
h := embedder.Vectorize(t.Subject)
r := embedder.Vectorize(t.Relation)
tv := embedder.Vectorize(t.Object)
hr := addVectors(h, r)
sim := vector.CosineSimilarity(hr, tv)
// 将 cos 映射到 [0, 1] 区间(原始可能在 [-1, 1]
t.VectorConf = (sim + 1.0) / 2.0
}
return triples
}
// ——— Phase 4: 融合裁决 ———
// fuseTriples 融合裁决:线性加权计算 final_score截断阈值降序输出
// 输入:候选三元组(带 syntax_conf + vector_conf
// 处理final_score = α * syntax_conf + β * vector_conf
// 输出:通过阈值且降序排列的最终三元组
func fuseTriples(triples []Triple, cfg ExtractorConfig) []Triple {
if len(triples) == 0 {
return triples
}
for i := range triples {
t := &triples[i]
finalScore := cfg.FusionAlpha*t.Score + cfg.FusionBeta*t.VectorConf
t.Score = finalScore
}
kept := make([]Triple, 0, len(triples))
for _, t := range triples {
if t.Score >= cfg.FusionThreshold {
kept = append(kept, t)
}
}
sort.Slice(kept, func(i, j int) bool {
return kept[i].Score > kept[j].Score
})
return kept
}
// ——— 向量工具 ———
// addVectors 向量加法 (h + r)
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
}