mirror of
https://gitcode.com/JianFeeeee/HomeAgent.git
synced 2026-09-21 09:28:14 +00:00
- 根目录清理: 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
278 lines
6.1 KiB
Go
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"
|
|
}
|