mirror of
https://gitcode.com/JianFeeeee/HomeAgent.git
synced 2026-09-21 09:28:14 +00:00
用户要求:后续发行版默认带 ONNX 模型能力。这条要求落到两处,并顺带修掉一个 被它**暴露出来**的真缺陷。 ## 1. 构建默认带 onnxruntime(deploy/packaging/build.sh) `HOMED_TAGS` 默认 `onnxruntime`,需要极简构建时显式 `HOMED_TAGS=` 关闭。 不带标签时 provider 仍注册、但打开即报「requires build tag」并优雅降级—— 不静默假装成功。运行期还需要 `libonnxruntime.so`(provider 按 /opt/onnxruntime、/usr/local/lib、/usr/lib 顺序查找),缺失时同样是 「日志里明确错误 + 降级」。 ## 2. 新装默认选 chineseclip(internal/config/registry.go) `SeedDefaults` 写入: core.memory.multimodal_space.provider = chineseclip core.memory.multimodal_space.options.model_dir = <dataDir>/models/chinese-clip-vit-b16-onnx 选它而不是 qwen3vl:后者实测常驻 9.4GB,多数机器装不下;chineseclip 是 1.99GB(实测,见下)。同时更新两个 ConfigDef 的默认值与描述(WebUI 显示用)。 **老安装不会自动拿到这两个默认值**,这是有意的:`seedDBValues` 对非空配置库 直接返回,`GetString` 缺键时回落到调用方默认值(main.go 传的是空串)。 升级就静默加载 ~1.8GB 模型不是无副作用的事,应由部署显式开启。已写进文档。 ## 3. 修掉 ORT 环境被重复初始化 + 误销毁(internal/nlp/onnx.go) 这是「默认带标签」才暴露的缺陷:此前不带标签时进程内不会有多个 ORT 消费者。 - `NewONNXParser` 无条件 `InitializeEnvironment()` → 若多模态 provider 先初始化, 这里报「The onnxruntime has already been initialized」并**降级**(实测日志: `ONNX parser init: init onnx env: ... using fallback`)。 - 更严重的是失败路径与 `Close()` 里的 `DestroyEnvironment()`:它会把别人 (多模态 provider)正在用的进程级环境一起拆掉,让对方的会话失效。 改为:初始化前先 `IsInitialized()`;**任何消费者都不销毁环境**(随进程存活), 只销毁自己的会话。providers/chineseclip 与 providers/qwen3vl 本来就是这个约定, 现在三处一致。 ## 验证(实测) - 全新数据目录启动:配置库出现上述两个默认值。 - 模型未安装:`multimodal space active` 不出现,代之以明确错误 (点名缺失的 embed_config.json 路径 + 已注册 provider 列表)+ 降级,不静默。 - 模型就位:`multimodal space active: provider=chineseclip dim=512 fp=cd2a495cf990 modalities=[text image]`。 - 内存:同一份 homed,启用时 RSS **1.99GB**(峰值 2.09GB),不启用 **0.17GB**。 - NLP 修复:日志由 `ONNX parser init: ... using fallback` 变为 `dep parser initialized`。 - 构建矩阵:`go build/vet ./...` 与 `-tags onnxruntime` 两种都过; `bash -n deploy/packaging/build.sh` 通过。 ## 未做(明确记录) - `libonnxruntime.so`(24MB)与 Chinese-CLIP 产物(754MB)目前都需自行安装/导出, 发行版尚未打包它们。若要让「默认启用」在干净机器上真正开箱可用,需要决定 是随包分发、安装时下载、还是保持文档指引。
335 lines
8.1 KiB
Go
335 lines
8.1 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 环境是**进程级单例**,同一进程里可能有多个消费者(多模态向量
|
||
// provider、本依存解析器)。onnxruntime_go 的行为是:第二次
|
||
// InitializeEnvironment 报「already been initialized」,而
|
||
// DestroyEnvironment 会把别人正在用的环境一起拆掉——先初始化的 provider
|
||
// 会因此拿到失效的会话。所以这里只在未初始化时初始化,并且**永不销毁**
|
||
// (与 providers/chineseclip、providers/qwen3vl 的约定一致):环境随进程存活。
|
||
//
|
||
// 这个缺陷是在「发行版默认带 onnxruntime 标签」后才暴露的:不带标签时
|
||
// 两个消费者不会同时存在,重复初始化与误销毁都无法发生。
|
||
if !ort.IsInitialized() {
|
||
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 {
|
||
// 不在这里 DestroyEnvironment:环境是进程级的,可能正被多模态 provider 使用。
|
||
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 环境:多模态 provider 可能仍在使用(见 NewONNXParser)。
|
||
})
|
||
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"
|
||
}
|