Files
HomeAgent/internal/nlp/onnx.go
JianFeeeee f91b20ee16 v0.7.3: 重构 Provider 层 + 计算层隔离 + Cleaner/NoMemory 架构
- 删除 OpenAIProvider/OllamaProvider 死代码,LuaAdaptedProvider 独存
- DisableThinking 从 ExtraBody 移到 CompletionRequest 顶层字段
- ContextWindow 从 Provider 签名移到 BaseConfig/ModelContextWindow() 统管
- 确认 CleanText 仅做基本空白 trim,QQ 模板剥离归插件 Cleaner
- Cleaner/NoMemory 仅作用于向量计算和 jieba 分词层,原文不变
- context.ContextEvent/Doc.Content 始终保存原文
- 删除 nlp/download.go 死代码
- media.go: context.Background() -> a.ctx 级联
- clawhubadapter: HTTP 超时
- cut.go: 跨平台 mod cache 路径 (GOMODCACHE->GOPATH->HomeDir)
- bridge_e2e_test: 移除未用 runtime import
- lua 适配器: disable_thinking 传参
2026-07-28 11:42:29 +08:00

324 lines
7.2 KiB
Go

//go:build onnxruntime
package nlp
import (
"embed"
"encoding/json"
"fmt"
"os"
"path/filepath"
"sync"
"gitcode.com/JianFeeeee/HomeAgent/internal/memory"
ort "github.com/yalue/onnxruntime_go"
)
//go:embed models/*
var onnxModelFS embed.FS
const maxSeqLen = 128
type ONNXParser struct {
rt *ort.DynamicAdvancedSession
vocab map[string]int64
posVocab map[string]int64
close sync.Once
}
type ONNXConfig struct {
ModelPath string
DataDir string
}
func NewONNXParser(cfg ONNXConfig) (*ONNXParser, error) {
vocab, err := loadWordMap("models/vocab.json")
if err != nil {
return nil, fmt.Errorf("load vocab: %w", err)
}
posVocab, err := loadWordMap("models/pos_vocab.json")
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(libPath())
if err := ort.InitializeEnvironment(); err != nil {
return nil, fmt.Errorf("init onnx env: %w", err)
}
inputNames := []string{"input_ids"}
outputNames := []string{"pos_logits", "head_logits", "rel_logits"}
session, err := ort.NewDynamicAdvancedSession(modelPath, inputNames, outputNames, nil)
if err != nil {
ort.DestroyEnvironment()
return nil, fmt.Errorf("create session: %w", err)
}
return &ONNXParser{
rt: session,
vocab: vocab,
posVocab: posVocab,
}, nil
}
func (p *ONNXParser) Close() error {
p.close.Do(func() {
p.rt.Destroy()
ort.DestroyEnvironment()
})
return nil
}
func (p *ONNXParser) Parse(text string) (*ParseResult, error) {
if text == "" {
return &ParseResult{}, nil
}
x := memory.GetJieba()
if x == nil {
return nil, fmt.Errorf("jieba unavailable")
}
words := x.Cut(text, true)
if len(words) == 0 {
return &ParseResult{}, nil
}
inIDs := p.wordsToIDs(words, maxSeqLen)
n := len(inIDs) - 1 // exclude <bos>
if n <= 0 {
return &ParseResult{}, nil
}
if n > len(words) {
n = len(words)
}
padded := padTo(inIDs, maxSeqLen)
inTensor, err := ort.NewTensor(ort.NewShape(1, maxSeqLen), padded)
if err != nil {
return nil, fmt.Errorf("create input tensor: %w", err)
}
defer inTensor.Destroy()
outputs := make([]ort.Value, 3)
if err := p.rt.Run([]ort.Value{inTensor}, outputs); err != nil {
return nil, fmt.Errorf("onnx run: %w", err)
}
posOut, ok := outputs[0].(*ort.Tensor[float32])
if !ok {
return nil, fmt.Errorf("pos output not Tensor[float32]")
}
headOut, ok := outputs[1].(*ort.Tensor[float32])
if !ok {
return nil, fmt.Errorf("head output not Tensor[float32]")
}
relOut, ok := outputs[2].(*ort.Tensor[float32])
if !ok {
return nil, fmt.Errorf("rel output not Tensor[float32]")
}
defer posOut.Destroy()
defer headOut.Destroy()
defer relOut.Destroy()
posShape := posOut.GetShape() // [1, seq, posDim]
headShape := headOut.GetShape() // [1, seq, seq]
relShape := relOut.GetShape() // [1, seq, seq, relDim]
if len(posShape) < 3 || len(headShape) < 3 || len(relShape) < 4 {
return nil, fmt.Errorf("unexpected output ranks: pos=%d head=%d rel=%d",
len(posShape), len(headShape), len(relShape))
}
seqDim := int(headShape[1])
posDim := int(posShape[2])
relDim := int(relShape[3])
if n > seqDim {
n = seqDim
}
rawPOS := posOut.GetData()
rawHeads := headOut.GetData()
rawRels := relOut.GetData()
pos := decodePOS(rawPOS, n, posDim, p.posVocab)
heads := decodeHeads(rawHeads, n, seqDim)
rels := decodeRels(rawRels, n, seqDim, relDim, heads)
return &ParseResult{
Tokens: words[:n],
POS: pos,
Heads: heads,
DepRels: rels,
}, nil
}
func (p *ONNXParser) wordsToIDs(words []string, maxLen int) []int64 {
ids := make([]int64, 0, maxLen)
if bos, ok := p.vocab["<bos>"]; ok {
ids = append(ids, bos)
}
for _, w := range words {
if len(ids) >= maxLen {
break
}
if id, ok := p.vocab[w]; ok {
ids = append(ids, id)
} else if unk, ok := p.vocab["<unk>"]; ok {
ids = append(ids, unk)
}
}
return ids
}
func padTo(ids []int64, length int) []int64 {
for len(ids) < length {
ids = append(ids, 0)
}
return ids
}
func decodePOS(raw []float32, n, posDim int, posVocab map[string]int64) []string {
rev := make(map[int64]string)
for k, v := range posVocab {
rev[v] = k
}
pos := make([]string, n)
for i := 0; i < n; i++ {
bestIdx := 0
bestVal := float32(-1e9)
for j := 0; j < posDim; j++ {
if v := raw[i*posDim+j]; v > bestVal {
bestVal = v
bestIdx = j
}
}
if tag, ok := rev[int64(bestIdx)]; ok {
pos[i] = tag
} else {
pos[i] = "X"
}
}
return pos
}
func decodeHeads(raw []float32, n, seqDim int) []int {
heads := make([]int, n)
for i := 0; i < n; i++ {
bestIdx := 0
bestVal := float32(-1e9)
for j := 0; j < seqDim; j++ {
if v := raw[i*seqDim+j]; v > bestVal {
bestVal = v
bestIdx = j
}
}
heads[i] = bestIdx
}
return heads
}
func decodeRels(raw []float32, n, seqDim, relDim int, heads []int) []string {
rels := make([]string, n)
stride := seqDim * relDim
for i := 0; i < n; i++ {
h := heads[i]
if h < 0 || h >= seqDim {
rels[i] = "dep"
continue
}
bestIdx := 0
bestVal := float32(-1e9)
for r := 0; r < relDim; r++ {
if v := raw[i*stride+h*relDim+r]; v > bestVal {
bestVal = v
bestIdx = r
}
}
rels[i] = depRelLabel(bestIdx)
}
return rels
}
func loadWordMap(path string) (map[string]int64, error) {
data, err := onnxModelFS.ReadFile(path)
if err != nil {
return nil, err
}
var raw struct {
Word map[string]int64 `json:"word"`
}
if err := json.Unmarshal(data, &raw); err != nil {
var flat map[string]int64
if err2 := json.Unmarshal(data, &flat); err2 != nil {
return nil, err
}
return flat, 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 libPath() string {
for _, env := range []string{"ONNXRUNTIME_DIR", "ONNX_ML_DIR"} {
if d := os.Getenv(env); d != "" {
for _, name := range []string{"libonnxruntime.so", "libonnxruntime.dylib", "onnxruntime.dll"} {
if candidate := filepath.Join(d, name); fileExists(candidate) {
return candidate
}
}
}
}
for _, name := range []string{"libonnxruntime.so", "libonnxruntime.dylib", "onnxruntime.dll"} {
if fileExists(name) {
abs, _ := filepath.Abs(name)
return abs
}
}
return "onnxruntime.dll"
}
func fileExists(p string) bool {
_, err := os.Stat(p)
return err == nil
}
func depRelLabel(id int) string {
labels := []string{"root", "nsubj", "obj", "iobj", "obl", "vocative", "expl", "csubj", "ccomp", "xcomp",
"advcl", "advmod", "amod", "appos", "nmod", "acl", "det", "clf", "case", "mark",
"nummod", "discourse", "aux", "cop", "cc", "conj", "fixed", "flat", "list", "parataxis",
"orphan", "goeswith", "reparandum", "punct", "dep"}
if id >= 0 && id < len(labels) {
return labels[id]
}
return "dep"
}