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

278 lines
6.1 KiB
Go

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