v0.7.2: 根目录清理 + Agent 心跳重构 + 内嵌 ONNX 模型

- 根目录清理: branding/docs/knowledge -> assets/, package/tools/deploy -> deploy/
- meta.go: Version 0.7.2, SDKCompatibleVersion 语义改为最高兼容
- Makefile: 版本回退 0.7.2
- registry.go: 系统提示词改用 meta.Version 格式化
- Agent 心跳: reorgGraph 拆分为三个独立循环(archive/merge/review),各自可配间隔
- GraphDB: 新增 sentences 表 + 关系句子溯源 + ClearSentenceID + CleanupOrphanedSentences
- Knowledge: 支持词嵌入向量化器
- NLP 四阶段流水线: Parse -> Extract -> Verify -> Fuse + SentenceRef
- 移除远程 HTTP 解析器(remote_parser.go)
- 新增内嵌 ONNX 模型(vocab + dep_parser.onnx):
  +build onnxruntime: 全量 ONNX Runtime 推理
  !build onnxruntime: 内嵌词表规则式降级解析器
- config: core.agent.onnx_model_path 替代 dep_parser_url
This commit is contained in:
JianFeeeee
2026-07-28 09:56:26 +08:00
parent 61fbc55274
commit 5006712c8f
47 changed files with 26141 additions and 242 deletions

View File

@ -5,9 +5,10 @@ import "gitcode.com/JianFeeeee/HomeAgent/internal/memory"
// ToMemoryTriple 将 nlp.Triple 转为 memory.Triple
func ToMemoryTriple(t Triple) memory.Triple {
return memory.Triple{
Subject: t.Subject,
Relation: t.Relation,
Object: t.Object,
Confidence: t.Score,
Subject: t.Subject,
Relation: t.Relation,
Object: t.Object,
Confidence: t.Score,
SentenceText: t.SentenceRef,
}
}

View File

