mirror of
https://gitcode.com/JianFeeeee/HomeAgent.git
synced 2026-09-21 17:38:10 +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:
13
internal/nlp/bridge.go
Normal file
13
internal/nlp/bridge.go
Normal file
@ -0,0 +1,13 @@
|
||||
package nlp
|
||||
|
||||
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,
|
||||
}
|
||||
}
|
||||
103
internal/nlp/download.go
Normal file
103
internal/nlp/download.go
Normal file
@ -0,0 +1,103 @@
|
||||
package nlp
|
||||
|
||||
import (
|
||||
"crypto/md5"
|
||||
"fmt"
|
||||
"io"
|
||||
"log"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
)
|
||||
|
||||
// ModelSource 模型来源:本地路径或远程 URL
|
||||
type ModelSource struct {
|
||||
Path string // 本地路径(优先)
|
||||
URL string // 远程下载地址
|
||||
}
|
||||
|
||||
// EnsureModel 确保模型文件存在,返回最终路径
|
||||
func EnsureModel(dstDir string, src ModelSource, filename string) (string, error) {
|
||||
if err := os.MkdirAll(dstDir, 0755); err != nil {
|
||||
return "", fmt.Errorf("create dir %s: %w", dstDir, err)
|
||||
}
|
||||
|
||||
dst := filepath.Join(dstDir, filename)
|
||||
|
||||
// 1. 本地路径优先
|
||||
if src.Path != "" {
|
||||
if _, err := os.Stat(src.Path); err == nil {
|
||||
if err := copyFile(src.Path, dst); err != nil {
|
||||
return "", fmt.Errorf("copy from %s: %w", src.Path, err)
|
||||
}
|
||||
log.Printf("[nlp] model ready (local): %s", dst)
|
||||
return dst, nil
|
||||
}
|
||||
log.Printf("[nlp] local path %s not found, trying remote...", src.Path)
|
||||
}
|
||||
|
||||
// 2. 远程下载
|
||||
if src.URL != "" {
|
||||
if _, err := os.Stat(dst); err == nil {
|
||||
return dst, nil // 已存在
|
||||
}
|
||||
log.Printf("[nlp] downloading model from %s ...", src.URL)
|
||||
if err := downloadFile(dst, src.URL); err != nil {
|
||||
return "", fmt.Errorf("download from %s: %w", src.URL, err)
|
||||
}
|
||||
return dst, nil
|
||||
}
|
||||
|
||||
return "", fmt.Errorf("model not found: no local path or remote URL")
|
||||
}
|
||||
|
||||
func downloadFile(dst, url string) error {
|
||||
tmp := dst + ".download." + fmt.Sprintf("%x", md5.Sum([]byte(url)))
|
||||
|
||||
resp, err := http.Get(url)
|
||||
if err != nil {
|
||||
return fmt.Errorf("http get %s: %w", url, err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return fmt.Errorf("http status %s", resp.Status)
|
||||
}
|
||||
|
||||
f, err := os.Create(tmp)
|
||||
if err != nil {
|
||||
return fmt.Errorf("create temp %s: %w", tmp, err)
|
||||
}
|
||||
|
||||
written, err := io.Copy(f, resp.Body)
|
||||
f.Close()
|
||||
if err != nil {
|
||||
os.Remove(tmp)
|
||||
return fmt.Errorf("write: %w", err)
|
||||
}
|
||||
|
||||
if err := os.Rename(tmp, dst); err != nil {
|
||||
os.Remove(tmp)
|
||||
return fmt.Errorf("rename: %w", err)
|
||||
}
|
||||
|
||||
log.Printf("[nlp] downloaded %d bytes to %s", written, dst)
|
||||
return nil
|
||||
}
|
||||
|
||||
func copyFile(src, dst string) error {
|
||||
in, err := os.Open(src)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer in.Close()
|
||||
|
||||
out, err := os.Create(dst)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer out.Close()
|
||||
|
||||
_, err = io.Copy(out, in)
|
||||
return err
|
||||
}
|
||||
383
internal/nlp/extractor.go
Normal file
383
internal/nlp/extractor.go
Normal file
@ -0,0 +1,383 @@
|
||||
package nlp
|
||||
|
||||
import "strings"
|
||||
|
||||
// ——— 分句 ———
|
||||
|
||||
func splitSentences(text string) []string {
|
||||
var sentences []string
|
||||
buf := strings.Builder{}
|
||||
for _, r := range text {
|
||||
buf.WriteRune(r)
|
||||
if r == '。' || r == '!' || r == '?' || r == ';' || r == '\n' {
|
||||
s := strings.TrimSpace(buf.String())
|
||||
if s != "" {
|
||||
sentences = append(sentences, s)
|
||||
}
|
||||
buf.Reset()
|
||||
}
|
||||
}
|
||||
if tail := strings.TrimSpace(buf.String()); tail != "" {
|
||||
sentences = append(sentences, tail)
|
||||
}
|
||||
return sentences
|
||||
}
|
||||
|
||||
// ——— 依存句法模板 ———
|
||||
|
||||
type depTemplate struct {
|
||||
subjRel string
|
||||
objRel string
|
||||
score float64
|
||||
}
|
||||
|
||||
var depTemplates = []depTemplate{
|
||||
{subjRel: "SBV", objRel: "VOB", score: 0.9},
|
||||
{subjRel: "SBV", objRel: "IOB", score: 0.85},
|
||||
{subjRel: "SBV", objRel: "FOB", score: 0.8},
|
||||
{subjRel: "SBV", objRel: "POB", score: 0.75},
|
||||
}
|
||||
|
||||
// extractFromDep 基于依存句法树提取三元组
|
||||
func extractFromDep(result *ParseResult) []Triple {
|
||||
if len(result.Tokens) < 2 {
|
||||
return nil
|
||||
}
|
||||
var triples []Triple
|
||||
|
||||
verbIndices := findPredicates(result.POS, result.Tokens)
|
||||
for _, vi := range verbIndices {
|
||||
var subj, obj string
|
||||
var objIdx int
|
||||
|
||||
for i, head := range result.Heads {
|
||||
if head == 0 {
|
||||
continue
|
||||
}
|
||||
parentIdx := head - 1
|
||||
if parentIdx != vi {
|
||||
continue
|
||||
}
|
||||
rel := result.DepRels[i]
|
||||
|
||||
if isSubjRel(rel) && subj == "" {
|
||||
subj = result.Tokens[i]
|
||||
} else if isObjRel(rel) && obj == "" {
|
||||
obj = result.Tokens[i]
|
||||
objIdx = i
|
||||
}
|
||||
}
|
||||
|
||||
if subj == "" {
|
||||
for j := vi - 1; j >= 0; j-- {
|
||||
if isNounLike(result.POS[j]) {
|
||||
subj = result.Tokens[j]
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if subj != "" && obj != "" {
|
||||
relLabel := result.Tokens[vi]
|
||||
score := 0.8
|
||||
if objIdx < len(result.Heads) && result.Heads[objIdx] == vi+1 {
|
||||
for _, t := range depTemplates {
|
||||
if t.objRel == result.DepRels[objIdx] {
|
||||
score = t.score
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
triples = append(triples, Triple{
|
||||
Subject: subj,
|
||||
Relation: relLabel,
|
||||
Object: obj,
|
||||
Score: score,
|
||||
Src: "dep",
|
||||
})
|
||||
}
|
||||
|
||||
// COO 链扩展:如果宾语有并列结构,为每个并列项生成三元组
|
||||
if obj != "" {
|
||||
cooExpanded := expandCOO(result, objIdx, vi)
|
||||
for _, cooObj := range cooExpanded {
|
||||
if cooObj == obj {
|
||||
continue
|
||||
}
|
||||
relLabel := result.Tokens[vi]
|
||||
triples = append(triples, Triple{
|
||||
Subject: subj,
|
||||
Relation: relLabel,
|
||||
Object: cooObj,
|
||||
Score: 0.7,
|
||||
Src: "dep_coo",
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
triples = mergeAttTriples(result, triples)
|
||||
return triples
|
||||
}
|
||||
|
||||
// expandCOO 从宾语开始沿 COO 链展开所有并列项
|
||||
func expandCOO(result *ParseResult, startIdx, excludeParent int) []string {
|
||||
var expanded []string
|
||||
seen := make(map[int]bool)
|
||||
|
||||
var walk func(idx int)
|
||||
walk = func(idx int) {
|
||||
if idx < 0 || idx >= len(result.Tokens) || seen[idx] {
|
||||
return
|
||||
}
|
||||
seen[idx] = true
|
||||
expanded = append(expanded, result.Tokens[idx])
|
||||
for i, head := range result.Heads {
|
||||
if head == 0 {
|
||||
continue
|
||||
}
|
||||
if result.DepRels[i] == "COO" && head-1 == idx && i != excludeParent {
|
||||
walk(i)
|
||||
}
|
||||
}
|
||||
}
|
||||
walk(startIdx)
|
||||
return expanded
|
||||
}
|
||||
|
||||
// ——— POS 序列模板(降级) ———
|
||||
|
||||
type posTemplate struct {
|
||||
pattern []string
|
||||
subj int // 主语在 pattern 中的绝对索引
|
||||
verb int // 谓语在 pattern 中的绝对索引
|
||||
obj int // 宾语在 pattern 中的绝对索引
|
||||
score float64
|
||||
}
|
||||
|
||||
var posTemplates = []posTemplate{
|
||||
// 我/r 吃/v 苹果/n
|
||||
{pattern: []string{"r", "v", "n"}, subj: 0, verb: 1, obj: 2, score: 0.7},
|
||||
// 我/r 吃/v 苹果/n
|
||||
{pattern: []string{"r", "v", "nr"}, subj: 0, verb: 1, obj: 2, score: 0.7},
|
||||
// 我/r 是/v 学生/n
|
||||
{pattern: []string{"r", "v", "n"}, subj: 0, verb: 1, obj: 2, score: 0.7},
|
||||
// 小明/nr 喜欢/v 篮球/n
|
||||
{pattern: []string{"nr", "v", "n"}, subj: 0, verb: 1, obj: 2, score: 0.7},
|
||||
// 小明/nr 打/v 篮球/n
|
||||
{pattern: []string{"nr", "v", "nr"}, subj: 0, verb: 1, obj: 2, score: 0.65},
|
||||
// 我/r 在/p 杭州/ns 读书/v
|
||||
{pattern: []string{"r", "p", "ns", "v"}, subj: 0, verb: 3, obj: 2, score: 0.65},
|
||||
// 我/r 在/p 杭州/ns 读书/n(读书被标为 n)
|
||||
{pattern: []string{"r", "p", "ns", "n"}, subj: 0, verb: 3, obj: 2, score: 0.55},
|
||||
// 我/r 在/p 杭州/ns 工作/vn
|
||||
{pattern: []string{"r", "p", "ns", "vn"}, subj: 0, verb: 3, obj: 2, score: 0.6},
|
||||
// 小明/nr 在/p 杭州/ns 读书/v
|
||||
{pattern: []string{"nr", "p", "ns", "v"}, subj: 0, verb: 3, obj: 2, score: 0.65},
|
||||
// 我/r 住在/p 杭州/ns ("住"被标为 v,"在"是 p)
|
||||
{pattern: []string{"r", "v", "p", "ns"}, subj: 0, verb: 1, obj: 3, score: 0.6},
|
||||
// 小明/nr 住在/p 北京/ns
|
||||
{pattern: []string{"nr", "v", "p", "ns"}, subj: 0, verb: 1, obj: 3, score: 0.6},
|
||||
// 天气/n 很/d 好/a
|
||||
{pattern: []string{"n", "d", "a"}, subj: 0, verb: 2, obj: 2, score: 0.5},
|
||||
// 天气/n 很/zg 好/a(很 被标为 zg 而非 d)
|
||||
{pattern: []string{"n", "zg", "a"}, subj: 0, verb: 2, obj: 2, score: 0.45},
|
||||
// 今天/t 天气/n 好/a
|
||||
{pattern: []string{"t", "n", "a"}, subj: 1, verb: 2, obj: 2, score: 0.5},
|
||||
// 我/r 喜欢/v 跑步/vn
|
||||
{pattern: []string{"r", "v", "vn"}, subj: 0, verb: 1, obj: 2, score: 0.6},
|
||||
// 我/r 喜欢/v 游泳/vn
|
||||
{pattern: []string{"r", "v", "v"}, subj: 0, verb: 1, obj: 2, score: 0.65},
|
||||
// 我/r 叫/v 小明/nr
|
||||
{pattern: []string{"r", "v", "nr"}, subj: 0, verb: 1, obj: 2, score: 0.7},
|
||||
// 通用:代词/名词 + 动词 + 名词
|
||||
{pattern: []string{"r", "v", "ns"}, subj: 0, verb: 1, obj: 2, score: 0.6},
|
||||
{pattern: []string{"n", "v", "n"}, subj: 0, verb: 1, obj: 2, score: 0.65},
|
||||
// 我/r 吃/v 了/u 苹果/n
|
||||
{pattern: []string{"r", "v", "u", "n"}, subj: 0, verb: 1, obj: 3, score: 0.6},
|
||||
// 名词跟在代词后作为谓语(打球/n 在 我/r 后)
|
||||
{pattern: []string{"r", "n"}, subj: 0, verb: 1, obj: 1, score: 0.5},
|
||||
// 小明/x 喜欢/v 吃/v 苹果/n(x 为人名,连动结构)
|
||||
{pattern: []string{"x", "v", "v", "n"}, subj: 0, verb: 1, obj: 3, score: 0.55},
|
||||
// 小明/x 喜欢/v 苹果/n
|
||||
{pattern: []string{"x", "v", "n"}, subj: 0, verb: 1, obj: 2, score: 0.55},
|
||||
// 我/r 喜欢/v 吃/v 苹果/n
|
||||
{pattern: []string{"r", "v", "v", "n"}, subj: 0, verb: 1, obj: 3, score: 0.6},
|
||||
// 通用:x 标签代词 + 动词 + vn
|
||||
{pattern: []string{"x", "v", "vn"}, subj: 0, verb: 1, obj: 2, score: 0.5},
|
||||
// 小明/x 在/p 北京/ns 工作/v
|
||||
{pattern: []string{"x", "p", "ns", "v"}, subj: 0, verb: 3, obj: 2, score: 0.55},
|
||||
// 小明/x 在/p 北京/ns 上班/vn
|
||||
{pattern: []string{"x", "p", "ns", "vn"}, subj: 0, verb: 3, obj: 2, score: 0.5},
|
||||
// 名词/n + 动词/v + 动词/v + 名词/n(连动)
|
||||
{pattern: []string{"n", "v", "v", "n"}, subj: 0, verb: 1, obj: 3, score: 0.55},
|
||||
}
|
||||
|
||||
// extractFromPOS 基于 POS 序列匹配模板提取三元组
|
||||
func extractFromPOS(result *ParseResult) []Triple {
|
||||
if len(result.Tokens) < 2 {
|
||||
return nil
|
||||
}
|
||||
|
||||
var triples []Triple
|
||||
pos := result.POS
|
||||
tokens := result.Tokens
|
||||
|
||||
for _, tpl := range posTemplates {
|
||||
pat := tpl.pattern
|
||||
if len(pat) > len(pos) {
|
||||
continue
|
||||
}
|
||||
for i := 0; i <= len(pos)-len(pat); i++ {
|
||||
if !matchPOS(pos[i:i+len(pat)], pat) {
|
||||
continue
|
||||
}
|
||||
|
||||
subj := tokens[i+tpl.subj]
|
||||
verb := tokens[i+tpl.verb]
|
||||
obj := tokens[i+tpl.obj]
|
||||
if subj == "" || verb == "" || obj == "" {
|
||||
continue
|
||||
}
|
||||
// 跳过自指谓语/无宾语谓语
|
||||
if subj == obj {
|
||||
continue
|
||||
}
|
||||
// 跳过谓语等于宾语(形容词谓语等无实际宾语的情况)
|
||||
if verb == obj {
|
||||
continue
|
||||
}
|
||||
triples = append(triples, Triple{
|
||||
Subject: subj,
|
||||
Relation: verb,
|
||||
Object: obj,
|
||||
Score: tpl.score,
|
||||
Src: "pos",
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// 去重(相同 subj/rel/obj 只保留一个)
|
||||
triples = dedupTriples(triples)
|
||||
return triples
|
||||
}
|
||||
|
||||
func matchPOS(got, want []string) bool {
|
||||
if len(got) != len(want) {
|
||||
return false
|
||||
}
|
||||
for i := range got {
|
||||
if got[i] != want[i] {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func dedupTriples(triples []Triple) []Triple {
|
||||
seen := make(map[string]bool)
|
||||
var out []Triple
|
||||
for _, t := range triples {
|
||||
key := t.Subject + "\x00" + t.Relation + "\x00" + t.Object
|
||||
if seen[key] {
|
||||
continue
|
||||
}
|
||||
seen[key] = true
|
||||
out = append(out, t)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// ——— ATT 链合并 ———
|
||||
|
||||
// mergeAttTriples ATT 链合并:将定语合并到被修饰词
|
||||
func mergeAttTriples(result *ParseResult, triples []Triple) []Triple {
|
||||
attMap := make(map[int][]int)
|
||||
for i, head := range result.Heads {
|
||||
if head == 0 {
|
||||
continue
|
||||
}
|
||||
if i >= len(result.DepRels) {
|
||||
continue
|
||||
}
|
||||
if result.DepRels[i] == "ATT" {
|
||||
parentIdx := head - 1
|
||||
attMap[parentIdx] = append(attMap[parentIdx], i)
|
||||
}
|
||||
}
|
||||
if len(attMap) == 0 {
|
||||
return triples
|
||||
}
|
||||
for i := range triples {
|
||||
for headIdx, attIds := range attMap {
|
||||
if headIdx >= len(result.Tokens) {
|
||||
continue
|
||||
}
|
||||
headWord := result.Tokens[headIdx]
|
||||
var attWords []string
|
||||
for _, aid := range attIds {
|
||||
if aid < len(result.Tokens) {
|
||||
attWords = append(attWords, result.Tokens[aid])
|
||||
}
|
||||
}
|
||||
if len(attWords) == 0 {
|
||||
continue
|
||||
}
|
||||
expanded := strings.Join(attWords, "") + headWord
|
||||
if triples[i].Subject == headWord {
|
||||
triples[i].Subject = expanded
|
||||
}
|
||||
if triples[i].Object == headWord {
|
||||
triples[i].Object = expanded
|
||||
}
|
||||
}
|
||||
}
|
||||
return triples
|
||||
}
|
||||
|
||||
// ——— helper ———
|
||||
|
||||
func findPredicates(pos []string, tokens []string) []int {
|
||||
var indices []int
|
||||
for i, p := range pos {
|
||||
if isVerb(p) || isAdj(p) {
|
||||
indices = append(indices, i)
|
||||
continue
|
||||
}
|
||||
if isNounLike(p) && i > 0 && isPronoun(pos[i-1]) {
|
||||
indices = append(indices, i)
|
||||
continue
|
||||
}
|
||||
if isNounLike(p) && i > 0 && isNounLike(pos[i-1]) {
|
||||
indices = append(indices, i)
|
||||
continue
|
||||
}
|
||||
}
|
||||
return indices
|
||||
}
|
||||
|
||||
func isVerb(p string) bool {
|
||||
return p == "v" || p == "vd" || strings.HasPrefix(p, "v")
|
||||
}
|
||||
|
||||
func isNounLike(p string) bool {
|
||||
return p == "n" || p == "nr" || p == "ns" || p == "nt" || p == "nz" ||
|
||||
p == "an" || p == "vn" || p == "x" ||
|
||||
strings.HasPrefix(p, "n")
|
||||
}
|
||||
|
||||
func isPronoun(p string) bool {
|
||||
return p == "r"
|
||||
}
|
||||
|
||||
func isAdj(p string) bool {
|
||||
return p == "a"
|
||||
}
|
||||
|
||||
func isSubjRel(rel string) bool {
|
||||
return rel == "SBV"
|
||||
}
|
||||
|
||||
func isObjRel(rel string) bool {
|
||||
return rel == "VOB" || rel == "IOB" || rel == "FOB" || rel == "POB"
|
||||
}
|
||||
99
internal/nlp/extractor_test.go
Normal file
99
internal/nlp/extractor_test.go
Normal file
@ -0,0 +1,99 @@
|
||||
package nlp
|
||||
|
||||
import (
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestFallbackParseDebug(t *testing.T) {
|
||||
cases := []string{"我打球", "我在杭州读书", "小明喜欢吃苹果", "天气很好", "我住在杭州"}
|
||||
p := newFallbackParser()
|
||||
for _, c := range cases {
|
||||
result, err := p.Parse(c)
|
||||
if err != nil || result == nil || len(result.Tokens) == 0 {
|
||||
t.Skip("jieba not available")
|
||||
}
|
||||
t.Logf("%q → tokens=%v pos=%v", c, result.Tokens, result.POS)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractFromPOS(t *testing.T) {
|
||||
p := newFallbackParser()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
input string
|
||||
}{
|
||||
{"pronoun_prep_ns_noun", "我在杭州读书"},
|
||||
{"pronoun_verb_noun", "我打球"},
|
||||
{"name_verb_noun", "小明喜欢吃苹果"},
|
||||
{"adj_predicate", "天气很好"},
|
||||
{"pronoun_verb_prep_ns", "我住在杭州"},
|
||||
{"empty", ""},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if tt.input == "" {
|
||||
result, _ := p.Parse("")
|
||||
triples := extractFromPOS(result)
|
||||
if len(triples) != 0 {
|
||||
t.Errorf("expected 0 triples for empty, got %d", len(triples))
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
result, err := p.Parse(tt.input)
|
||||
if err != nil || result == nil || len(result.Tokens) == 0 {
|
||||
t.Skip("jieba not available")
|
||||
}
|
||||
|
||||
t.Logf("input=%q tokens=%v pos=%v", tt.input, result.Tokens, result.POS)
|
||||
triples := extractFromPOS(result)
|
||||
|
||||
for _, tr := range triples {
|
||||
if tr.Subject == "" || tr.Relation == "" || tr.Object == "" {
|
||||
t.Errorf("triple has empty field: %+v", tr)
|
||||
}
|
||||
t.Logf("triple: Subject=%q Relation=%q Object=%q score=%.2f", tr.Subject, tr.Relation, tr.Object, tr.Score)
|
||||
}
|
||||
|
||||
if len(triples) == 0 {
|
||||
t.Logf("no triples extracted (may be expected depending on jieba POS tagging)")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractorFallback(t *testing.T) {
|
||||
e := NewExtractor(nil)
|
||||
result := e.Extract("我住在杭州")
|
||||
if result == nil {
|
||||
t.Fatal("expected result")
|
||||
}
|
||||
if result.Src == "" {
|
||||
t.Skip("jieba not available")
|
||||
}
|
||||
if len(result.Triples) > 0 {
|
||||
tr := result.Triples[0]
|
||||
t.Logf("extracted: Subject=%q Relation=%q Object=%q (score=%.2f, src=%s)",
|
||||
tr.Subject, tr.Relation, tr.Object, tr.Score, tr.Src)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractorWithDepStub(t *testing.T) {
|
||||
dummy := &dummyParser{}
|
||||
e := NewExtractor(dummy)
|
||||
result := e.Extract("我今天去北京")
|
||||
if result == nil {
|
||||
t.Fatal("expected result")
|
||||
}
|
||||
if len(result.Triples) > 0 {
|
||||
t.Logf("result: src=%s, triples=%+v", result.Src, result.Triples)
|
||||
}
|
||||
}
|
||||
|
||||
type dummyParser struct{}
|
||||
|
||||
func (d *dummyParser) Parse(text string) (*ParseResult, error) {
|
||||
return &ParseResult{}, nil
|
||||
}
|
||||
53
internal/nlp/fallback.go
Normal file
53
internal/nlp/fallback.go
Normal file
@ -0,0 +1,53 @@
|
||||
package nlp
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"gitcode.com/JianFeeeee/HomeAgent/internal/memory"
|
||||
)
|
||||
|
||||
// fallbackParser 使用 gojieba 分词 + POS 做降级句法分析
|
||||
// 返回解析结果中只填充 Tokens 和 POS,Heads/DepRels 留空
|
||||
type fallbackParser struct{}
|
||||
|
||||
func newFallbackParser() *fallbackParser {
|
||||
return &fallbackParser{}
|
||||
}
|
||||
|
||||
func (p *fallbackParser) Parse(text string) (*ParseResult, error) {
|
||||
if text == "" {
|
||||
return &ParseResult{}, nil
|
||||
}
|
||||
|
||||
x := memory.GetJieba()
|
||||
if x == nil {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
tagged := x.Tag(text)
|
||||
|
||||
var tokens, pos []string
|
||||
for _, t := range tagged {
|
||||
// Tag() 返回 "word/POS" 格式
|
||||
idx := strings.LastIndex(t, "/")
|
||||
if idx < 0 {
|
||||
continue
|
||||
}
|
||||
word := t[:idx]
|
||||
tag := t[idx+1:]
|
||||
if word == "" {
|
||||
continue
|
||||
}
|
||||
tokens = append(tokens, word)
|
||||
pos = append(pos, tag)
|
||||
}
|
||||
|
||||
if len(tokens) == 0 {
|
||||
return &ParseResult{}, nil
|
||||
}
|
||||
|
||||
return &ParseResult{
|
||||
Tokens: tokens,
|
||||
POS: pos,
|
||||
}, nil
|
||||
}
|
||||
25
internal/nlp/model.go
Normal file
25
internal/nlp/model.go
Normal file
@ -0,0 +1,25 @@
|
||||
package nlp
|
||||
|
||||
// ParseResult 依存句法分析结果
|
||||
type ParseResult struct {
|
||||
Tokens []string
|
||||
POS []string
|
||||
Heads []int // 父节点索引,0=ROOT
|
||||
DepRels []string // 依存关系标签
|
||||
}
|
||||
|
||||
// Triple 三元组 (subject, relation, object)
|
||||
type Triple struct {
|
||||
Subject string
|
||||
Relation string
|
||||
Object string
|
||||
Score float64
|
||||
Src string // "dep" / "fallback"
|
||||
}
|
||||
|
||||
// TripleSet 提取结果
|
||||
type TripleSet struct {
|
||||
Triples []Triple
|
||||
Src string // "dep_parser" / "fallback" / ""
|
||||
Err error
|
||||
}
|
||||
28
internal/nlp/onnx_stub.go
Normal file
28
internal/nlp/onnx_stub.go
Normal file
@ -0,0 +1,28 @@
|
||||
//go:build !onnxruntime
|
||||
|
||||
package nlp
|
||||
|
||||
import "fmt"
|
||||
|
||||
// ONNXParserStub 占位 — 编译时未启用 onnxruntime
|
||||
type ONNXParser struct{}
|
||||
|
||||
type ONNXConfig struct {
|
||||
ModelPath string
|
||||
VocabPath string
|
||||
POSVocPath string
|
||||
}
|
||||
|
||||
func NewONNXParser(cfg ONNXConfig) (*ONNXParser, error) {
|
||||
return nil, fmt.Errorf("onnxparser: build with -tags onnxruntime to enable")
|
||||
}
|
||||
|
||||
func (p *ONNXParser) Close() {}
|
||||
|
||||
func (p *ONNXParser) Parse(text string) (*ParseResult, error) {
|
||||
return nil, fmt.Errorf("onnxparser: not available (build with -tags onnxruntime)")
|
||||
}
|
||||
|
||||
func (p *ONNXParser) EnsureModel(dataDir string) error {
|
||||
return fmt.Errorf("onnxparser: not available")
|
||||
}
|
||||
121
internal/nlp/parser.go
Normal file
121
internal/nlp/parser.go
Normal file
@ -0,0 +1,121 @@
|
||||
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
|
||||
}
|
||||
Reference in New Issue
Block a user