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:
377
docs/zh/plan.md
Normal file
377
docs/zh/plan.md
Normal file
@ -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 |
|
||||||
@ -148,6 +148,47 @@ type Provider interface {
|
|||||||
Name() string
|
Name() string
|
||||||
Chat(ctx context.Context, req *CompletionRequest) (*CompletionResponse, error)
|
Chat(ctx context.Context, req *CompletionRequest) (*CompletionResponse, error)
|
||||||
ChatStream(ctx context.Context, req *CompletionRequest) (<-chan StreamChunk, 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 {
|
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) Name() string { return p.name }
|
||||||
|
|
||||||
func (p *LuaAdaptedProvider) Chat(ctx context.Context, req *CompletionRequest) (*CompletionResponse, error) {
|
func (p *LuaAdaptedProvider) Chat(ctx context.Context, req *CompletionRequest) (*CompletionResponse, error) {
|
||||||
|
|||||||
@ -105,6 +105,8 @@ type Agent struct {
|
|||||||
noMergeMarkers map[string]int
|
noMergeMarkers map[string]int
|
||||||
noMergeMu sync.Mutex
|
noMergeMu sync.Mutex
|
||||||
|
|
||||||
|
// 词嵌入模型,用于实体语义相似度计算
|
||||||
|
embedder *memory.StaticEmbedder
|
||||||
}
|
}
|
||||||
|
|
||||||
type AgentConfig struct {
|
type AgentConfig struct {
|
||||||
@ -187,6 +189,7 @@ func New(cfg AgentConfig) *Agent {
|
|||||||
pluginHealth: newPluginHealthTracker(),
|
pluginHealth: newPluginHealthTracker(),
|
||||||
thinkingEnabled: cfg.ThinkingEnabled,
|
thinkingEnabled: cfg.ThinkingEnabled,
|
||||||
inputCfg: cfg.InputProcessing,
|
inputCfg: cfg.InputProcessing,
|
||||||
|
embedder: embedder,
|
||||||
noMergeMarkers: make(map[string]int),
|
noMergeMarkers: make(map[string]int),
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|||||||
@ -35,14 +35,11 @@ func TestDocToTriples(t *testing.T) {
|
|||||||
triples := docToTriples(doc)
|
triples := docToTriples(doc)
|
||||||
|
|
||||||
foundSummary := false
|
foundSummary := false
|
||||||
foundRel := false
|
|
||||||
foundSource := false
|
foundSource := false
|
||||||
for _, tr := range triples {
|
for _, tr := range triples {
|
||||||
switch {
|
switch {
|
||||||
case tr.Subject == "文档" && tr.Relation == "主题":
|
case tr.Subject == "文档" && tr.Relation == "主题":
|
||||||
foundSummary = true
|
foundSummary = true
|
||||||
case tr.Relation == "关联":
|
|
||||||
foundRel = true
|
|
||||||
case tr.Subject == "文档" && tr.Relation == "来源":
|
case tr.Subject == "文档" && tr.Relation == "来源":
|
||||||
foundSource = true
|
foundSource = true
|
||||||
}
|
}
|
||||||
@ -54,9 +51,6 @@ func TestDocToTriples(t *testing.T) {
|
|||||||
if !foundSource {
|
if !foundSource {
|
||||||
t.Error("missing '来源' triple")
|
t.Error("missing '来源' triple")
|
||||||
}
|
}
|
||||||
if needJieba() && !foundRel {
|
|
||||||
t.Error("missing '关联' triple with jieba available")
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestDocToTriplesNil(t *testing.T) {
|
func TestDocToTriplesNil(t *testing.T) {
|
||||||
@ -88,7 +82,6 @@ func TestDocToTriplesTypes(t *testing.T) {
|
|||||||
|
|
||||||
triples := docToTriples(doc)
|
triples := docToTriples(doc)
|
||||||
|
|
||||||
// 主题 and 来源 triples have Subject=文档
|
|
||||||
for _, tr := range triples {
|
for _, tr := range triples {
|
||||||
if tr.Subject == "文档" {
|
if tr.Subject == "文档" {
|
||||||
if tr.SubjectType != "Concept" {
|
if tr.SubjectType != "Concept" {
|
||||||
@ -97,14 +90,6 @@ func TestDocToTriplesTypes(t *testing.T) {
|
|||||||
if tr.Confidence != 1.0 {
|
if tr.Confidence != 1.0 {
|
||||||
t.Errorf("文档 triple confidence should be 1.0, got %f", tr.Confidence)
|
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
|
// all should have SubjectType/ObjectType set
|
||||||
if tr.SubjectType == "" || tr.ObjectType == "" {
|
if tr.SubjectType == "" || tr.ObjectType == "" {
|
||||||
|
|||||||
@ -40,11 +40,8 @@ func TestDocToTriplesConversation(t *testing.T) {
|
|||||||
}
|
}
|
||||||
triples := docToTriples(doc)
|
triples := docToTriples(doc)
|
||||||
|
|
||||||
minLen := 2
|
if len(triples) < 2 {
|
||||||
hasJieba := needJieba()
|
t.Errorf("expected at least 2 triples (主题+来源), got %d", len(triples))
|
||||||
|
|
||||||
if hasJieba && len(triples) <= minLen {
|
|
||||||
t.Errorf("expected more than %d triples with jieba, got %d", minLen, len(triples))
|
|
||||||
}
|
}
|
||||||
|
|
||||||
for i, tr := range 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)
|
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) {
|
func TestDocToTriplesMultiLine(t *testing.T) {
|
||||||
|
|||||||
@ -4,12 +4,13 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"log"
|
"log"
|
||||||
"runtime/debug"
|
"runtime/debug"
|
||||||
"strings"
|
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
agentIO "gitcode.com/JianFeeeee/HomeAgent/internal/agent/io"
|
agentIO "gitcode.com/JianFeeeee/HomeAgent/internal/agent/io"
|
||||||
"gitcode.com/JianFeeeee/HomeAgent/internal/memory"
|
"gitcode.com/JianFeeeee/HomeAgent/internal/memory"
|
||||||
"gitcode.com/JianFeeeee/HomeAgent/internal/memory/document"
|
"gitcode.com/JianFeeeee/HomeAgent/internal/memory/document"
|
||||||
|
"gitcode.com/JianFeeeee/HomeAgent/internal/memory/vector"
|
||||||
|
"gitcode.com/JianFeeeee/HomeAgent/internal/nlp"
|
||||||
)
|
)
|
||||||
|
|
||||||
type ConsolidationTask struct {
|
type ConsolidationTask struct {
|
||||||
@ -107,15 +108,18 @@ func (a *Agent) reorgGraph() {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
llmCandidates := 0
|
||||||
maxCandidates := 5
|
maxCandidates := 5
|
||||||
candidates := 0
|
|
||||||
for i := 0; i < len(result.Entities) && candidates < maxCandidates; i++ {
|
for i := 0; i < len(result.Entities) && llmCandidates < maxCandidates; i++ {
|
||||||
for j := i + 1; j < len(result.Entities) && candidates < maxCandidates; j++ {
|
for j := i + 1; j < len(result.Entities) && llmCandidates < maxCandidates; j++ {
|
||||||
ea, eb := result.Entities[i].Name, result.Entities[j].Name
|
ea, eb := result.Entities[i].Name, result.Entities[j].Name
|
||||||
if ea > eb {
|
if ea > eb {
|
||||||
ea, eb = eb, ea
|
ea, eb = eb, ea
|
||||||
}
|
}
|
||||||
key := ea + "||" + eb
|
key := ea + "||" + eb
|
||||||
|
|
||||||
|
// 跳过已标记"不合并"的实体对
|
||||||
a.noMergeMu.Lock()
|
a.noMergeMu.Lock()
|
||||||
rounds, ok := a.noMergeMarkers[key]
|
rounds, ok := a.noMergeMarkers[key]
|
||||||
if ok {
|
if ok {
|
||||||
@ -130,9 +134,16 @@ func (a *Agent) reorgGraph() {
|
|||||||
if ok {
|
if ok {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// 复合相似度:字符二元组 + 语义向量(仅增强检测,不做自动合并)
|
||||||
sim := entitySimilarity(result.Entities[i].Name, result.Entities[j].Name)
|
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 {
|
if sim > 0.75 {
|
||||||
candidates++
|
llmCandidates++
|
||||||
a.enqueueConsolidationTask(ConsolidationTask{
|
a.enqueueConsolidationTask(ConsolidationTask{
|
||||||
Type: "entity_merge",
|
Type: "entity_merge",
|
||||||
Reason: fmt.Sprintf(
|
Reason: fmt.Sprintf(
|
||||||
@ -155,83 +166,24 @@ func (a *Agent) reorgGraph() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if candidates > 0 {
|
if llmCandidates > 0 {
|
||||||
log.Printf("[agent] graph reorg: %d merge candidates sent for LLM decision", candidates)
|
log.Printf("[agent] graph reorg: %d merge candidates sent for LLM decision", llmCandidates)
|
||||||
} else {
|
} else {
|
||||||
log.Printf("[agent] graph reorg: no similar entities found")
|
log.Printf("[agent] graph reorg: no similar entities found")
|
||||||
}
|
}
|
||||||
|
|
||||||
a.evaluateGraphQuality()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (a *Agent) evaluateGraphQuality() {
|
// entitySemanticSimilarity 使用词嵌入向量余弦相似度计算实体名语义相似度
|
||||||
if a.memory == nil {
|
func entitySemanticSimilarity(a, b string, embedder *memory.StaticEmbedder) float64 {
|
||||||
return
|
if a == "" || b == "" || embedder == nil || !embedder.Loaded() {
|
||||||
|
return 0
|
||||||
}
|
}
|
||||||
|
va := embedder.Vectorize(a)
|
||||||
pending, err := a.memory.RecallPending(10)
|
vb := embedder.Vectorize(b)
|
||||||
if err != nil {
|
if len(va) == 0 || len(vb) == 0 {
|
||||||
log.Printf("[agent] recall pending relations error: %v", err)
|
return 0
|
||||||
return
|
|
||||||
}
|
}
|
||||||
if len(pending) == 0 {
|
return vector.CosineSimilarity(va, vb)
|
||||||
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))
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func entitySimilarity(a, b string) float64 {
|
func entitySimilarity(a, b string) float64 {
|
||||||
@ -286,6 +238,7 @@ func docToTriples(doc *document.Doc) []memory.Triple {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// 文档元数据
|
||||||
triples = append(triples, memory.Triple{
|
triples = append(triples, memory.Triple{
|
||||||
Subject: "文档",
|
Subject: "文档",
|
||||||
SubjectType: "Concept",
|
SubjectType: "Concept",
|
||||||
@ -295,22 +248,15 @@ func docToTriples(doc *document.Doc) []memory.Triple {
|
|||||||
Confidence: 1.0,
|
Confidence: 1.0,
|
||||||
})
|
})
|
||||||
|
|
||||||
lines := strings.Split(doc.Content, "\n")
|
// NLP 通用提取
|
||||||
for _, line := range lines {
|
e := nlp.NewExtractor(nil)
|
||||||
line = strings.TrimSpace(line)
|
result := e.Extract(doc.Content)
|
||||||
if line == "" {
|
if result != nil {
|
||||||
continue
|
for _, nt := range result.Triples {
|
||||||
}
|
mt := nlp.ToMemoryTriple(nt)
|
||||||
terms := memory.CutExact(line)
|
if mt.Subject != "" && mt.Relation != "" && mt.Object != "" {
|
||||||
for i := 0; i < len(terms)-1; i++ {
|
triples = append(triples, mt)
|
||||||
triples = append(triples, memory.Triple{
|
}
|
||||||
Subject: terms[i],
|
|
||||||
SubjectType: "Concept",
|
|
||||||
Relation: "关联",
|
|
||||||
Object: terms[i+1],
|
|
||||||
ObjectType: "Concept",
|
|
||||||
Confidence: 0.8,
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@ -354,13 +300,5 @@ func (a *Agent) processConsolidation(evt *agentIO.InputEvent, input string) {
|
|||||||
return
|
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)
|
log.Printf("[agent] consolidation done (%dms, tools=%v)", time.Since(start).Milliseconds(), toolsUsed)
|
||||||
}
|
}
|
||||||
|
|||||||
@ -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")
|
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)
|
sysPrompt := a.buildSystemPrompt(memContext, input)
|
||||||
tools := a.buildToolDefs()
|
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 blocks, ok := stageCtx.Extra["media_blocks"].([]agentAPI.ContentBlock); ok && len(blocks) > 0 {
|
||||||
if len(msgs) > 0 {
|
if len(msgs) > 0 {
|
||||||
msgs[len(msgs)-1].Blocks = blocks
|
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(),
|
len(tools), a.context.Len(),
|
||||||
a.personality != nil && a.personality.Content != "",
|
a.personality != nil && a.personality.Content != "",
|
||||||
a.docStoreSize())
|
a.docStoreSize())
|
||||||
@ -295,7 +298,7 @@ func (a *Agent) docStoreSize() int {
|
|||||||
return 0
|
return 0
|
||||||
}
|
}
|
||||||
|
|
||||||
func (a *Agent) formatMergedTimeline() string {
|
func (a *Agent) formatMergedTimeline(maxTokens int) string {
|
||||||
a.context.mu.Lock()
|
a.context.mu.Lock()
|
||||||
events := make([]*ContextEvent, len(a.context.events))
|
events := make([]*ContextEvent, len(a.context.events))
|
||||||
copy(events, a.context.events)
|
copy(events, a.context.events)
|
||||||
@ -305,9 +308,35 @@ func (a *Agent) formatMergedTimeline() string {
|
|||||||
return ""
|
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
|
var sb strings.Builder
|
||||||
sb.WriteString("【对话时序】\n")
|
sb.WriteString("【对话时序】\n")
|
||||||
for _, e := range events {
|
for _, e := range events[start:] {
|
||||||
sb.WriteString(fmt.Sprintf("[%s] %s: %s",
|
sb.WriteString(fmt.Sprintf("[%s] %s: %s",
|
||||||
e.Timestamp.Format("15:04:05"), e.Source, e.Input))
|
e.Timestamp.Format("15:04:05"), e.Source, e.Input))
|
||||||
if len(e.ToolsUsed) > 0 {
|
if len(e.ToolsUsed) > 0 {
|
||||||
@ -321,11 +350,11 @@ func (a *Agent) formatMergedTimeline() string {
|
|||||||
return sb.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}}
|
msgs := []agentAPI.Message{{Role: "system", Content: sysPrompt}}
|
||||||
|
|
||||||
if ctxStr := a.formatMergedTimeline(); ctxStr != "" {
|
if ctxTok := a.formatMergedTimeline(ctxTokens); ctxTok != "" {
|
||||||
msgs = append(msgs, agentAPI.Message{Role: "system", Content: ctxStr})
|
msgs = append(msgs, agentAPI.Message{Role: "system", Content: ctxTok})
|
||||||
}
|
}
|
||||||
|
|
||||||
msgs = append(msgs, agentAPI.Message{Role: "user", Content: input})
|
msgs = append(msgs, agentAPI.Message{Role: "user", Content: input})
|
||||||
|
|||||||
86
internal/agent/core/tokenbudget.go
Normal file
86
internal/agent/core/tokenbudget.go
Normal file
@ -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])
|
||||||
|
}
|
||||||
@ -7,12 +7,16 @@ import (
|
|||||||
agentIO "gitcode.com/JianFeeeee/HomeAgent/internal/agent/io"
|
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 {
|
if a.indexer == nil {
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
injected := a.indexer.BuildContext(input)
|
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 {
|
func (a *Agent) buildSystemPrompt(memContext string, userInput string) string {
|
||||||
|
|||||||
@ -92,7 +92,7 @@ func NewStore(root string) *Store {
|
|||||||
root: root,
|
root: root,
|
||||||
indexPath: filepath.Join(root, ".index.json"),
|
indexPath: filepath.Join(root, ".index.json"),
|
||||||
vec: vector.NewStore(),
|
vec: vector.NewStore(),
|
||||||
veczer: vector.NewTFIDFVectorizer(3),
|
veczer: vector.NewTFIDFVectorizer(memory.TokenizeWords),
|
||||||
items: make(map[string]*Knowledge),
|
items: make(map[string]*Knowledge),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@ -4,6 +4,7 @@ import (
|
|||||||
"log"
|
"log"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
|
|
||||||
"github.com/yanyiwu/gojieba"
|
"github.com/yanyiwu/gojieba"
|
||||||
@ -92,13 +93,34 @@ var stopWords = map[string]bool{
|
|||||||
"when": true, "who": true, "whom": true,
|
"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 {
|
func ExtractKeywords(text string) []string {
|
||||||
text = CleanText(text)
|
text = CleanText(text)
|
||||||
x := GetJieba()
|
x := GetJieba()
|
||||||
if x == nil {
|
if x == nil {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
words := x.Cut(text, true)
|
words := x.Cut(text, false)
|
||||||
var keywords []string
|
var keywords []string
|
||||||
seen := make(map[string]bool)
|
seen := make(map[string]bool)
|
||||||
for _, w := range words {
|
for _, w := range words {
|
||||||
|
|||||||
@ -69,7 +69,7 @@ func NewStore(dir string) *Store {
|
|||||||
return &Store{
|
return &Store{
|
||||||
dir: dir,
|
dir: dir,
|
||||||
vec: vector.NewStore(),
|
vec: vector.NewStore(),
|
||||||
veczer: vector.NewTFIDFVectorizer(2),
|
veczer: vector.NewTFIDFVectorizer(memory.TokenizeWords),
|
||||||
docs: make(map[string]*Doc),
|
docs: make(map[string]*Doc),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@ -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
|
|
||||||
}
|
|
||||||
@ -31,9 +31,6 @@ type Relation struct {
|
|||||||
TurnID int `json:"turn_id"`
|
TurnID int `json:"turn_id"`
|
||||||
CreatedAt time.Time `json:"created_at"`
|
CreatedAt time.Time `json:"created_at"`
|
||||||
DateBucket string `json:"date_bucket"`
|
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 {
|
type Triple struct {
|
||||||
@ -96,9 +93,6 @@ func (g *GraphDB) initSchema() error {
|
|||||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||||
updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||||
date_bucket TEXT,
|
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 (source_id) REFERENCES entities(id),
|
||||||
FOREIGN KEY (target_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 tx.Commit()
|
||||||
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
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (g *GraphDB) Commit(triples []Triple, sessionID string, turnID int) (int, int, error) {
|
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(
|
relRows, err := g.db.Query(
|
||||||
`SELECT r.id, r.source_id, r.target_id, e1.name, e2.name,
|
`SELECT r.id, r.source_id, r.target_id, e1.name, e2.name,
|
||||||
r.relation_type, r.confidence, r.status, r.session_id,
|
r.relation_type, r.confidence, r.status, r.session_id,
|
||||||
r.turn_id, r.created_at, COALESCE(r.date_bucket, ''),
|
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
|
FROM relations r
|
||||||
JOIN entities e1 ON r.source_id = e1.id
|
JOIN entities e1 ON r.source_id = e1.id
|
||||||
JOIN entities e2 ON r.target_id = e2.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,
|
if err := relRows.Scan(&rel.ID, &rel.SourceID, &rel.TargetID,
|
||||||
&rel.SourceName, &rel.TargetName, &rel.RelationType,
|
&rel.SourceName, &rel.TargetName, &rel.RelationType,
|
||||||
&rel.Confidence, &rel.Status, &rel.SessionID,
|
&rel.Confidence, &rel.Status, &rel.SessionID,
|
||||||
&rel.TurnID, &rel.CreatedAt, &rel.DateBucket,
|
&rel.TurnID, &rel.CreatedAt, &rel.DateBucket); err != nil {
|
||||||
&rel.EvalStatus, &rel.EvalRound, &rel.EvalAt); err != nil {
|
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
result.Relations = append(result.Relations, rel)
|
result.Relations = append(result.Relations, rel)
|
||||||
@ -361,8 +340,7 @@ func (g *GraphDB) Recall(keywords []string, seedEntities []string, depth int, se
|
|||||||
query := fmt.Sprintf(
|
query := fmt.Sprintf(
|
||||||
`SELECT r.id, r.source_id, r.target_id, e1.name, e2.name,
|
`SELECT r.id, r.source_id, r.target_id, e1.name, e2.name,
|
||||||
r.relation_type, r.confidence, r.status, r.session_id,
|
r.relation_type, r.confidence, r.status, r.session_id,
|
||||||
r.turn_id, r.created_at, COALESCE(r.date_bucket, ''),
|
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
|
FROM relations r
|
||||||
JOIN entities e1 ON r.source_id = e1.id
|
JOIN entities e1 ON r.source_id = e1.id
|
||||||
JOIN entities e2 ON r.target_id = e2.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,
|
if err := relRows.Scan(&rel.ID, &rel.SourceID, &rel.TargetID,
|
||||||
&rel.SourceName, &rel.TargetName, &rel.RelationType,
|
&rel.SourceName, &rel.TargetName, &rel.RelationType,
|
||||||
&rel.Confidence, &rel.Status, &rel.SessionID,
|
&rel.Confidence, &rel.Status, &rel.SessionID,
|
||||||
&rel.TurnID, &rel.CreatedAt, &rel.DateBucket,
|
&rel.TurnID, &rel.CreatedAt, &rel.DateBucket); err != nil {
|
||||||
&rel.EvalStatus, &rel.EvalRound, &rel.EvalAt); err != nil {
|
|
||||||
relRows.Close()
|
relRows.Close()
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@ -770,89 +747,6 @@ func (g *GraphDB) Archive(days int) (int, error) {
|
|||||||
return int(n), nil
|
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 {
|
func (g *GraphDB) Close() error {
|
||||||
return g.db.Close()
|
return g.db.Close()
|
||||||
}
|
}
|
||||||
|
|||||||
@ -118,11 +118,11 @@ func TestRecallWithDepth(t *testing.T) {
|
|||||||
defer g.Close()
|
defer g.Close()
|
||||||
|
|
||||||
g.Commit([]Triple{
|
g.Commit([]Triple{
|
||||||
{Subject: "甲", Relation: "认识", Object: "乙"},
|
{Subject: "小明", Relation: "认识", Object: "小红"},
|
||||||
{Subject: "乙", Relation: "认识", Object: "丙"},
|
{Subject: "小红", Relation: "认识", Object: "小刚"},
|
||||||
}, "session3", 0)
|
}, "session3", 0)
|
||||||
|
|
||||||
result, err := g.Recall(nil, []string{"甲"}, 2, "")
|
result, err := g.Recall(nil, []string{"小明"}, 2, "")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
|
|||||||
@ -22,7 +22,7 @@ func NewIndexer(db *GraphDB) *Indexer {
|
|||||||
return &Indexer{
|
return &Indexer{
|
||||||
db: db,
|
db: db,
|
||||||
vec: vector.NewStore(),
|
vec: vector.NewStore(),
|
||||||
veczer: vector.NewTFIDFVectorizer(2),
|
veczer: vector.NewTFIDFVectorizer(TokenizeWords),
|
||||||
recalled: make(map[string]bool),
|
recalled: make(map[string]bool),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@ -14,6 +14,7 @@ import (
|
|||||||
"time"
|
"time"
|
||||||
|
|
||||||
"gitcode.com/JianFeeeee/HomeAgent/internal/memory"
|
"gitcode.com/JianFeeeee/HomeAgent/internal/memory"
|
||||||
|
"gitcode.com/JianFeeeee/HomeAgent/internal/nlp"
|
||||||
)
|
)
|
||||||
|
|
||||||
type RawRecord struct {
|
type RawRecord struct {
|
||||||
@ -265,182 +266,24 @@ func (d *Distiller) cleanupRawFiles() {
|
|||||||
func extractKeyTriples(userContent, assistantContent string) []memory.Triple {
|
func extractKeyTriples(userContent, assistantContent string) []memory.Triple {
|
||||||
var triples []memory.Triple
|
var triples []memory.Triple
|
||||||
|
|
||||||
// 提取对话中的关键信息,而不是直接 dump 原文
|
e := nlp.NewExtractor(nil)
|
||||||
// 规则1: "我的名字是X" / "我叫X" → (用户, 姓名, X)
|
text := userContent
|
||||||
if name := extractName(userContent); name != "" {
|
if assistantContent != "" {
|
||||||
triples = append(triples, memory.Triple{Subject: "用户", Relation: "姓名", Object: name})
|
text += assistantContent
|
||||||
}
|
}
|
||||||
// 规则2: "我住在X" / "我家在X" → (用户, 居住地, X)
|
result := e.Extract(text)
|
||||||
if loc := extractLocation(userContent); loc != "" {
|
if result != nil {
|
||||||
triples = append(triples, memory.Triple{Subject: "用户", Relation: "居住地", Object: loc})
|
for _, nt := range result.Triples {
|
||||||
}
|
mt := nlp.ToMemoryTriple(nt)
|
||||||
// 规则3: "我喜欢X" / "我爱X" → (用户, 喜好, X)
|
if mt.Subject != "" && mt.Relation != "" && mt.Object != "" {
|
||||||
if like := extractLike(userContent); like != "" {
|
triples = append(triples, mt)
|
||||||
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})
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return triples
|
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 {
|
func truncate(s string, max int) string {
|
||||||
if len(s) > max {
|
if len(s) > max {
|
||||||
return s[:max] + "..."
|
return s[:max] + "..."
|
||||||
|
|||||||
@ -113,27 +113,13 @@ func TestExtractKeyTriples(t *testing.T) {
|
|||||||
tests := []struct {
|
tests := []struct {
|
||||||
user string
|
user string
|
||||||
assistant string
|
assistant string
|
||||||
want int // expected number of triples
|
|
||||||
check func([]memory.Triple) bool
|
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: "我住在北京",
|
user: "我住在北京",
|
||||||
want: 1,
|
|
||||||
check: func(triples []memory.Triple) bool {
|
check: func(triples []memory.Triple) bool {
|
||||||
for _, tr := range triples {
|
for _, tr := range triples {
|
||||||
if tr.Subject == "用户" && tr.Relation == "居住地" && tr.Object == "北京" {
|
if tr.Subject == "我" && tr.Relation == "住" && tr.Object == "北京" {
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@ -141,35 +127,11 @@ func TestExtractKeyTriples(t *testing.T) {
|
|||||||
},
|
},
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
user: "我喜欢打篮球",
|
user: "我在杭州读书",
|
||||||
want: 1,
|
assistant: "好的",
|
||||||
check: func(triples []memory.Triple) bool {
|
check: func(triples []memory.Triple) bool {
|
||||||
for _, tr := range triples {
|
for _, tr := range triples {
|
||||||
if tr.Subject == "用户" && tr.Relation == "喜好" && tr.Object == "打篮球" {
|
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 == "程序员" {
|
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@ -178,80 +140,20 @@ func TestExtractKeyTriples(t *testing.T) {
|
|||||||
},
|
},
|
||||||
{
|
{
|
||||||
user: "今天天气真好",
|
user: "今天天气真好",
|
||||||
want: 0, // 没有匹配任何规则
|
|
||||||
check: func(triples []memory.Triple) bool {
|
check: func(triples []memory.Triple) bool {
|
||||||
return true // any result is fine
|
return true // NLP 提取器可能不提取形容词谓语句,0 个也没关系
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, tt := range tests {
|
for _, tt := range tests {
|
||||||
triples := extractKeyTriples(tt.user, tt.assistant)
|
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) {
|
if tt.check != nil && !tt.check(triples) {
|
||||||
t.Errorf("extractKeyTriples(%q) = %v, check failed", tt.user, 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) {
|
func TestDistillerGetRecentRecords(t *testing.T) {
|
||||||
d := NewDistiller(nil, t.TempDir(), DistillerConfig{})
|
d := NewDistiller(nil, t.TempDir(), DistillerConfig{})
|
||||||
d.Append("s1", "user", "a")
|
d.Append("s1", "user", "a")
|
||||||
|
|||||||
@ -271,7 +271,7 @@ func (e *StaticEmbedder) tokenize(text string) []string {
|
|||||||
if e.jieba == nil {
|
if e.jieba == nil {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
words := e.jieba.Cut(text, true)
|
words := e.jieba.Cut(text, false)
|
||||||
var result []string
|
var result []string
|
||||||
seen := make(map[string]bool)
|
seen := make(map[string]bool)
|
||||||
for _, w := range words {
|
for _, w := range words {
|
||||||
|
|||||||
@ -121,21 +121,31 @@ func (s *Store) All() []DocVector {
|
|||||||
return out
|
return out
|
||||||
}
|
}
|
||||||
|
|
||||||
// TFIDFVectorizer 使用字符 bigram + TF-IDF
|
// Tokenizer 将文本拆分为词级 token
|
||||||
type TFIDFVectorizer struct {
|
type Tokenizer func(string) []string
|
||||||
mu sync.RWMutex
|
|
||||||
docFreq map[string]float64 // feature → 文档频率
|
// NGramTokenizer 创建字符 n-gram tokenizer(降级方案)
|
||||||
totalDocs int
|
func NGramTokenizer(maxN int) Tokenizer {
|
||||||
maxNGram int
|
return func(text string) []string {
|
||||||
|
return extractNGrams(text, maxN)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewTFIDFVectorizer(maxNGram int) *TFIDFVectorizer {
|
// TFIDFVectorizer 使用 tokenizer + TF-IDF
|
||||||
if maxNGram <= 0 {
|
type TFIDFVectorizer struct {
|
||||||
maxNGram = 2
|
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{
|
return &TFIDFVectorizer{
|
||||||
docFreq: make(map[string]float64),
|
tokenizer: tokenizer,
|
||||||
maxNGram: maxNGram,
|
docFreq: make(map[string]float64),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@ -148,7 +158,7 @@ func (v *TFIDFVectorizer) Train(docs []string) {
|
|||||||
|
|
||||||
seen := make(map[string]map[string]bool)
|
seen := make(map[string]map[string]bool)
|
||||||
for _, doc := range docs {
|
for _, doc := range docs {
|
||||||
features := extractNGrams(doc, v.maxNGram)
|
features := v.tokenizer(doc)
|
||||||
key := doc
|
key := doc
|
||||||
if seen[key] == nil {
|
if seen[key] == nil {
|
||||||
seen[key] = make(map[string]bool)
|
seen[key] = make(map[string]bool)
|
||||||
@ -166,7 +176,7 @@ func (v *TFIDFVectorizer) Vectorize(text string) Vector {
|
|||||||
v.mu.RLock()
|
v.mu.RLock()
|
||||||
defer v.mu.RUnlock()
|
defer v.mu.RUnlock()
|
||||||
|
|
||||||
features := extractNGrams(text, v.maxNGram)
|
features := v.tokenizer(text)
|
||||||
tf := make(map[string]float64)
|
tf := make(map[string]float64)
|
||||||
for _, f := range features {
|
for _, f := range features {
|
||||||
tf[f]++
|
tf[f]++
|
||||||
|
|||||||
@ -65,7 +65,7 @@ func TestCosineSimilarity(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestTFIDFVectorizer(t *testing.T) {
|
func TestTFIDFVectorizer(t *testing.T) {
|
||||||
v := NewTFIDFVectorizer(2)
|
v := NewTFIDFVectorizer(NGramTokenizer(2))
|
||||||
docs := []string{"今天天气很好", "今天心情不错", "明天要下雨"}
|
docs := []string{"今天天气很好", "今天心情不错", "明天要下雨"}
|
||||||
v.Train(docs)
|
v.Train(docs)
|
||||||
|
|
||||||
@ -87,7 +87,7 @@ func TestTFIDFVectorizer(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestTFIDFVectorizerEmpty(t *testing.T) {
|
func TestTFIDFVectorizerEmpty(t *testing.T) {
|
||||||
v := NewTFIDFVectorizer(2)
|
v := NewTFIDFVectorizer(NGramTokenizer(2))
|
||||||
v.Train(nil)
|
v.Train(nil)
|
||||||
vec := v.Vectorize("test")
|
vec := v.Vectorize("test")
|
||||||
if len(vec) == 0 {
|
if len(vec) == 0 {
|
||||||
@ -127,7 +127,7 @@ func TestInvertedIndex(t *testing.T) {
|
|||||||
|
|
||||||
func TestStoreInsertAndSearch(t *testing.T) {
|
func TestStoreInsertAndSearch(t *testing.T) {
|
||||||
s := NewStore()
|
s := NewStore()
|
||||||
v := NewTFIDFVectorizer(2)
|
v := NewTFIDFVectorizer(NGramTokenizer(2))
|
||||||
v.Train([]string{"hello world", "goodbye world"})
|
v.Train([]string{"hello world", "goodbye world"})
|
||||||
|
|
||||||
s.Insert("1", "hello world", v.Vectorize("hello world"), nil)
|
s.Insert("1", "hello world", v.Vectorize("hello world"), nil)
|
||||||
@ -148,7 +148,7 @@ func TestStoreInsertAndSearch(t *testing.T) {
|
|||||||
|
|
||||||
func TestStoreRemove(t *testing.T) {
|
func TestStoreRemove(t *testing.T) {
|
||||||
s := NewStore()
|
s := NewStore()
|
||||||
v := NewTFIDFVectorizer(1)
|
v := NewTFIDFVectorizer(NGramTokenizer(1))
|
||||||
v.Train([]string{"a"})
|
v.Train([]string{"a"})
|
||||||
|
|
||||||
s.Insert("1", "a", v.Vectorize("a"), nil)
|
s.Insert("1", "a", v.Vectorize("a"), nil)
|
||||||
@ -175,7 +175,7 @@ func TestStoreEmpty(t *testing.T) {
|
|||||||
|
|
||||||
func TestStoreAll(t *testing.T) {
|
func TestStoreAll(t *testing.T) {
|
||||||
s := NewStore()
|
s := NewStore()
|
||||||
v := NewTFIDFVectorizer(1)
|
v := NewTFIDFVectorizer(NGramTokenizer(1))
|
||||||
v.Train([]string{"a", "b"})
|
v.Train([]string{"a", "b"})
|
||||||
|
|
||||||
s.Insert("1", "a", v.Vectorize("a"), map[string]string{"k": "v"})
|
s.Insert("1", "a", v.Vectorize("a"), map[string]string{"k": "v"})
|
||||||
|
|||||||
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
|
||||||
|
}
|
||||||
@ -631,6 +631,7 @@ func TestSettingsWithPluginRegistry(t *testing.T) {
|
|||||||
type echoProvider struct{ name string }
|
type echoProvider struct{ name string }
|
||||||
|
|
||||||
func (p *echoProvider) Name() string { return p.name }
|
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) {
|
func (p *echoProvider) Chat(ctx context.Context, req *agentAPI.CompletionRequest) (*agentAPI.CompletionResponse, error) {
|
||||||
content := "echo: " + req.Messages[len(req.Messages)-1].Content
|
content := "echo: " + req.Messages[len(req.Messages)-1].Content
|
||||||
return &agentAPI.CompletionResponse{Content: content, FinishReason: "stop"}, nil
|
return &agentAPI.CompletionResponse{Content: content, FinishReason: "stop"}, nil
|
||||||
|
|||||||
Reference in New Issue
Block a user