@ -36,19 +36,32 @@ var depTemplates = []depTemplate{
{subjRel: "SBV", objRel: "IOB", score: 0.85},
{subjRel: "SBV", objRel: "FOB", score: 0.8},
{subjRel: "SBV", objRel: "POB", score: 0.75},
{subjRel: "ATT", objRel: "VOB", score: 0.7},
{subjRel: "ATT", objRel: "IOB", score: 0.65},
{subjRel: "ATT", objRel: "FOB", score: 0.6},
{subjRel: "ATT", objRel: "POB", score: 0.55},
}
// extractFromDep 基于依存句法树提取三元组
func extractFromDep(result *ParseResult) []Triple {
// extractFromDep 基于依存句法树提取三元组 (Phase 2: 结构初筛)
// 输入:Token 序列(含依存关系)
// 处理:标记名词性节点 → 遍历谓词中心 → 收集 SBV/ATT 主语、VOB/IOB/POB 宾语 → 笛卡尔积 → 赋句法置信度 → ATT 链合并
// 输出:候选三元组列表(带 syntax_conf)
func extractFromDep(result *ParseResult, sentence string) []Triple {
if len(result.Tokens) < 2 {
return nil
}
// Step 1: 标记所有名词性节点为候选实体(供后续 ATT 合并等使用)
// (隐式使用,通过 isNounLike 判断)
// Step 2: 遍历所有动词节点作为谓词中心
verbIndices := findPredicates(result.POS, result.Tokens)
var triples []Triple
verbIndices := findPredicates(result.POS, result.Tokens)
for _, vi := range verbIndices {
var subj, obj string
var objIdx int
// Step 3: 沿依存弧收集主语(SBV/ATT)和宾语(VOB/IOB/FOB/POB)
var subjIndices, objIndices []int
var subjRels, objRels []string
for i, head := range result.Heads {
if head == 0 {
@ -60,63 +73,102 @@ func extractFromDep(result *ParseResult) []Triple {
}
rel := result.DepRels[i]
if isSubjRel(rel) && subj == "" {
subj = result.Tokens[i]
} else if isObjRel(rel) && obj == "" {
obj = result.Tokens[i]
objIdx = i
if isSubjRel(rel) {
subjIndices = append(subjIndices, i)
subjRels = append(subjRels, rel)
} else if isObjRel(rel) {
objIndices = append(objIndices, i)
objRels = append(objRels, rel)
}
}
if subj == "" {
// 主语降级:无 SBV/ATT 主语时向左查找最近的名词性节点
if len(subjIndices) == 0 {
for j := vi - 1; j >= 0; j-- {
if isNounLike(result.POS[j]) {
subj = result.Tokens[j]
subjIndices = append(subjIndices, j)
subjRels = append(subjRels, "SBV_IMPLICIT")
break
}
}
}
if subj != "" && obj != "" {
relLabel := result.Tokens[vi]
score := 0.8
if objIdx < len(result.Heads) && result.Heads[objIdx] == vi+1 {
// 宾语降级:无显式宾语时查找动词的其他名词性依赖
if len(objIndices) == 0 {
for i, head := range result.Heads {
if head == 0 {
continue
}
if head-1 == vi && isNounLike(result.POS[i]) && !isSubjRel(result.DepRels[i]) {
objIndices = append(objIndices, i)
objRels = append(objRels, "OBJ_IMPLICIT")
}
}
}
if len(subjIndices) == 0 || len(objIndices) == 0 {
continue
}
// Step 4: 笛卡尔积生成候选对,按模板赋予句法置信度
relLabel := result.Tokens[vi]
for _, si := range subjIndices {
for _, oi := range objIndices {
if si == oi {
continue
}
subj := result.Tokens[si]
obj := result.Tokens[oi]
score := 0.8 // 默认句法置信度
// 匹配模板查询精确置信度
for _, t := range depTemplates {
if t.objRel == result.DepRels[objIdx] {
if si < len(result.Heads) && result.Heads[si] == vi+1 &&
oi < len(result.Heads) && result.Heads[oi] == vi+1 &&
t.subjRel == result.DepRels[si] && t.objRel == result.DepRels[oi] {
score = t.score
break
}
}
triples = append(triples, Triple{
Subject: subj,
Relation: relLabel,
Object: obj,
Score: score,
Src: "dep",
SentenceRef: sentence,
})
}
triples = append(triples, Triple{
Subject: subj,
Relation: relLabel,
Object: obj,
Score: score,
Src: "dep",
})
}
// COO 链扩展:如果宾语有并列结构,为每个并列项生成三元组
if obj != "" {
cooExpanded := expandCOO(result, objIdx, vi)
// COO 链扩展:为每个宾语所在的并列结构生成额外三元组
for _, oi := range objIndices {
cooExpanded := expandCOO(result, oi, vi)
for _, cooObj := range cooExpanded {
if cooObj == obj {
if cooObj == result.Tokens[oi] {
continue
}
relLabel := result.Tokens[vi]
triples = append(triples, Triple{
Subject: subj,
Relation: relLabel,
Object: cooObj,
Score: 0.7,
Src: "dep_coo",
})
for _, si := range subjIndices {
subj := result.Tokens[si]
triples = append(triples, Triple{
Subject: subj,
Relation: relLabel,
Object: cooObj,
Score: 0.7,
Src: "dep_coo",
SentenceRef: sentence,
})
}
}
}
}
// Step 5: ATT 链合并多词实体
triples = mergeAttTriples(result, triples)
// 去重
triples = dedupTriples(triples)
return triples
}
@ -213,8 +265,8 @@ var posTemplates = []posTemplate{
{pattern: []string{"n", "v", "v", "n"}, subj: 0, verb: 1, obj: 3, score: 0.55},
}
// extractFromPOS 基于 POS 序列匹配模板提取三元组
func extractFromPOS(result *ParseResult) []Triple {
// extractFromPOS 基于 POS 序列匹配模板提取三元组 (Phase 2 降级路径)
func extractFromPOS(result *ParseResult, sentence string) []Triple {
if len(result.Tokens) < 2 {
return nil
}
@ -248,11 +300,12 @@ func extractFromPOS(result *ParseResult) []Triple {
continue
}
triples = append(triples, Triple{
Subject: subj,
Relation: verb,
Object: obj,
Score: tpl.score,
Src: "pos",
Subject: subj,
Relation: verb,
Object: obj,
Score: tpl.score,
Src: "pos",
SentenceRef: sentence,
})
}
}
@ -375,7 +428,7 @@ func isAdj(p string) bool {
}
func isSubjRel(rel string) bool {
return rel == "SBV"
return rel == "SBV" || rel == "ATT"
}
func isObjRel(rel string) bool {

View File

@ -35,7 +35,7 @@ func TestExtractFromPOS(t *testing.T) {
t.Run(tt.name, func(t *testing.T) {
if tt.input == "" {
result, _ := p.Parse("")
triples := extractFromPOS(result)
triples := extractFromPOS(result, "")
if len(triples) != 0 {
t.Errorf("expected 0 triples for empty, got %d", len(triples))
}
@ -48,7 +48,7 @@ func TestExtractFromPOS(t *testing.T) {
}
t.Logf("input=%q tokens=%v pos=%v", tt.input, result.Tokens, result.POS)
triples := extractFromPOS(result)
triples := extractFromPOS(result, tt.input)
for _, tr := range triples {
if tr.Subject == "" || tr.Relation == "" || tr.Object == "" {

View File

@ -10,11 +10,13 @@ type ParseResult struct {
// Triple 三元组 (subject, relation, object)
type Triple struct {
Subject string
Relation string
Object string
Score float64
Src string // "dep" / "fallback"
Subject string
Relation string
Object string
Score float64 // syntax_conf:句法置信度(Phase 2 输出)
VectorConf float64 // vector_conf:语义向量置信度(Phase 3 输出)
Src string // "dep" / "dep_coo" / "pos" / "fallback"
SentenceRef string // 原始句子,用于LLM复审时修正
}
// TripleSet 提取结果

Binary file not shown.

View File

@ -0,0 +1,20 @@
{
"<bos>": 0,
"ADJ": 1,
"ADP": 2,
"ADV": 3,
"AUX": 4,
"CCONJ": 5,
"DET": 6,
"INTJ": 7,
"NOUN": 8,
"NUM": 9,
"PART": 10,
"PRON": 11,
"PROPN": 12,
"PUNCT": 13,
"SCONJ": 14,
"SYM": 15,
"VERB": 16,
"X": 17
}

24949
internal/nlp/models/vocab.json Normal file

File diff suppressed because it is too large Load Diff

277
internal/nlp/onnx.go Normal file
View File

@ -0,0 +1,277 @@
//go:build onnxruntime
package nlp
import (
"embed"
"encoding/json"
"fmt"
"os"
"path/filepath"
"gitcode.com/JianFeeeee/HomeAgent/internal/config"
ort "github.com/yalue/onnxruntime_go"
)
//go:embed models/*
var onnxModelFS embed.FS
type ONNXParser struct {
rt *ort.AdvancedSession
vocab map[string]int64
posVocab map[string]int64
Release func()
}
type ONNXConfig struct {
ModelPath string // 留空使用内嵌模型
DataDir string // 模型解压/缓存目录
}
func NewONNXParser(cfg ONNXConfig) (*ONNXParser, error) {
vocab, err := loadJSONMap[int64]("models/vocab.json", onnxModelFS)
if err != nil {
return nil, fmt.Errorf("load vocab: %w", err)
}
posVocab, err := loadJSONMap[int64]("models/pos_vocab.json", onnxModelFS)
if err != nil {
return nil, fmt.Errorf("load pos_vocab: %w", err)
}
modelPath := cfg.ModelPath
if modelPath == "" {
modelPath, err = extractEmbeddedModel(cfg.DataDir)
if err != nil {
return nil, fmt.Errorf("extract model: %w", err)
}
}
ort.SetSharedLibraryPath(findONNXRuntime())
if err := ort.InitializeEnvironment(); err != nil {
return nil, fmt.Errorf("init onnx env: %w", err)
}
inputs := ort.NewInputDetails()
inputs.Append("input_ids", []int64{1, 128})
outputs := ort.NewOutputDetails()
outputs.Append("pos_logits", []int64{1, 128, 18})
outputs.Append("head_logits", []int64{1, 128, 128})
outputs.Append("rel_logits", []int64{1, 128, 128, 18})
session, err := ort.NewAdvancedSession(modelPath, inputs, outputs, nil)
if err != nil {
return nil, fmt.Errorf("create session: %w", err)
}
release := func() {
session.Destroy()
ort.DestroyEnvironment()
}
return &ONNXParser{
rt: session,
vocab: vocab,
posVocab: posVocab,
Release: release,
}, nil
}
func (p *ONNXParser) Parse(text string) (*ParseResult, error) {
if text == "" {
return &ParseResult{}, nil
}
inputIDs := tokenize(text, p.vocab, 128)
inputIDs = padTo(inputIDs, 128)
inputTensor, err := ort.NewTensor(ort.NewShape(1, 128), inputIDs)
if err != nil {
return nil, fmt.Errorf("create input tensor: %w", err)
}
defer inputTensor.Destroy()
outputs, err := p.rt.Call(inputTensor)
if err != nil {
return nil, fmt.Errorf("onnx call: %w", err)
}
rawPOS := outputs[0].GetData().([]float32)
rawHeads := outputs[1].GetData().([]float32)
rawRels := outputs[2].GetData().([]float32)
seqLen := actualLen(inputIDs)
tokens := idsToTokens(inputIDs[:seqLen], p.vocab)
pos := decodePOS(rawPOS, seqLen, p.posVocab)
heads := decodeHeads(rawHeads, seqLen)
rels := decodeRels(rawRels, seqLen)
return &ParseResult{Tokens: tokens, POS: pos, Heads: heads, DepRels: rels}, nil
}
func loadJSONMap[T ~int64 | ~string](path string, fs embed.FS) (map[string]T, error) {
data, err := fs.ReadFile(path)
if err != nil {
return nil, err
}
var raw struct {
Word map[string]T `json:"word"`
}
if err := json.Unmarshal(data, &raw); err != nil {
result := make(map[string]T)
if err2 := json.Unmarshal(data, &result); err2 != nil {
return nil, err
}
return result, nil
}
return raw.Word, nil
}
func extractEmbeddedModel(dataDir string) (string, error) {
if dataDir == "" {
dataDir = filepath.Join(os.TempDir(), "homeagent-nlp")
}
os.MkdirAll(dataDir, 0755)
dst := filepath.Join(dataDir, "dep_parser.onnx")
if _, err := os.Stat(dst); err == nil {
return dst, nil
}
data, err := onnxModelFS.ReadFile("models/dep_parser.onnx")
if err != nil {
return "", err
}
if err := os.WriteFile(dst, data, 0644); err != nil {
return "", err
}
return dst, nil
}
func findONNXRuntime() string {
candidates := []string{
"onnxruntime.dll",
"libonnxruntime.so",
"libonnxruntime.dylib",
filepath.Join(os.Getenv("ONNXRUNTIME_DIR"), "libonnxruntime.so"),
filepath.Join(os.Getenv("ONNXRUNTIME_DIR"), "onnxruntime.dll"),
}
for _, c := range candidates {
if _, err := os.Stat(c); err == nil {
abs, _ := filepath.Abs(c)
return abs
}
}
return "onnxruntime.dll"
}
func tokenize(text string, vocab map[string]int64, maxLen int) []int64 {
ids := []int64{vocab["<bos>"]}
runes := []rune(text)
for i := 0; i < len(runes) && len(ids) < maxLen; i++ {
if id, ok := vocab[string(runes[i])]; ok {
ids = append(ids, id)
} else {
ids = append(ids, vocab["<unk>"])
}
}
return ids
}
func padTo(ids []int64, length int) []int64 {
for len(ids) < length {
ids = append(ids, 0)
}
return ids
}
func actualLen(ids []int64) int {
for i, id := range ids {
if id == 0 {
return i
}
}
return len(ids)
}
func idsToTokens(ids []int64, vocab map[string]int64) []string {
rev := make(map[int64]string)
for k, v := range vocab {
rev[v] = k
}
var tokens []string
for _, id := range ids {
if t, ok := rev[id]; ok {
tokens = append(tokens, t)
}
}
return tokens
}
func decodePOS(raw []float32, seqLen int, posVocab map[string]int64) []string {
rev := make(map[int64]string)
for k, v := range posVocab {
rev[v] = k
}
pos := make([]string, seqLen)
for i := 0; i < seqLen; i++ {
bestIdx := 0
bestVal := float32(-1e9)
for j := 0; j < 18; j++ {
v := raw[i*18+j]
if v > bestVal {
bestVal = v
bestIdx = j
}
}
if tag, ok := rev[int64(bestIdx)]; ok {
pos[i] = tag
}
}
return pos
}
func decodeHeads(raw []float32, seqLen int) []int {
heads := make([]int, seqLen)
for i := 0; i < seqLen; i++ {
bestIdx := 0
bestVal := float32(-1e9)
for j := 0; j < seqLen; j++ {
v := raw[i*seqLen+j]
if v > bestVal {
bestVal = v
bestIdx = j
}
}
heads[i] = bestIdx
}
return heads
}
func decodeRels(raw []float32, seqLen int) []string {
rels := make([]string, seqLen)
for i := 0; i < seqLen; i++ {
bestIdx := 0
bestVal := float32(-1e9)
for j := 0; j < 18; j++ {
// average over head dimension for argmax
var sum float32
for k := 0; k < seqLen; k++ {
sum += raw[i*seqLen*18+k*18+j]
}
avg := sum / float32(seqLen)
if avg > bestVal {
bestVal = avg
bestIdx = j
}
}
rels[i] = posIDToTag(bestIdx)
}
return rels
}
func posIDToTag(id int) string {
tags := []string{"<bos>", "ADJ", "ADP", "ADV", "AUX", "CCONJ", "DET", "INTJ", "NOUN", "NUM", "PART", "PRON", "PROPN", "PUNCT", "SCONJ", "SYM", "VERB", "X"}
if id >= 0 && id < len(tags) {
return tags[id]
}
return "X"
}

View File

@ -2,27 +2,264 @@
package nlp
import "fmt"
import (
"embed"
"encoding/json"
"fmt"
"strings"
)
// ONNXParserStub 占位 — 编译时未启用 onnxruntime
type ONNXParser struct{}
//go:embed models/vocab.json models/pos_vocab.json
var vocabFS embed.FS
// ONNXParser 在未启用 onnxruntime 时作为规则式降级解析器。
// 使用内嵌词表实现基于词典的 POS 标注 + 基于 POS 序列的依存关系推断。
type ONNXParser struct {
vocab map[string]int
posVocab map[string]int
}
type ONNXConfig struct {
ModelPath string
VocabPath string
POSVocPath string
ModelPath string // 留空使用内嵌规则引擎
DataDir string // 仅在 onnxruntime 启用时使用
}
func NewONNXParser(cfg ONNXConfig) (*ONNXParser, error) {
return nil, fmt.Errorf("onnxparser: build with -tags onnxruntime to enable")
}
vocab := make(map[string]int)
data, err := vocabFS.ReadFile("models/vocab.json")
if err != nil {
return nil, fmt.Errorf("read vocab: %w", err)
}
func (p *ONNXParser) Close() {}
var raw struct {
Word map[string]int `json:"word"`
}
if err := json.Unmarshal(data, &raw); err != nil {
// 尝试直接解析为 flat map
var flat map[string]int
if err2 := json.Unmarshal(data, &flat); err2 != nil {
return nil, fmt.Errorf("parse vocab: %w", err)
}
vocab = flat
} else {
vocab = raw.Word
}
posVocab := make(map[string]int)
data, err = vocabFS.ReadFile("models/pos_vocab.json")
if err != nil {
return nil, fmt.Errorf("read pos_vocab: %w", err)
}
if err := json.Unmarshal(data, &posVocab); err != nil {
return nil, fmt.Errorf("parse pos_vocab: %w", err)
}
return &ONNXParser{vocab: vocab, posVocab: posVocab}, nil
}
func (p *ONNXParser) Parse(text string) (*ParseResult, error) {
return nil, fmt.Errorf("onnxparser: not available (build with -tags onnxruntime)")
if text == "" {
return &ParseResult{}, nil
}
// Phase 1: 基于词表的最大匹配分词
tokens := p.tokenize(text)
if len(tokens) == 0 {
return &ParseResult{}, nil
}
// Phase 2: 基于词表的规则式 POS 标注
pos := p.tagPOS(tokens)
// Phase 3: 基于 POS 序列的依存头推断
heads := p.inferHeads(tokens, pos)
// Phase 4: 关系标签推断
rels := p.inferRels(tokens, pos, heads)
return &ParseResult{
Tokens: tokens,
POS: pos,
Heads: heads,
DepRels: rels,
}, nil
}
func (p *ONNXParser) EnsureModel(dataDir string) error {
return fmt.Errorf("onnxparser: not available")
func (p *ONNXParser) tokenize(text string) []string {
runes := []rune(text)
var tokens []string
buf := []rune{}
for _, r := range runes {
if r == ' ' || r == '\t' || r == '\n' || r == '\r' {
if len(buf) > 0 {
tokens = append(tokens, string(buf))
buf = buf[:0]
}
continue
}
buf = append(buf, r)
// 最长匹配:检查当前 buf 是否在词表中
if _, ok := p.vocab[string(buf)]; !ok && len(buf) > 0 {
// 回退:取 buf[:-1] 作为词,继续
if _, ok2 := p.vocab[string(buf[:len(buf)-1])]; ok2 && len(buf) > 2 {
tokens = append(tokens, string(buf[:len(buf)-1]))
buf = buf[len(buf)-1:]
}
}
}
if len(buf) > 0 {
tokens = append(tokens, string(buf))
}
if len(tokens) == 0 {
tokens = strings.Fields(text)
}
return tokens
}
func (p *ONNXParser) tagPOS(tokens []string) []string {
pos := make([]string, len(tokens))
for i, t := range tokens {
pos[i] = p.guessPOS(t)
}
return pos
}
func (p *ONNXParser) guessPOS(word string) string {
if _, ok := p.vocab[word]; !ok {
// OOV: 基于启发式
if len(word) == 0 {
return "X"
}
if isPunct([]rune(word)[0]) {
return "PUNCT"
}
if isDigit(word) {
return "NUM"
}
return "X"
}
// 对词表中的词,基于可用特征判断
runes := []rune(word)
if len(runes) == 0 {
return "X"
}
first := runes[0]
if isPunct(first) {
return "PUNCT"
}
return "NOUN"
}
func isPunct(r rune) bool {
return (r >= 0x3000 && r <= 0x303F) || // CJK 标点
(r >= 0xFF00 && r <= 0xFFEF) || // 全角
r == '.' || r == ',' || r == '!' || r == '?' ||
r == ';' || r == ':' || r == '"' || r == '\'' ||
r == '(' || r == ')' || r == '[' || r == ']' ||
r == '{' || r == '}' || r == '。' || r == ',' ||
r == '!' || r == '?' || r == ';' || r == ':' ||
r == '、' || r == '‘' || r == '’' || r == '“' || r == '”'
}
func isDigit(s string) bool {
for _, r := range s {
if r < '0' || r > '9' {
if r < 0xFF10 || r > 0xFF19 { // 全角数字
return false
}
}
}
return len(s) > 0
}
// inferHeads 基于 POS 序列的规则式依存头推断。
// 动词通常作为根(head=0),名词依附于动词,形容词依附于名词。
func (p *ONNXParser) inferHeads(tokens []string, pos []string) []int {
n := len(tokens)
heads := make([]int, n)
// 找到第一个动词作为根
rootIdx := -1
for i, tag := range pos {
if tag == "VERB" {
rootIdx = i
break
}
}
if rootIdx < 0 {
rootIdx = 0
}
heads[rootIdx] = 0
for i := 0; i < n; i++ {
if i == rootIdx {
continue
}
switch pos[i] {
case "NOUN", "PROPN":
// 名词指向最近的动词或前一个名词
if i < rootIdx {
heads[i] = rootIdx
} else {
heads[i] = rootIdx
}
case "ADJ", "ADV":
// 修饰语指向前一个名词或动词
if i > 0 {
heads[i] = i - 1
} else {
heads[i] = rootIdx
}
case "NUM", "DET":
// 限定词指向前一个名词
if i > 0 {
heads[i] = i - 1
} else {
heads[i] = rootIdx
}
case "PUNCT":
heads[i] = rootIdx
default:
heads[i] = rootIdx
}
}
return heads
}
// inferRels 基于 POS 对的关系标签推断。
func (p *ONNXParser) inferRels(tokens []string, pos []string, heads []int) []string {
n := len(tokens)
rels := make([]string, n)
for i := 0; i < n; i++ {
if heads[i] == 0 {
rels[i] = "ROOT"
continue
}
h := heads[i]
if h < 0 || h >= n {
rels[i] = "dep"
continue
}
rels[i] = posToRel(pos[h], pos[i])
}
return rels
}
func posToRel(headPOS, depPOS string) string {
switch {
case depPOS == "NOUN" || depPOS == "PROPN":
return "nsubj"
case depPOS == "ADJ":
return "amod"
case depPOS == "ADV":
return "advmod"
case depPOS == "NUM" || depPOS == "DET":
return "det"
case depPOS == "VERB":
return "xcomp"
case depPOS == "PUNCT":
return "punct"
default:
return "dep"
}
}

View File

@ -1,12 +1,30 @@
package nlp
import "gitcode.com/JianFeeeee/HomeAgent/internal/memory/vector"
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
@ -19,8 +37,13 @@ type Extractor struct {
embedder Vectorizer // 可选:用于 TransE 语义验证
}
// NewExtractor 创建提取器,parser 为 nil 时纯用 fallback
// NewExtractor 创建提取器。
// parser 为 nil 时尝试使用包级默认解析器 (SetDefaultParser),
// 若仍未设置则纯用 fallback (POS 模板匹配)。
func NewExtractor(parser Parser) *Extractor {
if parser == nil {
parser = defaultParser
}
return &Extractor{
parser: parser,
fallack: newFallbackParser(),
@ -32,9 +55,11 @@ func (e *Extractor) SetEmbedder(ev Vectorizer) {
e.embedder = ev
}
// Extract 从文本中提取三元组
// 优先使用 parser,失败/无结果时自动降级到 fallback
// 如果设置了 embedder,还会做 h+r≈t 向量验证过滤
// 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}
@ -50,11 +75,11 @@ func (e *Extractor) Extract(text string) *TripleSet {
}
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)
triples = extractFromDep(result, sentence)
if len(triples) > 0 {
src = "dep_parser"
}
@ -65,18 +90,23 @@ func (e *Extractor) Extract(text string) *TripleSet {
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)
triples = extractFromPOS(result, sentence)
if len(triples) > 0 {
src = "fallback"
}
}
}
// 向量验证(可选):用 h+r≈t 过滤不合理三元组
// ——— 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)
}
allTriples = append(allTriples, triples...)
}
@ -86,10 +116,15 @@ func (e *Extractor) Extract(text string) *TripleSet {
return &TripleSet{Src: src}
}
// verifyTriples 使用 TransE 打分 (h+r≈t) 验证三元组,过滤低分项
// ——— 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 {
var kept []Triple
for _, t := range triples {
for i := range triples {
t := &triples[i]
h := embedder.Vectorize(t.Subject)
r := embedder.Vectorize(t.Relation)
tv := embedder.Vectorize(t.Object)
@ -97,18 +132,55 @@ func verifyTriples(triples []Triple, embedder Vectorizer) []Triple {
hr := addVectors(h, r)
sim := vector.CosineSimilarity(hr, tv)
// 语义一致性过低 → 过滤(除非 fallback 无其他候选)
if sim >= 0.25 {
t.Score *= (0.5 + 0.5*sim)
// 将 cos 映射到 [0, 1] 区间(原始可能在 [-1, 1])
t.VectorConf = (sim + 1.0) / 2.0
}
return triples
}
// ——— Phase 4: 融合裁决 ———
const (
fusionAlpha = 0.4 // syntax_conf 权重
fusionBeta = 0.6 // vector_conf 权重
fusionThreshold = 0.3 // 最终阈值
)
// fuseTriples 融合裁决:线性加权计算 final_score,截断阈值,降序输出
// 输入:候选三元组(带 syntax_conf + vector_conf)
// 处理:final_score = α * syntax_conf + β * vector_conf
// 输出:通过阈值且降序排列的最终三元组
func fuseTriples(triples []Triple) []Triple {
if len(triples) == 0 {
return triples
}
// 计算 final_score 并更新 Score 字段
for i := range triples {
t := &triples[i]
finalScore := fusionAlpha*t.Score + fusionBeta*t.VectorConf
t.Score = finalScore
}
// 截断低分项
kept := make([]Triple, 0, len(triples))
for _, t := range triples {
if t.Score >= fusionThreshold {
kept = append(kept, t)
}
}
if len(kept) == 0 {
return triples
}
// 降序排列
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 {
@ -118,4 +190,4 @@ func addVectors(a, b vector.Vector) vector.Vector {
out[k] += v
}
return out
}
}