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 传参
This commit is contained in:
JianFeeeee
2026-07-28 11:42:29 +08:00
parent 2c5f9ff262
commit f91b20ee16
24 changed files with 334 additions and 824 deletions

View File

@ -1,103 +0,0 @@
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
}

View File

@ -8,32 +8,35 @@ import (
"fmt"
"os"
"path/filepath"
"sync"
"gitcode.com/JianFeeeee/HomeAgent/internal/config"
"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.AdvancedSession
vocab map[string]int64
rt *ort.DynamicAdvancedSession
vocab map[string]int64
posVocab map[string]int64
Release func()
close sync.Once
}
type ONNXConfig struct {
ModelPath string // 留空使用内嵌模型
DataDir string // 模型解压/缓存目录
ModelPath string
DataDir string
}
func NewONNXParser(cfg ONNXConfig) (*ONNXParser, error) {
vocab, err := loadJSONMap[int64]("models/vocab.json", onnxModelFS)
vocab, err := loadWordMap("models/vocab.json")
if err != nil {
return nil, fmt.Errorf("load vocab: %w", err)
}
posVocab, err := loadJSONMap[int64]("models/pos_vocab.json", onnxModelFS)
posVocab, err := loadWordMap("models/pos_vocab.json")
if err != nil {
return nil, fmt.Errorf("load pos_vocab: %w", err)
}
@ -46,83 +49,221 @@ func NewONNXParser(cfg ONNXConfig) (*ONNXParser, error) {
}
}
ort.SetSharedLibraryPath(findONNXRuntime())
ort.SetSharedLibraryPath(libPath())
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})
inputNames := []string{"input_ids"}
outputNames := []string{"pos_logits", "head_logits", "rel_logits"}
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)
session, err := ort.NewDynamicAdvancedSession(modelPath, inputNames, outputNames, nil)
if err != nil {
return nil, fmt.Errorf("create session: %w", err)
}
release := func() {
session.Destroy()
ort.DestroyEnvironment()
return nil, fmt.Errorf("create session: %w", err)
}
return &ONNXParser{
rt: session,
vocab: vocab,
posVocab: posVocab,
Release: release,
}, 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
}
inputIDs := tokenize(text, p.vocab, 128)
inputIDs = padTo(inputIDs, 128)
x := memory.GetJieba()
if x == nil {
return nil, fmt.Errorf("jieba unavailable")
}
words := x.Cut(text, true)
if len(words) == 0 {
return &ParseResult{}, nil
}
inputTensor, err := ort.NewTensor(ort.NewShape(1, 128), inputIDs)
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 inputTensor.Destroy()
defer inTensor.Destroy()
outputs, err := p.rt.Call(inputTensor)
if err != nil {
return nil, fmt.Errorf("onnx call: %w", err)
outputs := make([]ort.Value, 3)
if err := p.rt.Run([]ort.Value{inTensor}, outputs); err != nil {
return nil, fmt.Errorf("onnx run: %w", err)
}
rawPOS := outputs[0].GetData().([]float32)
rawHeads := outputs[1].GetData().([]float32)
rawRels := outputs[2].GetData().([]float32)
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()
seqLen := actualLen(inputIDs)
tokens := idsToTokens(inputIDs[:seqLen], p.vocab)
pos := decodePOS(rawPOS, seqLen, p.posVocab)
heads := decodeHeads(rawHeads, seqLen)
rels := decodeRels(rawRels, seqLen)
posShape := posOut.GetShape() // [1, seq, posDim]
headShape := headOut.GetShape() // [1, seq, seq]
relShape := relOut.GetShape() // [1, seq, seq, relDim]
return &ParseResult{Tokens: tokens, POS: pos, Heads: heads, DepRels: rels}, nil
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 loadJSONMap[T ~int64 | ~string](path string, fs embed.FS) (map[string]T, error) {
data, err := fs.ReadFile(path)
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]T `json:"word"`
Word map[string]int64 `json:"word"`
}
if err := json.Unmarshal(data, &raw); err != nil {
result := make(map[string]T)
if err2 := json.Unmarshal(data, &result); err2 != nil {
var flat map[string]int64
if err2 := json.Unmarshal(data, &flat); err2 != nil {
return nil, err
}
return result, nil
return flat, nil
}
return raw.Word, nil
}
@ -146,132 +287,37 @@ func extractEmbeddedModel(dataDir string) (string, error) {
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"),
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 _, c := range candidates {
if _, err := os.Stat(c); err == nil {
abs, _ := filepath.Abs(c)
for _, name := range []string{"libonnxruntime.so", "libonnxruntime.dylib", "onnxruntime.dll"} {
if fileExists(name) {
abs, _ := filepath.Abs(name)
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 fileExists(p string) bool {
_, err := os.Stat(p)
return err == nil
}
func padTo(ids []int64, length int) []int64 {
for len(ids) < length {
ids = append(ids, 0)
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 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"
return "dep"
}

View File

@ -2,264 +2,19 @@
package nlp
import (
"embed"
"encoding/json"
"fmt"
"strings"
)
import "fmt"
//go:embed models/vocab.json models/pos_vocab.json
var vocabFS embed.FS
// ONNXParser 在未启用 onnxruntime 时作为规则式降级解析器。
// 使用内嵌词表实现基于词典的 POS 标注 + 基于 POS 序列的依存关系推断。
type ONNXParser struct {
vocab map[string]int
posVocab map[string]int
}
type ONNXParser struct{}
type ONNXConfig struct {
ModelPath string // 留空使用内嵌规则引擎
DataDir string // 仅在 onnxruntime 启用时使用
ModelPath string
DataDir string
}
func NewONNXParser(cfg ONNXConfig) (*ONNXParser, error) {
vocab := make(map[string]int)
data, err := vocabFS.ReadFile("models/vocab.json")
if err != nil {
return nil, fmt.Errorf("read vocab: %w", err)
}
var raw struct {
Word map[string]int `json:"word"`
}
if err := json.Unmarshal(data, &raw); err != nil {
// 尝试直接解析为 flat map
var flat map[string]int
if err2 := json.Unmarshal(data, &flat); err2 != nil {
return nil, fmt.Errorf("parse vocab: %w", err)
}
vocab = flat
} else {
vocab = raw.Word
}
posVocab := make(map[string]int)
data, err = vocabFS.ReadFile("models/pos_vocab.json")
if err != nil {
return nil, fmt.Errorf("read pos_vocab: %w", err)
}
if err := json.Unmarshal(data, &posVocab); err != nil {
return nil, fmt.Errorf("parse pos_vocab: %w", err)
}
return &ONNXParser{vocab: vocab, posVocab: posVocab}, nil
func NewONNXParser(_ ONNXConfig) (*ONNXParser, error) {
return nil, fmt.Errorf("ONNX parser requires build tag 'onnxruntime' (go build -tags onnxruntime)")
}
func (p *ONNXParser) Parse(text string) (*ParseResult, error) {
if text == "" {
return &ParseResult{}, nil
}
// Phase 1: 基于词表的最大匹配分词
tokens := p.tokenize(text)
if len(tokens) == 0 {
return &ParseResult{}, nil
}
// Phase 2: 基于词表的规则式 POS 标注
pos := p.tagPOS(tokens)
// Phase 3: 基于 POS 序列的依存头推断
heads := p.inferHeads(tokens, pos)
// Phase 4: 关系标签推断
rels := p.inferRels(tokens, pos, heads)
return &ParseResult{
Tokens: tokens,
POS: pos,
Heads: heads,
DepRels: rels,
}, nil
}
func (p *ONNXParser) tokenize(text string) []string {
runes := []rune(text)
var tokens []string
buf := []rune{}
for _, r := range runes {
if r == ' ' || r == '\t' || r == '\n' || r == '\r' {
if len(buf) > 0 {
tokens = append(tokens, string(buf))
buf = buf[:0]
}
continue
}
buf = append(buf, r)
// 最长匹配:检查当前 buf 是否在词表中
if _, ok := p.vocab[string(buf)]; !ok && len(buf) > 0 {
// 回退:取 buf[:-1] 作为词,继续
if _, ok2 := p.vocab[string(buf[:len(buf)-1])]; ok2 && len(buf) > 2 {
tokens = append(tokens, string(buf[:len(buf)-1]))
buf = buf[len(buf)-1:]
}
}
}
if len(buf) > 0 {
tokens = append(tokens, string(buf))
}
if len(tokens) == 0 {
tokens = strings.Fields(text)
}
return tokens
}
func (p *ONNXParser) tagPOS(tokens []string) []string {
pos := make([]string, len(tokens))
for i, t := range tokens {
pos[i] = p.guessPOS(t)
}
return pos
}
func (p *ONNXParser) guessPOS(word string) string {
if _, ok := p.vocab[word]; !ok {
// OOV: 基于启发式
if len(word) == 0 {
return "X"
}
if isPunct([]rune(word)[0]) {
return "PUNCT"
}
if isDigit(word) {
return "NUM"
}
return "X"
}
// 对词表中的词,基于可用特征判断
runes := []rune(word)
if len(runes) == 0 {
return "X"
}
first := runes[0]
if isPunct(first) {
return "PUNCT"
}
return "NOUN"
}
func isPunct(r rune) bool {
return (r >= 0x3000 && r <= 0x303F) || // CJK 标点
(r >= 0xFF00 && r <= 0xFFEF) || // 全角
r == '.' || r == ',' || r == '!' || r == '?' ||
r == ';' || r == ':' || r == '"' || r == '\'' ||
r == '(' || r == ')' || r == '[' || r == ']' ||
r == '{' || r == '}' || r == '。' || r == ',' ||
r == '!' || r == '?' || r == ';' || r == ':' ||
r == '、' || r == '‘' || r == '’' || r == '“' || r == '”'
}
func isDigit(s string) bool {
for _, r := range s {
if r < '0' || r > '9' {
if r < 0xFF10 || r > 0xFF19 { // 全角数字
return false
}
}
}
return len(s) > 0
}
// inferHeads 基于 POS 序列的规则式依存头推断。
// 动词通常作为根(head=0),名词依附于动词,形容词依附于名词。
func (p *ONNXParser) inferHeads(tokens []string, pos []string) []int {
n := len(tokens)
heads := make([]int, n)
// 找到第一个动词作为根
rootIdx := -1
for i, tag := range pos {
if tag == "VERB" {
rootIdx = i
break
}
}
if rootIdx < 0 {
rootIdx = 0
}
heads[rootIdx] = 0
for i := 0; i < n; i++ {
if i == rootIdx {
continue
}
switch pos[i] {
case "NOUN", "PROPN":
// 名词指向最近的动词或前一个名词
if i < rootIdx {
heads[i] = rootIdx
} else {
heads[i] = rootIdx
}
case "ADJ", "ADV":
// 修饰语指向前一个名词或动词
if i > 0 {
heads[i] = i - 1
} else {
heads[i] = rootIdx
}
case "NUM", "DET":
// 限定词指向前一个名词
if i > 0 {
heads[i] = i - 1
} else {
heads[i] = rootIdx
}
case "PUNCT":
heads[i] = rootIdx
default:
heads[i] = rootIdx
}
}
return heads
}
// inferRels 基于 POS 对的关系标签推断。
func (p *ONNXParser) inferRels(tokens []string, pos []string, heads []int) []string {
n := len(tokens)
rels := make([]string, n)
for i := 0; i < n; i++ {
if heads[i] == 0 {
rels[i] = "ROOT"
continue
}
h := heads[i]
if h < 0 || h >= n {
rels[i] = "dep"
continue
}
rels[i] = posToRel(pos[h], pos[i])
}
return rels
}
func posToRel(headPOS, depPOS string) string {
switch {
case depPOS == "NOUN" || depPOS == "PROPN":
return "nsubj"
case depPOS == "ADJ":
return "amod"
case depPOS == "ADV":
return "advmod"
case depPOS == "NUM" || depPOS == "DET":
return "det"
case depPOS == "VERB":
return "xcomp"
case depPOS == "PUNCT":
return "punct"
default:
return "dep"
}
func (p *ONNXParser) Parse(_ string) (*ParseResult, error) {
return nil, fmt.Errorf("ONNX parser not available: rebuild with -tags onnxruntime")
}