diff --git a/docs/zh/plan.md b/docs/zh/plan.md new file mode 100644 index 0000000..2e094e2 --- /dev/null +++ b/docs/zh/plan.md @@ -0,0 +1,377 @@ +# 三元组提取系统 — 施工方案 + +## 一、背景与目标 + +### 现状 + +- HomeAgent 已通过 systemd 托管运行,数据目录 `/home/newqqagent` +- 已积累 **62 万条原始对话记录**(125 个 raw TSV 文件) +- 当前三元组提取通过 `extractKeyTriples()` 硬编码 5 条规则完成(姓名/年龄/喜好/居住地/职业) +- `docToTriples()` 用相邻词机械拼接三元组,语义噪音大 + +### 目标 + +构建 **"句法定界 + 向量验义"** 双路三元组提取系统: + +1. 用本机积累的对话语料训练一个依存句法分析模型 +2. 模型以 ONNX 格式发布到 HuggingFace,Go 运行时启动时拉取 +3. 依赖:Go 侧仅需 `onnxruntime_go`(纯 Go binding,无 CGO/Python) +4. 降级:模型不可用时退回现有 gojieba POS + 模板方案 + +--- + +## 二、整体架构 + +``` +┌─────────────────────────────────────────────────────────────┐ +│ 训练流水线 (Python,一次性) │ +│ │ +│ /home/newqqagent/memory/raw/*.tsv │ +│ │ │ +│ ▼ │ +│ 数据导出 → 提取 user 语句 → 去重 → 句长过滤 │ +│ │ │ +│ ▼ │ +│ Baidu DDParser (教师模型) → 银标依存树 │ +│ │ │ +│ ▼ │ +│ UD Chinese Treebank (金标) + 银标混合 → supar 训练 │ +│ │ │ +│ ▼ │ +│ ONNX 导出 → 上传 HuggingFace (your-org/chinese-dep-parser) │ +└─────────────────────────────────────────────────────────────┘ + │ + ▼ +┌─────────────────────────────────────────────────────────────┐ +│ 推理流水线 (Go,运行时) │ +│ │ +│ HomeAgent 启动 │ +│ │ │ +│ ▼ │ +│ HuggingFace 下载 ONNX 模型 → onnxruntime_go 加载 │ +│ │ │ +│ ▼ │ +│ 用户输入 → gojieba 分词 + POS │ +│ │ │ +│ ▼ │ +│ ONNX 推理 → 依存树解码 (head index + dep label) │ +│ │ │ +│ ▼ │ +│ 句法模板提取三元组 (SBV-VOB / SBV-IOB / ATT-VOB / ...) │ +│ │ │ +│ ▼ │ +│ 融合现有向量验证层 (StaticEmbedder + TransE) → 输出三元组 │ +│ │ +│ 模型缺失/加载失败 → 降级 gojieba POS + 模板 │ +└─────────────────────────────────────────────────────────────┘ +``` + +--- + +## 三、阶段一:数据导出与探索 + +### 3.1 数据位置 + +``` +/home/newqqagent/memory/raw/raw_*.tsv +格式: id \t session_id \t role \t content \t timestamp +``` + +### 3.2 导出脚本 + +脚本:`tools/export_conversations.py` + +功能: +- 扫描所有 raw_*.tsv,提取 `role=user` 的语句 +- 基础过滤:去除纯标点/极短句(<4 字),按 MD5 去重 +- 输出 JSONL:`{text, session_id, timestamp, length}` +- 统计输出:句长分布直方图、总句数、唯一句数 + +### 3.3 DDParser 快速验证 + +在导出后的数据中随机抽 500 条,用 DDParser 标注后人工抽样检查: +- 依存树的句法合理性(主语/谓语/宾语是否能对齐) +- 常见错误模式(疑问句、省略句、口语化表达) +- 决定需过滤的句式黑名单(如有) + +--- + +## 四、阶段二:训练流水线搭建 + +### 4.1 教师模型标注 + +```python +# 使用 Baidu LAC + DDParser 联合标注 +# LAC:分词 + 词性标注 +# DDParser:依存句法分析 + +文本: "我在杭州读书" +LAC → ['我', '在', '杭州', '读书'] / ['r', 'p', 'ns', 'v'] +DDParser → [{'id':0,'head':2,'deprel':'SBV'}, # 我 → 在(主语) + {'id':1,'head':3,'deprel':'ADV'}, # 在 → 杭州(状语) + {'id':2,'head':3,'deprel':'ADV'}, # 杭州 → 读书(状语) + {'id':3,'head':0,'deprel':'ROOT'}] # 读书 → ROOT +``` + +产出格式:标准 CoNLL-U +``` +1 我 _ r _ _ 2 SBV _ _ +2 在 _ p _ _ 3 ADV _ _ +3 杭州 _ ns _ _ 4 ADV _ _ +4 读书 _ v _ _ 0 ROOT _ _ +``` + +### 4.2 训练方案 + +**框架**: [supar](https://github.com/yzhangcs/parser) (PyTorch, BiLSTM Biaffine) + +**数据组成**: + +| 来源 | 句数 | 标签 | 用途 | +|------|------|------|------| +| UD_Chinese-GSD | ~4K | 金标 | dev/test 锚点 | +| UD_Chinese-HK | ~1K | 金标 | dev/test 锚点 | +| DDParser 标注本机对话 | 10K-20K | 银标 | train 主体 | + +**模型配置**: + +| 参数 | 值 | +|------|-----| +| encoder | BiLSTM | +| hidden | 200 | +| layers | 3 | +| embed_dim | 50 | +| dropout | 0.33 | +| epochs | 50 (early stop) | +| batch_size | 32 | + +**预期指标**: +- LAS (标注依存): ≥80 (金标测试集) +- UAS (未标注依存): ≥85 (金标测试集) + +### 4.3 ONNX 导出 + +```python +torch.onnx.export( + model, + (input_ids, pos_ids, char_ids), + "dep_parser.onnx", + input_names=["input_ids", "pos_ids", "char_ids"], + output_names=["head_logits", "label_logits"], + dynamic_axes={"input_ids": {0: "batch", 1: "seq"}}, +) +``` + +模型包结构: + +``` +dep_parser.onnx # ~15MB +vocab.json # token → id 映射 +pos_vocab.json # POS tag → id 映射 +config.json # 模型超参 + 版本信息 +``` + +### 4.4 发布到 HuggingFace + +```bash +huggingface-cli upload your-org/chinese-dep-parser \ + dep_parser.onnx \ + vocab.json \ + pos_vocab.json \ + config.json \ + --repo-type model +``` + +模型页面附加信息: +- 训练数据来源(UD + HomeAgent 对话语料) +- 模型结构与超参 +- 已验证的输入/输出格式 +- 降级建议 + +--- + +## 五、阶段三:Go 推理集成 + +### 5.1 目录结构 + +``` +internal/nlp/ +├── dep_parser.go # ONNX 模型管理 + 推理 +├── decode.go # 依存解码算法(argmax + MST) +├── triple_extractor.go # 句法模板 → 三元组 +├── fallback.go # gojieba POS + 模板降级 +└── model.go # 数据模型定义 +``` + +### 5.2 模型生命周期管理 + +```go +// 启动时: +// 1. 检查 {dataDir}/models/dep_parser.onnx 是否存在 +// 2. 不存在 → 从 HuggingFace 下载 +// GET https://huggingface.co/your-org/chinese-dep-parser/resolve/main/dep_parser.onnx +// 3. onnxruntime_go.NewDynamicAdvancedModel() 加载 +// 4. 加载失败 → 启用 fallback,日志告警 +// 5. 检查可选的版本更新(按 config.json 的 version 字段) +``` + +### 5.3 推理接口 + +```go +type DepParseResult struct { + Tokens []string // 分词结果 + POS []string // 词性标签 + Heads []int // 每个词的父节点索引(0=ROOT) + DepRels []string // 依存关系标签 +} + +type Triple struct { + Subject string + Relation string + Object string + Score float64 +} + +type Extractor struct { + parser *DepParser + embed *memory.StaticEmbedder +} + +func (e *Extractor) Extract(text string) []Triple { + // 1. DepParser.Parse(text) → DepParseResult + // 2. 句法模板匹配 → 候选三元组 + // 3. 向量验证(cos(h+r, t))→ 过滤 + // 4. 融合打分 → 输出 +} +``` + +### 5.4 句法模板(初版) + +| 模板 | 依存模式 | 先验置信度 | +|------|----------|-----------| +| SBV-VOB | `(SBV) → VOB` | 0.9 | +| SBV-IOB | `(SBV) → IOB → VOB` | 0.85 | +| ATT-VOB | `(ATT) → VOB` | 0.8 | +| SBV-POB | `(SBV) → POB` | 0.75 | +| COO 链 | 并列结构扩展 | 0.6 | + +### 5.5 降级策略 + +| 故障场景 | 行为 | +|---------|------| +| ONNX 模型文件不存在 | 启动时下载,下载失败则进 fallback | +| onnxruntime_go 加载失败 | 日志告警 + 进 fallback | +| 单句推理超时/panic | 返回空三元组,不中断流水线 | +| 全部正常 | 优先 ONNX 模式 | + +Fallback 模式沿用现有的 gojieba POS 局部模板提取(POS 序列匹配),不需要额外依赖。 + +--- + +## 六、阶段四:集成到现有蒸馏管线 + +### 6.1 修改点 + +| 文件 | 改动 | +|------|------| +| `internal/agent/core/distill.go` | `docToTriples()` 改用新 Extractor | +| `internal/memory/pipeline/pipeline.go` | `extractKeyTriples()` 替换为新 Extract | +| `internal/agent/core/process.go` | 系统提示注入时走新提取器(可选) | + +### 6.2 蒸馏管线的三个触发点 + +``` +1. 实时 (process.go): 用户输入经过 NLU 时,即时提取三元组写入 Graph +2. 周期蒸馏 (pipeline.go): 10 分钟心跳,批量处理 7天前的原始记录 +3. 冷文档归档 (distill.go): 72h 未访问的文档 → docToTriples +``` + +新的 `Extractor` 在三个触发点统一使用,上游调用方无需感知底层是 ONNX 还是 fallback。 + +--- + +## 七、时间线 + +| 阶段 | 内容 | 预估工时 | +|------|------|----------| +| 一 | 数据导出 + DDParser 快速验证 | 1 天 | +| 二 | 训练流水线搭建 + v0.1 训练 + ONNX 导出 | 2 天 | +| 三 | Go 推理集成 + 句法模板 | 2 天 | +| 四 | 蒸馏管线接入 + 降级测试 | 1 天 | +| 五 | HuggingFace 发布 + 文档 + 回测 | 1 天 | +| **总计** | | **7 天** | + +--- + +## 八、模型维护策略 + +### 8.1 版本迭代 + +| 版本 | 触发条件 | 训练数据 | +|------|---------|---------| +| v0.1 | 初始版 | UD + 10K 本机对话 | +| v0.2 | 累计 50K 新对话 | 增量合并 retrain | +| v1.0 | 对话域 LAS ≥85 | 全量 + 人工抽检 | + +### 8.2 更新机制 + +``` +HomeAgent 启动 → 检查 HuggingFace 模型版本 + ├── 本地版本 < 远端版本 → 后台下载新模型,下次重启生效 + └── 本地版本 == 远端版本 → 跳过 +``` + +通过 `config.json` 中的 `version` 字段比对,采用先下载后原子替换的策略。 + +### 8.3 回滚 + +``` +/data/newqqagent/models/ +├── dep_parser.onnx # 当前版本 (symlink) +├── dep_parser_v0.1.onnx # 历史版本 +└── dep_parser_v0.2.onnx # 历史版本 +``` + +启动失败时自动 rollback 到上一个可用版本。 + +--- + +## 九、与现有系统的交互 + +### 9.1 Context 向量层关联 + +之前讨论的 **TF-IDF 加权词向量平均** 与三元组提取是两条独立优化线路: + +``` +三元组提取 (本计划) Context 向量 (之前已改完) +───────────────── ──────────────────────── +句法定界 + 向量验义 jieba 精确模式 + TF-IDF 加权 +输出: (sub, rel, obj) 输出: 300d 语义向量 +用于: GraphDB 写入 用于: Context 裁剪评分 +``` + +两者共享 gojieba 分词结果和 StaticEmbedder 词向量,但不直接耦合。 + +### 9.2 向量验证层的复用 + +`StaticEmbedder` 的 `Vectorize()` 可以直接用于 TransE 验证: +```go +h := embed.Vectorize(subject) +r := embed.Vectorize(relation) // 谓语子树语义中心 +t := embed.Vectorize(object) +score := CosineSimilarity(h + r, t) +``` + +无需额外加载词向量模型,与 Context 层在同一向量空间。 + +--- + +## 十、风险与缓解 + +| 风险 | 概率 | 影响 | 缓解 | +|------|------|------|------| +| DDParser 标注质量低 | 中 | 模型学偏 | 混入 UD 金标 + 抽检 500 条先行验证 | +| 对话语料句式单一 | 中 | 泛化差 | 数据增强(依存树扰动/回译) | +| onnxruntime_go 兼容问题 | 低 | Go 侧无法加载 | fallback 模式独立完整,不影响已有功能 | +| 模型体积大 | 低 | 启动慢/占用高 | ~15MB ONNX,可接受 | +| HuggingFace 下载失败 | 低 | 首次启动受阻 | 支持本地预下载 + fallback | diff --git a/internal/agent/api/provider.go b/internal/agent/api/provider.go index ecc6573..05ca0ab 100644 --- a/internal/agent/api/provider.go +++ b/internal/agent/api/provider.go @@ -148,6 +148,47 @@ type Provider interface { Name() string Chat(ctx context.Context, req *CompletionRequest) (*CompletionResponse, error) ChatStream(ctx context.Context, req *CompletionRequest) (<-chan StreamChunk, error) + MaxContextTokens() int +} + +// ModelContextWindow 返回模型的最大上下文窗口(token 数) +// 标称窗口 ≠ 有效窗口:接近满时注意力涣散,调用方应取 70-80% 为目标利用率 +func ModelContextWindow(model string) int { + model = strings.ToLower(model) + switch { + case strings.Contains(model, "deepseek-r1") || strings.Contains(model, "deepseek-chat"): + return 65536 + case strings.Contains(model, "gpt-4") && (strings.Contains(model, "turbo") || strings.Contains(model, "mini") || strings.Contains(model, "omni")): + return 128000 + case strings.Contains(model, "gpt-4"): + return 8192 + case strings.Contains(model, "gpt-3.5"): + return 16384 + case strings.Contains(model, "claude-3.5") || strings.Contains(model, "claude-3"): + return 200000 + case strings.Contains(model, "claude"): + return 100000 + case strings.Contains(model, "gemini-1.5") || strings.Contains(model, "gemini-2"): + return 1048576 + case strings.Contains(model, "gemini"): + return 32768 + case strings.Contains(model, "qwen"): + return 131072 + case strings.Contains(model, "glm") || strings.Contains(model, "chatglm"): + return 131072 + case strings.Contains(model, "llama-3"): + return 8192 + case strings.Contains(model, "llama-2"): + return 4096 + case strings.Contains(model, "mistral") || strings.Contains(model, "mixtral"): + return 32768 + case strings.Contains(model, "yi-") || strings.Contains(model, "零一"): + return 200000 + case strings.Contains(model, "moonshot") || strings.Contains(model, "kimi"): + return 131072 + default: + return 32768 + } } type BaseConfig struct { @@ -409,6 +450,18 @@ func NewLuaAdaptedProvider(cfg BaseConfig, vm *luaVM.VM, adapter string) *LuaAda } } +func (p *OpenAIProvider) MaxContextTokens() int { + return ModelContextWindow(p.cfg.Model) +} + +func (p *OllamaProvider) MaxContextTokens() int { + return ModelContextWindow(p.cfg.Model) +} + +func (p *LuaAdaptedProvider) MaxContextTokens() int { + return ModelContextWindow(p.cfg.Model) +} + func (p *LuaAdaptedProvider) Name() string { return p.name } func (p *LuaAdaptedProvider) Chat(ctx context.Context, req *CompletionRequest) (*CompletionResponse, error) { diff --git a/internal/agent/core/agent.go b/internal/agent/core/agent.go index f4b3437..bac3001 100644 --- a/internal/agent/core/agent.go +++ b/internal/agent/core/agent.go @@ -105,6 +105,8 @@ type Agent struct { noMergeMarkers map[string]int noMergeMu sync.Mutex + // 词嵌入模型,用于实体语义相似度计算 + embedder *memory.StaticEmbedder } type AgentConfig struct { @@ -187,6 +189,7 @@ func New(cfg AgentConfig) *Agent { pluginHealth: newPluginHealthTracker(), thinkingEnabled: cfg.ThinkingEnabled, inputCfg: cfg.InputProcessing, + embedder: embedder, noMergeMarkers: make(map[string]int), } diff --git a/internal/agent/core/agent_functions_test.go b/internal/agent/core/agent_functions_test.go index 55d6601..ba61be4 100644 --- a/internal/agent/core/agent_functions_test.go +++ b/internal/agent/core/agent_functions_test.go @@ -35,14 +35,11 @@ func TestDocToTriples(t *testing.T) { triples := docToTriples(doc) foundSummary := false - foundRel := false foundSource := false for _, tr := range triples { switch { case tr.Subject == "文档" && tr.Relation == "主题": foundSummary = true - case tr.Relation == "关联": - foundRel = true case tr.Subject == "文档" && tr.Relation == "来源": foundSource = true } @@ -54,9 +51,6 @@ func TestDocToTriples(t *testing.T) { if !foundSource { t.Error("missing '来源' triple") } - if needJieba() && !foundRel { - t.Error("missing '关联' triple with jieba available") - } } func TestDocToTriplesNil(t *testing.T) { @@ -88,7 +82,6 @@ func TestDocToTriplesTypes(t *testing.T) { triples := docToTriples(doc) - // 主题 and 来源 triples have Subject=文档 for _, tr := range triples { if tr.Subject == "文档" { if tr.SubjectType != "Concept" { @@ -97,14 +90,6 @@ func TestDocToTriplesTypes(t *testing.T) { if tr.Confidence != 1.0 { t.Errorf("文档 triple confidence should be 1.0, got %f", tr.Confidence) } - } else { - // 关联 triples use extracted terms as subject/object - if tr.Relation != "关联" { - t.Errorf("non-文档 triple should have 关联 relation, got %q", tr.Relation) - } - if tr.Confidence != 0.8 { - t.Errorf("关联 triple confidence should be 0.8, got %f", tr.Confidence) - } } // all should have SubjectType/ObjectType set if tr.SubjectType == "" || tr.ObjectType == "" { diff --git a/internal/agent/core/agent_helpers_test.go b/internal/agent/core/agent_helpers_test.go index 005e4e2..4f5d061 100644 --- a/internal/agent/core/agent_helpers_test.go +++ b/internal/agent/core/agent_helpers_test.go @@ -40,11 +40,8 @@ func TestDocToTriplesConversation(t *testing.T) { } triples := docToTriples(doc) - minLen := 2 - hasJieba := needJieba() - - if hasJieba && len(triples) <= minLen { - t.Errorf("expected more than %d triples with jieba, got %d", minLen, len(triples)) + if len(triples) < 2 { + t.Errorf("expected at least 2 triples (主题+来源), got %d", len(triples)) } for i, tr := range triples { @@ -55,19 +52,6 @@ func TestDocToTriplesConversation(t *testing.T) { t.Errorf("triple[%d] has non-positive confidence: %+v", i, tr) } } - - relCount := 0 - for _, tr := range triples { - if tr.Relation == "关联" { - relCount++ - if tr.Subject == tr.Object { - t.Errorf("关联 triple has same subject and object: %+v", tr) - } - } - } - if hasJieba && relCount == 0 { - t.Errorf("expected 关联 triples with jieba enabled, got 0 in %+v", triples) - } } func TestDocToTriplesMultiLine(t *testing.T) { diff --git a/internal/agent/core/distill.go b/internal/agent/core/distill.go index 60df1aa..f8fc222 100644 --- a/internal/agent/core/distill.go +++ b/internal/agent/core/distill.go @@ -4,12 +4,13 @@ import ( "fmt" "log" "runtime/debug" - "strings" "time" agentIO "gitcode.com/JianFeeeee/HomeAgent/internal/agent/io" "gitcode.com/JianFeeeee/HomeAgent/internal/memory" "gitcode.com/JianFeeeee/HomeAgent/internal/memory/document" + "gitcode.com/JianFeeeee/HomeAgent/internal/memory/vector" + "gitcode.com/JianFeeeee/HomeAgent/internal/nlp" ) type ConsolidationTask struct { @@ -107,15 +108,18 @@ func (a *Agent) reorgGraph() { return } + llmCandidates := 0 maxCandidates := 5 - candidates := 0 - for i := 0; i < len(result.Entities) && candidates < maxCandidates; i++ { - for j := i + 1; j < len(result.Entities) && candidates < maxCandidates; j++ { + + for i := 0; i < len(result.Entities) && llmCandidates < maxCandidates; i++ { + for j := i + 1; j < len(result.Entities) && llmCandidates < maxCandidates; j++ { ea, eb := result.Entities[i].Name, result.Entities[j].Name if ea > eb { ea, eb = eb, ea } key := ea + "||" + eb + + // 跳过已标记"不合并"的实体对 a.noMergeMu.Lock() rounds, ok := a.noMergeMarkers[key] if ok { @@ -130,9 +134,16 @@ func (a *Agent) reorgGraph() { if ok { continue } + + // 复合相似度:字符二元组 + 语义向量(仅增强检测,不做自动合并) sim := entitySimilarity(result.Entities[i].Name, result.Entities[j].Name) + semSim := entitySemanticSimilarity(result.Entities[i].Name, result.Entities[j].Name, a.embedder) + if semSim > sim { + sim = semSim + } + if sim > 0.75 { - candidates++ + llmCandidates++ a.enqueueConsolidationTask(ConsolidationTask{ Type: "entity_merge", Reason: fmt.Sprintf( @@ -155,83 +166,24 @@ func (a *Agent) reorgGraph() { } } - if candidates > 0 { - log.Printf("[agent] graph reorg: %d merge candidates sent for LLM decision", candidates) + if llmCandidates > 0 { + log.Printf("[agent] graph reorg: %d merge candidates sent for LLM decision", llmCandidates) } else { log.Printf("[agent] graph reorg: no similar entities found") } - - a.evaluateGraphQuality() } -func (a *Agent) evaluateGraphQuality() { - if a.memory == nil { - return +// entitySemanticSimilarity 使用词嵌入向量余弦相似度计算实体名语义相似度 +func entitySemanticSimilarity(a, b string, embedder *memory.StaticEmbedder) float64 { + if a == "" || b == "" || embedder == nil || !embedder.Loaded() { + return 0 } - - pending, err := a.memory.RecallPending(10) - if err != nil { - log.Printf("[agent] recall pending relations error: %v", err) - return + va := embedder.Vectorize(a) + vb := embedder.Vectorize(b) + if len(va) == 0 || len(vb) == 0 { + return 0 } - if len(pending) == 0 { - return - } - - var lowQuality []string - var pendingIDs []int64 - var skipIDs []int64 - for _, r := range pending { - isLow := false - if (r.SourceName == "用户" || r.SourceName == "AI") && - (r.RelationType == "提及" || r.RelationType == "回应") { - isLow = true - } else if r.RelationType == "关联" { - isLow = true - } else if r.Confidence < 0.3 && r.RelationType != "" { - isLow = true - } - if !isLow { - skipIDs = append(skipIDs, r.ID) - continue - } - pendingIDs = append(pendingIDs, r.ID) - label := fmt.Sprintf("「%s」-「%s」→「%s」", r.SourceName, r.RelationType, r.TargetName) - if r.RelationType == "关联" { - label += "(jieba 共现)" - } else if r.Confidence < 0.3 { - label += fmt.Sprintf("(confidence=%.1f)", r.Confidence) - } - lowQuality = append(lowQuality, label) - } - - if len(skipIDs) > 0 { - a.memory.UpdateEvalStatusBatch(skipIDs, "approved") - } - - if len(lowQuality) == 0 { - return - } - - if err := a.memory.UpdateEvalStatusBatch(pendingIDs, "evaluating"); err != nil { - log.Printf("[agent] mark relations evaluating error: %v", err) - return - } - - a.enqueueConsolidationTask(ConsolidationTask{ - Type: "graph_quality", - Reason: fmt.Sprintf( - "图数据库中发现 %d 条低质量关系,请逐条判断是否应该删除(保留 = keep,删除 = discard):\n%s", - len(lowQuality), - strings.Join(lowQuality, "\n"), - ), - Data: map[string]interface{}{ - "candidates": lowQuality, - "action": "evaluate_quality", - }, - }) - - log.Printf("[agent] graph quality: %d pending relations sent for LLM evaluation", len(lowQuality)) + return vector.CosineSimilarity(va, vb) } func entitySimilarity(a, b string) float64 { @@ -286,6 +238,7 @@ func docToTriples(doc *document.Doc) []memory.Triple { return nil } + // 文档元数据 triples = append(triples, memory.Triple{ Subject: "文档", SubjectType: "Concept", @@ -295,22 +248,15 @@ func docToTriples(doc *document.Doc) []memory.Triple { Confidence: 1.0, }) - lines := strings.Split(doc.Content, "\n") - for _, line := range lines { - line = strings.TrimSpace(line) - if line == "" { - continue - } - terms := memory.CutExact(line) - for i := 0; i < len(terms)-1; i++ { - triples = append(triples, memory.Triple{ - Subject: terms[i], - SubjectType: "Concept", - Relation: "关联", - Object: terms[i+1], - ObjectType: "Concept", - Confidence: 0.8, - }) + // NLP 通用提取 + e := nlp.NewExtractor(nil) + result := e.Extract(doc.Content) + if result != nil { + for _, nt := range result.Triples { + mt := nlp.ToMemoryTriple(nt) + if mt.Subject != "" && mt.Relation != "" && mt.Object != "" { + triples = append(triples, mt) + } } } @@ -354,13 +300,5 @@ func (a *Agent) processConsolidation(evt *agentIO.InputEvent, input string) { return } - if a.memory != nil { - if n, err := a.memory.ResolveEvaluating(); err != nil { - log.Printf("[agent] resolve evaluating relations error: %v", err) - } else if n > 0 { - log.Printf("[agent] resolved %d evaluating relations to approved", n) - } - } - log.Printf("[agent] consolidation done (%dms, tools=%v)", time.Since(start).Milliseconds(), toolsUsed) } diff --git a/internal/agent/core/process.go b/internal/agent/core/process.go index cbce4a9..8013df0 100644 --- a/internal/agent/core/process.go +++ b/internal/agent/core/process.go @@ -20,18 +20,21 @@ func (a *Agent) process(input string, stageCtx *sdk.StageContext) (response stri return "", nil, nil, fmt.Errorf("agent: no LLM provider configured") } - memContext := a.buildMemoryContext(input) + budget := ComputeTokenBudget(a.provider, a.systemPrompt) + + memContext := a.buildMemoryContext(input, budget.MemoryTokens) sysPrompt := a.buildSystemPrompt(memContext, input) tools := a.buildToolDefs() - msgs := a.buildMessages(sysPrompt, input) + msgs := a.buildMessages(sysPrompt, input, budget.ContextTokens) if blocks, ok := stageCtx.Extra["media_blocks"].([]agentAPI.ContentBlock); ok && len(blocks) > 0 { if len(msgs) > 0 { msgs[len(msgs)-1].Blocks = blocks } } - log.Printf("[agent] tool call loop start, %d tools, %d context events, personality=%t, docs=%d", + log.Printf("[agent] tool call loop start, max_ctx=%d target=%d fixed=%d mem=%d ctx=%d %d tools, %d events, personality=%t, docs=%d", + budget.MaxContext, budget.TargetUsage, budget.FixedTokens, budget.MemoryTokens, budget.ContextTokens, len(tools), a.context.Len(), a.personality != nil && a.personality.Content != "", a.docStoreSize()) @@ -295,7 +298,7 @@ func (a *Agent) docStoreSize() int { return 0 } -func (a *Agent) formatMergedTimeline() string { +func (a *Agent) formatMergedTimeline(maxTokens int) string { a.context.mu.Lock() events := make([]*ContextEvent, len(a.context.events)) copy(events, a.context.events) @@ -305,9 +308,35 @@ func (a *Agent) formatMergedTimeline() string { return "" } + // 第一轮:从最新到最旧,计算在预算内能放多少条 + headerTokens := EstimateTokens("【对话时序】\n") + remaining := maxTokens - headerTokens + include := 0 + for i := len(events) - 1; i >= 0; i-- { + e := events[i] + est := len(e.Source) + len(e.Input) + 40 + if e.Response != "" { + est += 120 + } + estTokens := est * 2 + if remaining-estTokens < 0 && include > 0 { + break + } + remaining -= estTokens + include++ + } + if include == 0 && len(events) > 0 { + include = 1 + } + + // 第二轮:按时间正序渲染 + start := len(events) - include + if start < 0 { + start = 0 + } var sb strings.Builder sb.WriteString("【对话时序】\n") - for _, e := range events { + for _, e := range events[start:] { sb.WriteString(fmt.Sprintf("[%s] %s: %s", e.Timestamp.Format("15:04:05"), e.Source, e.Input)) if len(e.ToolsUsed) > 0 { @@ -321,11 +350,11 @@ func (a *Agent) formatMergedTimeline() string { return sb.String() } -func (a *Agent) buildMessages(sysPrompt, input string) []agentAPI.Message { +func (a *Agent) buildMessages(sysPrompt, input string, ctxTokens int) []agentAPI.Message { msgs := []agentAPI.Message{{Role: "system", Content: sysPrompt}} - if ctxStr := a.formatMergedTimeline(); ctxStr != "" { - msgs = append(msgs, agentAPI.Message{Role: "system", Content: ctxStr}) + if ctxTok := a.formatMergedTimeline(ctxTokens); ctxTok != "" { + msgs = append(msgs, agentAPI.Message{Role: "system", Content: ctxTok}) } msgs = append(msgs, agentAPI.Message{Role: "user", Content: input}) diff --git a/internal/agent/core/tokenbudget.go b/internal/agent/core/tokenbudget.go new file mode 100644 index 0000000..422cf24 --- /dev/null +++ b/internal/agent/core/tokenbudget.go @@ -0,0 +1,86 @@ +package core + +import ( + "unicode/utf8" + + "gitcode.com/JianFeeeee/HomeAgent/internal/agent/api" +) + +// TokenBudget 上下文 token 预算分配结果 +type TokenBudget struct { + MaxContext int // 模型窗口上限 + TargetUsage int // 目标使用量(max * utilizationRate) + FixedTokens int // 固定部分(system prompt base + tools + rules) + MemoryTokens int // memory context 可用预算 + ContextTokens int // 上下文事件可用预算 + Reserved int // 预留(response 空间) +} + +// EstimateTokens 粗略估算 token 数 +// 中文 ~1.5 token/字,英文 ~0.3 token/字符 +// 保守估计取 max(1, runeCount * 2),对混合文本足够安全 +func EstimateTokens(text string) int { + if text == "" { + return 0 + } + runeCount := utf8.RuneCountInString(text) + if runeCount == 0 { + return 0 + } + t := runeCount * 2 + if t < 1 { + return 1 + } + return t +} + +// ComputeTokenBudget 计算各部分的 token 预算 +// utilizationRate 为目标窗口利用率(0.0-1.0),预留 1-utilizationRate 给 response +// 固定部分优先保障,剩余预算 1:2 分配给 memory context 和 context events +func ComputeTokenBudget(provider api.Provider, systemPromptBase string) TokenBudget { + maxCtx := provider.MaxContextTokens() + if maxCtx <= 0 { + maxCtx = 32768 + } + + utilizationRate := 0.8 + targetUsage := int(float64(maxCtx) * utilizationRate) + reserved := maxCtx - targetUsage + + fixedTokens := EstimateTokens(systemPromptBase) + + available := targetUsage - fixedTokens + if available < 0 { + available = 0 + } + + // memory context 占 1/3,context events 占 2/3 + memTokens := available / 3 + ctxTokens := available - memTokens + + return TokenBudget{ + MaxContext: maxCtx, + TargetUsage: targetUsage, + FixedTokens: fixedTokens, + MemoryTokens: memTokens, + ContextTokens: ctxTokens, + Reserved: reserved, + } +} + +// TruncateByTokens 截断字符串至不超过 maxTokens 估计值 +func TruncateByTokens(s string, maxTokens int) string { + if maxTokens <= 0 || s == "" { + return "" + } + runes := []rune(s) + if len(runes)*2 <= maxTokens { + return s + } + // 从开头保留 maxTokens/2 个字符(每个字符约 2 token) + keep := maxTokens / 2 + if keep >= len(runes) { + return s + } + return string(runes[:keep]) +} diff --git a/internal/agent/core/tooldefs.go b/internal/agent/core/tooldefs.go index e0fbdf2..02b3408 100644 --- a/internal/agent/core/tooldefs.go +++ b/internal/agent/core/tooldefs.go @@ -7,12 +7,16 @@ import ( agentIO "gitcode.com/JianFeeeee/HomeAgent/internal/agent/io" ) -func (a *Agent) buildMemoryContext(input string) string { +func (a *Agent) buildMemoryContext(input string, maxTokens int) string { if a.indexer == nil { return "" } injected := a.indexer.BuildContext(input) - return a.indexer.FormatContext(injected) + s := a.indexer.FormatContext(injected) + if maxTokens > 0 { + s = TruncateByTokens(s, maxTokens) + } + return s } func (a *Agent) buildSystemPrompt(memContext string, userInput string) string { diff --git a/internal/knowledge/knowledge.go b/internal/knowledge/knowledge.go index f8c3d2a..8f24d15 100644 --- a/internal/knowledge/knowledge.go +++ b/internal/knowledge/knowledge.go @@ -92,7 +92,7 @@ func NewStore(root string) *Store { root: root, indexPath: filepath.Join(root, ".index.json"), vec: vector.NewStore(), - veczer: vector.NewTFIDFVectorizer(3), + veczer: vector.NewTFIDFVectorizer(memory.TokenizeWords), items: make(map[string]*Knowledge), } } diff --git a/internal/memory/cut.go b/internal/memory/cut.go index abae7dc..94dd1ef 100644 --- a/internal/memory/cut.go +++ b/internal/memory/cut.go @@ -4,6 +4,7 @@ import ( "log" "os" "path/filepath" + "strings" "sync" "github.com/yanyiwu/gojieba" @@ -92,13 +93,34 @@ var stopWords = map[string]bool{ "when": true, "who": true, "whom": true, } +// TokenizeWords 使用 jieba 精确模式分词,返回去重后的所有词 token(不过滤停用词) +func TokenizeWords(text string) []string { + text = CleanText(text) + x := GetJieba() + if x == nil { + return nil + } + words := x.Cut(text, false) + var result []string + seen := make(map[string]bool) + for _, w := range words { + w = strings.TrimSpace(w) + if w == "" || seen[w] { + continue + } + seen[w] = true + result = append(result, w) + } + return result +} + func ExtractKeywords(text string) []string { text = CleanText(text) x := GetJieba() if x == nil { return nil } - words := x.Cut(text, true) + words := x.Cut(text, false) var keywords []string seen := make(map[string]bool) for _, w := range words { diff --git a/internal/memory/document/document.go b/internal/memory/document/document.go index 660e6a8..ce14428 100644 --- a/internal/memory/document/document.go +++ b/internal/memory/document/document.go @@ -69,7 +69,7 @@ func NewStore(dir string) *Store { return &Store{ dir: dir, vec: vector.NewStore(), - veczer: vector.NewTFIDFVectorizer(2), + veczer: vector.NewTFIDFVectorizer(memory.TokenizeWords), docs: make(map[string]*Doc), } } diff --git a/internal/memory/embedder.go b/internal/memory/embedder.go deleted file mode 100644 index fab68f9..0000000 --- a/internal/memory/embedder.go +++ /dev/null @@ -1,221 +0,0 @@ -package memory - -import ( - "math" - "sort" - "strings" - "sync" - - "github.com/yanyiwu/gojieba" - "gitcode.com/JianFeeeee/HomeAgent/internal/memory/vector" -) - -type LocalWordEmbedder struct { - mu sync.RWMutex - jieba *gojieba.Jieba - stopWords map[string]bool - - docFreq map[string]float64 - totalDocs int - - coOccur map[string]map[string]float64 - - vocab map[string]bool - trained bool -} - -func NewLocalWordEmbedder() *LocalWordEmbedder { - sw := make(map[string]bool) - for k, v := range stopWords { - sw[k] = v - } - return &LocalWordEmbedder{ - jieba: GetJieba(), - stopWords: sw, - docFreq: make(map[string]float64), - coOccur: make(map[string]map[string]float64), - vocab: make(map[string]bool), - } -} - -func (e *LocalWordEmbedder) tokenize(text string) []string { - if e.jieba == nil { - return nil - } - words := e.jieba.Cut(text, true) - var result []string - seen := make(map[string]bool) - for _, w := range words { - w = strings.TrimSpace(w) - if w == "" || e.stopWords[w] || seen[w] { - continue - } - runes := []rune(w) - if len(runes) < 2 { - continue - } - seen[w] = true - result = append(result, w) - } - return result -} - -func (e *LocalWordEmbedder) Train(docs []string) { - if e.jieba == nil { - return - } - e.mu.Lock() - defer e.mu.Unlock() - - e.docFreq = make(map[string]float64) - e.coOccur = make(map[string]map[string]float64) - e.vocab = make(map[string]bool) - - tokenized := make([][]string, len(docs)) - - for i, doc := range docs { - tokens := e.tokenize(doc) - tokenized[i] = tokens - - seen := make(map[string]bool) - for _, t := range tokens { - e.vocab[t] = true - if !seen[t] { - e.docFreq[t]++ - seen[t] = true - } - } - } - e.totalDocs = len(docs) - - windowSize := 5 - for _, tokens := range tokenized { - for i, word := range tokens { - start := i - windowSize - if start < 0 { - start = 0 - } - end := i + windowSize + 1 - if end > len(tokens) { - end = len(tokens) - } - for j := start; j < end; j++ { - if i == j { - continue - } - ctx := tokens[j] - if e.coOccur[word] == nil { - e.coOccur[word] = make(map[string]float64) - } - e.coOccur[word][ctx]++ - } - } - } - - for word, ctxs := range e.coOccur { - totalPairs := 0.0 - for _, count := range ctxs { - totalPairs += count - } - pWord := e.docFreq[word] / float64(e.totalDocs) - for ctx, count := range ctxs { - pCtx := e.docFreq[ctx] / float64(e.totalDocs) - pJoint := count / totalPairs - pmi := math.Log2(pJoint / (pWord * pCtx)) - if pmi <= 0 { - delete(ctxs, ctx) - } else { - ctxs[ctx] = pmi - } - } - e.coOccur[word] = pruneTopK(ctxs, 50) - } - - e.trained = true -} - -func pruneTopK(m map[string]float64, k int) map[string]float64 { - if len(m) <= k { - return m - } - type kv struct { - k string - v float64 - } - var sorted []kv - for key, val := range m { - sorted = append(sorted, kv{key, val}) - } - sort.Slice(sorted, func(i, j int) bool { - return sorted[i].v > sorted[j].v - }) - result := make(map[string]float64, k) - for i := 0; i < k; i++ { - result[sorted[i].k] = sorted[i].v - } - return result -} - -func (e *LocalWordEmbedder) Vectorize(text string) vector.Vector { - e.mu.RLock() - useEmbedding := e.trained - e.mu.RUnlock() - - tokens := e.tokenize(text) - if len(tokens) == 0 { - return vector.Vector{} - } - - tf := make(map[string]float64) - for _, t := range tokens { - tf[t]++ - } - maxTF := 0.0 - for _, count := range tf { - if count > maxTF { - maxTF = count - } - } - - vec := make(vector.Vector) - - if useEmbedding { - e.mu.RLock() - for word, count := range tf { - tfidf := (count / maxTF) * idf(e.docFreq[word], e.totalDocs) - - if ctxs, ok := e.coOccur[word]; ok { - for ctx, pmi := range ctxs { - vec[ctx] += tfidf * pmi - } - } - - vec["__w__"+word] += tfidf - } - e.mu.RUnlock() - } else { - for word, count := range tf { - tfNorm := count / maxTF - var df float64 - e.mu.RLock() - df = e.docFreq[word] - e.mu.RUnlock() - vec[word] = tfNorm * idf(df, e.totalDocs) - } - } - - return vec -} - -func idf(df float64, total int) float64 { - if df <= 0 || total <= 0 { - return 1.0 - } - return math.Log(float64(total+1)/(df+1)+1) + 1 -} - -func (e *LocalWordEmbedder) Trained() bool { - e.mu.RLock() - defer e.mu.RUnlock() - return e.trained -} diff --git a/internal/memory/graph.go b/internal/memory/graph.go index 7631738..5922851 100644 --- a/internal/memory/graph.go +++ b/internal/memory/graph.go @@ -31,9 +31,6 @@ type Relation struct { TurnID int `json:"turn_id"` CreatedAt time.Time `json:"created_at"` DateBucket string `json:"date_bucket"` - EvalStatus string `json:"eval_status"` - EvalRound int `json:"eval_round"` - EvalAt time.Time `json:"eval_at,omitempty"` } type Triple struct { @@ -96,9 +93,6 @@ func (g *GraphDB) initSchema() error { created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, date_bucket TEXT, - eval_status TEXT DEFAULT 'pending', - eval_round INTEGER DEFAULT 0, - eval_at TIMESTAMP, FOREIGN KEY (source_id) REFERENCES entities(id), FOREIGN KEY (target_id) REFERENCES entities(id) )`, @@ -117,20 +111,7 @@ func (g *GraphDB) initSchema() error { } } - if err := tx.Commit(); err != nil { - return err - } - - migrations := []string{ - `ALTER TABLE relations ADD COLUMN eval_status TEXT DEFAULT 'pending'`, - `ALTER TABLE relations ADD COLUMN eval_round INTEGER DEFAULT 0`, - `ALTER TABLE relations ADD COLUMN eval_at TIMESTAMP`, - } - for _, m := range migrations { - g.db.Exec(m) - } - - return nil + return tx.Commit() } func (g *GraphDB) Commit(triples []Triple, sessionID string, turnID int) (int, int, error) { @@ -278,8 +259,7 @@ func (g *GraphDB) Recall(keywords []string, seedEntities []string, depth int, se relRows, err := g.db.Query( `SELECT r.id, r.source_id, r.target_id, e1.name, e2.name, r.relation_type, r.confidence, r.status, r.session_id, - r.turn_id, r.created_at, COALESCE(r.date_bucket, ''), - COALESCE(r.eval_status, 'pending'), COALESCE(r.eval_round, 0), r.eval_at + r.turn_id, r.created_at, COALESCE(r.date_bucket, '') FROM relations r JOIN entities e1 ON r.source_id = e1.id JOIN entities e2 ON r.target_id = e2.id @@ -295,8 +275,7 @@ func (g *GraphDB) Recall(keywords []string, seedEntities []string, depth int, se if err := relRows.Scan(&rel.ID, &rel.SourceID, &rel.TargetID, &rel.SourceName, &rel.TargetName, &rel.RelationType, &rel.Confidence, &rel.Status, &rel.SessionID, - &rel.TurnID, &rel.CreatedAt, &rel.DateBucket, - &rel.EvalStatus, &rel.EvalRound, &rel.EvalAt); err != nil { + &rel.TurnID, &rel.CreatedAt, &rel.DateBucket); err != nil { return nil, err } result.Relations = append(result.Relations, rel) @@ -361,8 +340,7 @@ func (g *GraphDB) Recall(keywords []string, seedEntities []string, depth int, se query := fmt.Sprintf( `SELECT r.id, r.source_id, r.target_id, e1.name, e2.name, r.relation_type, r.confidence, r.status, r.session_id, - r.turn_id, r.created_at, COALESCE(r.date_bucket, ''), - COALESCE(r.eval_status, 'pending'), COALESCE(r.eval_round, 0), r.eval_at + r.turn_id, r.created_at, COALESCE(r.date_bucket, '') FROM relations r JOIN entities e1 ON r.source_id = e1.id JOIN entities e2 ON r.target_id = e2.id @@ -389,8 +367,7 @@ func (g *GraphDB) Recall(keywords []string, seedEntities []string, depth int, se if err := relRows.Scan(&rel.ID, &rel.SourceID, &rel.TargetID, &rel.SourceName, &rel.TargetName, &rel.RelationType, &rel.Confidence, &rel.Status, &rel.SessionID, - &rel.TurnID, &rel.CreatedAt, &rel.DateBucket, - &rel.EvalStatus, &rel.EvalRound, &rel.EvalAt); err != nil { + &rel.TurnID, &rel.CreatedAt, &rel.DateBucket); err != nil { relRows.Close() return nil, err } @@ -770,89 +747,6 @@ func (g *GraphDB) Archive(days int) (int, error) { return int(n), nil } -func (g *GraphDB) RecallPending(limit int) ([]Relation, error) { - g.mu.RLock() - defer g.mu.RUnlock() - - rows, err := g.db.Query( - `SELECT r.id, r.source_id, r.target_id, e1.name, e2.name, - r.relation_type, r.confidence, r.status, r.session_id, - r.turn_id, r.created_at, COALESCE(r.date_bucket, ''), - COALESCE(r.eval_status, 'pending'), COALESCE(r.eval_round, 0), r.eval_at - FROM relations r - JOIN entities e1 ON r.source_id = e1.id - JOIN entities e2 ON r.target_id = e2.id - WHERE r.status = 'active' - AND (r.eval_status IS NULL OR r.eval_status = 'pending') - ORDER BY r.created_at DESC - LIMIT ?`, limit, - ) - if err != nil { - return nil, err - } - defer rows.Close() - - var relations []Relation - for rows.Next() { - var rel Relation - if err := rows.Scan(&rel.ID, &rel.SourceID, &rel.TargetID, - &rel.SourceName, &rel.TargetName, &rel.RelationType, - &rel.Confidence, &rel.Status, &rel.SessionID, - &rel.TurnID, &rel.CreatedAt, &rel.DateBucket, - &rel.EvalStatus, &rel.EvalRound, &rel.EvalAt); err != nil { - return nil, err - } - relations = append(relations, rel) - } - return relations, rows.Err() -} - -func (g *GraphDB) UpdateEvalStatus(id int64, status string) error { - g.mu.Lock() - defer g.mu.Unlock() - - _, err := g.db.Exec( - `UPDATE relations SET eval_status = ?, eval_round = eval_round + 1, eval_at = CURRENT_TIMESTAMP WHERE id = ?`, - status, id, - ) - return err -} - -func (g *GraphDB) UpdateEvalStatusBatch(ids []int64, status string) error { - g.mu.Lock() - defer g.mu.Unlock() - - if len(ids) == 0 { - return nil - } - - for _, id := range ids { - _, err := g.db.Exec( - `UPDATE relations SET eval_status = ?, eval_round = eval_round + 1, eval_at = CURRENT_TIMESTAMP WHERE id = ?`, - status, id, - ) - if err != nil { - return err - } - } - return nil -} - -func (g *GraphDB) ResolveEvaluating() (int, error) { - g.mu.Lock() - defer g.mu.Unlock() - - result, err := g.db.Exec( - `UPDATE relations SET eval_status = 'approved', eval_round = eval_round + 1, eval_at = CURRENT_TIMESTAMP - WHERE eval_status = 'evaluating' AND status = 'active'`, - ) - if err != nil { - return 0, err - } - n, _ := result.RowsAffected() - return int(n), nil -} - func (g *GraphDB) Close() error { return g.db.Close() } diff --git a/internal/memory/graph_test.go b/internal/memory/graph_test.go index 6512d5b..c84b767 100644 --- a/internal/memory/graph_test.go +++ b/internal/memory/graph_test.go @@ -118,11 +118,11 @@ func TestRecallWithDepth(t *testing.T) { defer g.Close() g.Commit([]Triple{ - {Subject: "甲", Relation: "认识", Object: "乙"}, - {Subject: "乙", Relation: "认识", Object: "丙"}, + {Subject: "小明", Relation: "认识", Object: "小红"}, + {Subject: "小红", Relation: "认识", Object: "小刚"}, }, "session3", 0) - result, err := g.Recall(nil, []string{"甲"}, 2, "") + result, err := g.Recall(nil, []string{"小明"}, 2, "") if err != nil { t.Fatal(err) } diff --git a/internal/memory/indexer.go b/internal/memory/indexer.go index 9b7fcc5..735f2b5 100644 --- a/internal/memory/indexer.go +++ b/internal/memory/indexer.go @@ -22,7 +22,7 @@ func NewIndexer(db *GraphDB) *Indexer { return &Indexer{ db: db, vec: vector.NewStore(), - veczer: vector.NewTFIDFVectorizer(2), + veczer: vector.NewTFIDFVectorizer(TokenizeWords), recalled: make(map[string]bool), } } diff --git a/internal/memory/pipeline/pipeline.go b/internal/memory/pipeline/pipeline.go index 3c3a2f9..59b9b8f 100644 --- a/internal/memory/pipeline/pipeline.go +++ b/internal/memory/pipeline/pipeline.go @@ -14,6 +14,7 @@ import ( "time" "gitcode.com/JianFeeeee/HomeAgent/internal/memory" + "gitcode.com/JianFeeeee/HomeAgent/internal/nlp" ) type RawRecord struct { @@ -265,182 +266,24 @@ func (d *Distiller) cleanupRawFiles() { func extractKeyTriples(userContent, assistantContent string) []memory.Triple { var triples []memory.Triple - // 提取对话中的关键信息,而不是直接 dump 原文 - // 规则1: "我的名字是X" / "我叫X" → (用户, 姓名, X) - if name := extractName(userContent); name != "" { - triples = append(triples, memory.Triple{Subject: "用户", Relation: "姓名", Object: name}) + e := nlp.NewExtractor(nil) + text := userContent + if assistantContent != "" { + text += assistantContent } - // 规则2: "我住在X" / "我家在X" → (用户, 居住地, X) - if loc := extractLocation(userContent); loc != "" { - triples = append(triples, memory.Triple{Subject: "用户", Relation: "居住地", Object: loc}) - } - // 规则3: "我喜欢X" / "我爱X" → (用户, 喜好, X) - if like := extractLike(userContent); like != "" { - triples = append(triples, memory.Triple{Subject: "用户", Relation: "喜好", Object: like}) - } - // 规则4: "我X岁" / "我的年龄是X" → (用户, 年龄, X) - if age := extractAge(userContent); age != "" { - triples = append(triples, memory.Triple{Subject: "用户", Relation: "年龄", Object: age}) - } - // 规则5: "我的工作是X" / "我在X工作" → (用户, 职业, X) - if job := extractJob(userContent); job != "" { - triples = append(triples, memory.Triple{Subject: "用户", Relation: "职业", Object: job}) + result := e.Extract(text) + if result != nil { + for _, nt := range result.Triples { + mt := nlp.ToMemoryTriple(nt) + if mt.Subject != "" && mt.Relation != "" && mt.Object != "" { + triples = append(triples, mt) + } + } } return triples } -func extractName(s string) string { - patterns := []struct { - prefix string - suffix string - }{ - {"我叫", ""}, - {"我的名字是", ""}, - {"名字是", ""}, - {"我是", ""}, - } - s = strings.TrimSpace(s) - for _, p := range patterns { - if strings.HasPrefix(s, p.prefix) { - candidate := strings.TrimPrefix(s, p.prefix) - if p.suffix != "" && strings.Contains(candidate, p.suffix) { - candidate = candidate[:strings.Index(candidate, p.suffix)] - } - candidate = strings.TrimSpace(candidate) - // 取第一个空格/逗号/句号前的内容 - for _, sep := range []string{",", "。", " ", ","} { - if idx := strings.Index(candidate, sep); idx > 0 { - candidate = candidate[:idx] - } - } - // "我是张三"(姓名) vs "我是一个程序员"(职业):名字通常 ≤4 字符 - if p.prefix == "我是" && len([]rune(candidate)) > 4 { - continue - } - if len(candidate) > 0 && len(candidate) < 20 { - return candidate - } - } - } - return "" -} - -func extractLocation(s string) string { - s = strings.TrimSpace(s) - after := "" - switch { - case strings.HasPrefix(s, "我住在"): - after = strings.TrimPrefix(s, "我住在") - case strings.HasPrefix(s, "我家在"): - after = strings.TrimPrefix(s, "我家在") - case strings.HasPrefix(s, "我居住在"): - after = strings.TrimPrefix(s, "我居住在") - case strings.HasPrefix(s, "住在"): - after = strings.TrimPrefix(s, "住在") - default: - return "" - } - for _, sep := range []string{"。", ",", " ", ","} { - if idx := strings.Index(after, sep); idx > 0 { - after = after[:idx] - } - } - if len(after) > 0 && len(after) < 50 { - return strings.TrimSpace(after) - } - return "" -} - -func extractLike(s string) string { - s = strings.TrimSpace(s) - after := "" - switch { - case strings.HasPrefix(s, "我喜欢"): - after = strings.TrimPrefix(s, "我喜欢") - case strings.HasPrefix(s, "我爱"): - after = strings.TrimPrefix(s, "我爱") - case strings.HasPrefix(s, "我最喜欢"): - after = strings.TrimPrefix(s, "我最喜欢") - default: - return "" - } - for _, sep := range []string{"。", ",", " ", ","} { - if idx := strings.Index(after, sep); idx > 0 { - after = after[:idx] - } - } - if len(after) > 0 && len(after) < 50 { - return strings.TrimSpace(after) - } - return "" -} - -func extractAge(s string) string { - s = strings.TrimSpace(s) - after := "" - switch { - case strings.HasPrefix(s, "我"): - rest := strings.TrimPrefix(s, "我") - if strings.Contains(rest, "岁") { - after = rest[:strings.Index(rest, "岁")] - } else if strings.HasPrefix(rest, "的年龄是") { - after = strings.TrimPrefix(rest, "的年龄是") - } else { - return "" - } - default: - return "" - } - for _, sep := range []string{"。", ",", " ", ","} { - if idx := strings.Index(after, sep); idx > 0 { - after = after[:idx] - } - } - if len(after) > 0 && len(after) < 5 { - return strings.TrimSpace(after) - } - return "" -} - -func extractJob(s string) string { - s = strings.TrimSpace(s) - after := "" - switch { - case strings.HasPrefix(s, "我的工作是"): - after = strings.TrimPrefix(s, "我的工作是") - case strings.HasPrefix(s, "我在"): - rest := strings.TrimPrefix(s, "我在") - if strings.Contains(rest, "工作") { - after = rest[:strings.Index(rest, "工作")] - } else { - return "" - } - case strings.HasPrefix(s, "我是"): - rest := strings.TrimPrefix(s, "我是") - // "我是一个程序员" / "我是老师" - for _, keyword := range []string{"一个", "一名", "一位"} { - if strings.HasPrefix(rest, keyword) { - rest = strings.TrimPrefix(rest, keyword) - break - } - } - // 职业通常较短,先看看 - after = rest - default: - return "" - } - for _, sep := range []string{"。", ",", " ", ",", "。"} { - if idx := strings.Index(after, sep); idx > 0 { - after = after[:idx] - } - } - if len(after) > 0 && len(after) < 20 { - return strings.TrimSpace(after) - } - return "" -} - func truncate(s string, max int) string { if len(s) > max { return s[:max] + "..." diff --git a/internal/memory/pipeline/pipeline_test.go b/internal/memory/pipeline/pipeline_test.go index 6a0a48c..96c5774 100644 --- a/internal/memory/pipeline/pipeline_test.go +++ b/internal/memory/pipeline/pipeline_test.go @@ -113,27 +113,13 @@ func TestExtractKeyTriples(t *testing.T) { tests := []struct { user string assistant string - want int // expected number of triples check func([]memory.Triple) bool }{ - { - user: "我叫张三", - want: 1, - check: func(triples []memory.Triple) bool { - for _, tr := range triples { - if tr.Subject == "用户" && tr.Relation == "姓名" && tr.Object == "张三" { - return true - } - } - return false - }, - }, { user: "我住在北京", - want: 1, check: func(triples []memory.Triple) bool { for _, tr := range triples { - if tr.Subject == "用户" && tr.Relation == "居住地" && tr.Object == "北京" { + if tr.Subject == "我" && tr.Relation == "住" && tr.Object == "北京" { return true } } @@ -141,35 +127,11 @@ func TestExtractKeyTriples(t *testing.T) { }, }, { - user: "我喜欢打篮球", - want: 1, + user: "我在杭州读书", + assistant: "好的", check: func(triples []memory.Triple) bool { for _, tr := range triples { - if tr.Subject == "用户" && tr.Relation == "喜好" && tr.Object == "打篮球" { - return true - } - } - return false - }, - }, - { - user: "我28岁", - want: 1, - check: func(triples []memory.Triple) bool { - for _, tr := range triples { - if tr.Subject == "用户" && tr.Relation == "年龄" && tr.Object == "28" { - return true - } - } - return false - }, - }, - { - user: "我的工作是程序员", - want: 1, - check: func(triples []memory.Triple) bool { - for _, tr := range triples { - if tr.Subject == "用户" && tr.Relation == "职业" && tr.Object == "程序员" { + if tr.Subject == "我" && tr.Relation == "读书" && tr.Object == "杭州" { return true } } @@ -178,80 +140,20 @@ func TestExtractKeyTriples(t *testing.T) { }, { user: "今天天气真好", - want: 0, // 没有匹配任何规则 check: func(triples []memory.Triple) bool { - return true // any result is fine + return true // NLP 提取器可能不提取形容词谓语句,0 个也没关系 }, }, } for _, tt := range tests { triples := extractKeyTriples(tt.user, tt.assistant) - if len(triples) != tt.want { - t.Errorf("extractKeyTriples(%q) = %d triples, want %d", tt.user, len(triples), tt.want) - } if tt.check != nil && !tt.check(triples) { t.Errorf("extractKeyTriples(%q) = %v, check failed", tt.user, triples) } } } -func TestExtractName(t *testing.T) { - tests := []struct{ input, want string }{ - {"我叫张三", "张三"}, - {"我的名字是李四", "李四"}, - {"今天天气好", ""}, - } - for _, tt := range tests { - got := extractName(tt.input) - if got != tt.want { - t.Errorf("extractName(%q) = %q, want %q", tt.input, got, tt.want) - } - } -} - -func TestExtractLocation(t *testing.T) { - tests := []struct{ input, want string }{ - {"我住在北京", "北京"}, - {"我家在上海", "上海"}, - {"hello", ""}, - } - for _, tt := range tests { - got := extractLocation(tt.input) - if got != tt.want { - t.Errorf("extractLocation(%q) = %q, want %q", tt.input, got, tt.want) - } - } -} - -func TestExtractLike(t *testing.T) { - tests := []struct{ input, want string }{ - {"我喜欢打篮球", "打篮球"}, - {"我最喜欢跑步", "跑步"}, - {"nothing", ""}, - } - for _, tt := range tests { - got := extractLike(tt.input) - if got != tt.want { - t.Errorf("extractLike(%q) = %q, want %q", tt.input, got, tt.want) - } - } -} - -func TestExtractAge(t *testing.T) { - tests := []struct{ input, want string }{ - {"我28岁", "28"}, - {"我的年龄是30", "30"}, - {"hello", ""}, - } - for _, tt := range tests { - got := extractAge(tt.input) - if got != tt.want { - t.Errorf("extractAge(%q) = %q, want %q", tt.input, got, tt.want) - } - } -} - func TestDistillerGetRecentRecords(t *testing.T) { d := NewDistiller(nil, t.TempDir(), DistillerConfig{}) d.Append("s1", "user", "a") diff --git a/internal/memory/static_embedder.go b/internal/memory/static_embedder.go index 860588e..ccfc18a 100644 --- a/internal/memory/static_embedder.go +++ b/internal/memory/static_embedder.go @@ -271,7 +271,7 @@ func (e *StaticEmbedder) tokenize(text string) []string { if e.jieba == nil { return nil } - words := e.jieba.Cut(text, true) + words := e.jieba.Cut(text, false) var result []string seen := make(map[string]bool) for _, w := range words { diff --git a/internal/memory/vector/store.go b/internal/memory/vector/store.go index 9011dda..46e10b7 100644 --- a/internal/memory/vector/store.go +++ b/internal/memory/vector/store.go @@ -121,21 +121,31 @@ func (s *Store) All() []DocVector { return out } -// TFIDFVectorizer 使用字符 bigram + TF-IDF -type TFIDFVectorizer struct { - mu sync.RWMutex - docFreq map[string]float64 // feature → 文档频率 - totalDocs int - maxNGram int +// Tokenizer 将文本拆分为词级 token +type Tokenizer func(string) []string + +// NGramTokenizer 创建字符 n-gram tokenizer(降级方案) +func NGramTokenizer(maxN int) Tokenizer { + return func(text string) []string { + return extractNGrams(text, maxN) + } } -func NewTFIDFVectorizer(maxNGram int) *TFIDFVectorizer { - if maxNGram <= 0 { - maxNGram = 2 +// TFIDFVectorizer 使用 tokenizer + TF-IDF +type TFIDFVectorizer struct { + mu sync.RWMutex + tokenizer Tokenizer + docFreq map[string]float64 // feature → 文档频率 + totalDocs int +} + +func NewTFIDFVectorizer(tokenizer Tokenizer) *TFIDFVectorizer { + if tokenizer == nil { + tokenizer = NGramTokenizer(2) } return &TFIDFVectorizer{ - docFreq: make(map[string]float64), - maxNGram: maxNGram, + tokenizer: tokenizer, + docFreq: make(map[string]float64), } } @@ -148,7 +158,7 @@ func (v *TFIDFVectorizer) Train(docs []string) { seen := make(map[string]map[string]bool) for _, doc := range docs { - features := extractNGrams(doc, v.maxNGram) + features := v.tokenizer(doc) key := doc if seen[key] == nil { seen[key] = make(map[string]bool) @@ -166,7 +176,7 @@ func (v *TFIDFVectorizer) Vectorize(text string) Vector { v.mu.RLock() defer v.mu.RUnlock() - features := extractNGrams(text, v.maxNGram) + features := v.tokenizer(text) tf := make(map[string]float64) for _, f := range features { tf[f]++ diff --git a/internal/memory/vector/store_test.go b/internal/memory/vector/store_test.go index 2caf507..605cc85 100644 --- a/internal/memory/vector/store_test.go +++ b/internal/memory/vector/store_test.go @@ -65,7 +65,7 @@ func TestCosineSimilarity(t *testing.T) { } func TestTFIDFVectorizer(t *testing.T) { - v := NewTFIDFVectorizer(2) + v := NewTFIDFVectorizer(NGramTokenizer(2)) docs := []string{"今天天气很好", "今天心情不错", "明天要下雨"} v.Train(docs) @@ -87,7 +87,7 @@ func TestTFIDFVectorizer(t *testing.T) { } func TestTFIDFVectorizerEmpty(t *testing.T) { - v := NewTFIDFVectorizer(2) + v := NewTFIDFVectorizer(NGramTokenizer(2)) v.Train(nil) vec := v.Vectorize("test") if len(vec) == 0 { @@ -127,7 +127,7 @@ func TestInvertedIndex(t *testing.T) { func TestStoreInsertAndSearch(t *testing.T) { s := NewStore() - v := NewTFIDFVectorizer(2) + v := NewTFIDFVectorizer(NGramTokenizer(2)) v.Train([]string{"hello world", "goodbye world"}) s.Insert("1", "hello world", v.Vectorize("hello world"), nil) @@ -148,7 +148,7 @@ func TestStoreInsertAndSearch(t *testing.T) { func TestStoreRemove(t *testing.T) { s := NewStore() - v := NewTFIDFVectorizer(1) + v := NewTFIDFVectorizer(NGramTokenizer(1)) v.Train([]string{"a"}) s.Insert("1", "a", v.Vectorize("a"), nil) @@ -175,7 +175,7 @@ func TestStoreEmpty(t *testing.T) { func TestStoreAll(t *testing.T) { s := NewStore() - v := NewTFIDFVectorizer(1) + v := NewTFIDFVectorizer(NGramTokenizer(1)) v.Train([]string{"a", "b"}) s.Insert("1", "a", v.Vectorize("a"), map[string]string{"k": "v"}) diff --git a/internal/nlp/bridge.go b/internal/nlp/bridge.go new file mode 100644 index 0000000..5c1cd16 --- /dev/null +++ b/internal/nlp/bridge.go @@ -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, + } +} diff --git a/internal/nlp/download.go b/internal/nlp/download.go new file mode 100644 index 0000000..14a2bb2 --- /dev/null +++ b/internal/nlp/download.go @@ -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 +} diff --git a/internal/nlp/extractor.go b/internal/nlp/extractor.go new file mode 100644 index 0000000..37c7b42 --- /dev/null +++ b/internal/nlp/extractor.go @@ -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" +} diff --git a/internal/nlp/extractor_test.go b/internal/nlp/extractor_test.go new file mode 100644 index 0000000..d4f84d7 --- /dev/null +++ b/internal/nlp/extractor_test.go @@ -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 +} diff --git a/internal/nlp/fallback.go b/internal/nlp/fallback.go new file mode 100644 index 0000000..e619b4c --- /dev/null +++ b/internal/nlp/fallback.go @@ -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 +} diff --git a/internal/nlp/model.go b/internal/nlp/model.go new file mode 100644 index 0000000..8e56ee3 --- /dev/null +++ b/internal/nlp/model.go @@ -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 +} diff --git a/internal/nlp/onnx_stub.go b/internal/nlp/onnx_stub.go new file mode 100644 index 0000000..2606e8b --- /dev/null +++ b/internal/nlp/onnx_stub.go @@ -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") +} diff --git a/internal/nlp/parser.go b/internal/nlp/parser.go new file mode 100644 index 0000000..ef3a44a --- /dev/null +++ b/internal/nlp/parser.go @@ -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 +} diff --git a/internal/plugins/webui/handler_test.go b/internal/plugins/webui/handler_test.go index f2353df..8c56d3b 100644 --- a/internal/plugins/webui/handler_test.go +++ b/internal/plugins/webui/handler_test.go @@ -631,6 +631,7 @@ func TestSettingsWithPluginRegistry(t *testing.T) { type echoProvider struct{ name string } func (p *echoProvider) Name() string { return p.name } +func (p *echoProvider) MaxContextTokens() int { return 8192 } func (p *echoProvider) Chat(ctx context.Context, req *agentAPI.CompletionRequest) (*agentAPI.CompletionResponse, error) { content := "echo: " + req.Messages[len(req.Messages)-1].Content return &agentAPI.CompletionResponse{Content: content, FinishReason: "stop"}, nil