mirror of
https://gitcode.com/JianFeeeee/HomeAgent.git
synced 2026-10-03 15:53:56 +00:00
## 为什么 用户决定「本轮不覆盖 video,先支持 text+image」。这一刀正好解锁了此前 「小 + 可商用 + 覆盖视频」三者不可兼得的僵局:不要求视频后,唯一同时满足 **小、可商用、中文原生** 的选项是 Chinese-CLIP ViT-B/16。 实测对比(同机、真实跑出来的数字): | | Chinese-CLIP | jina-v5-omni-nano | Qwen3-VL-Emb-2B | |---|---|---|---| | 参数量 | 188M | 1.04B | 2B | | 产物 / 常驻内存 | 754MB / **1.15GB** | ~2GB / 2.23GB | 8GB / 9.4GB | | 维度 | 512 | 768 | 2048 | | 许可 | **Apache-2.0** | CC BY-NC(不可商用) | Apache-2.0 | | 视频 | 无 | 有 | 有 | 本机可用内存只有 5.3GB,Qwen 的 9.4GB 无法进程内使用;而 ORT format + mmap 那条路被证实当前不通(转换器对三段图段错误;走通还需同时升 ORT 运行时与 Go 绑定,v1.36 要求 API 29 而本机只有 28)。1.15GB 则可以直接进程内跑。 **代价已写进包注释与文档**:CLIP 是双塔对比学习,text↔image 是强项,但纯文本 语义明显弱于 MLLM 型嵌入器;文本检索仍由既有词向量/TF-IDF 路径兜底。 需要更强文本语义或视频时切回 qwen3vl。 ## 内容 - `providers/chineseclip/`:按公共 SPI 实现的 provider(注册名 `chineseclip`), 含 BERT WordPiece 分词器、图像预处理、ONNX 双塔推理、无标签 stub。 - `scripts/export_chineseclip_onnx.py`:从官方权重导出规范产物 + 冻结参考, 自带逐用例 PyTorch 对比与覆盖度断言(计划集合≠执行集合即非零退出)。 - `cmd/homed/main.go`:空白导入两个 provider,由配置选其一。 - `go.mod`:`golang.org/x/text` 由间接依赖转为直接依赖(删音标需要 NFD)。 ## 实现要点 - **分词器逐 token 对齐官方**。第一版探针自己拼 BertTokenizer(只给 vocab.txt、 没删音标、中文没逐字切),中文被整体切成 [UNK],三个不同句子产出几乎相同的 向量(余弦 0.98)——差点把「模型坏了」当成结论。官方配置是 do_lower_case=true + 删音标生效 + 中文逐字切分;`TestTokenizerMatchesOfficialReference` 钉住 逐 token 一致。 - **图像缩放自写 bicubic**(复刻 PIL 的 precompute_coeffs + a=-0.5 核),不引 golang.org/x/image:它未进本机模块缓存,且最新版要求把整个工具链升到 Go 1.26, 为一个缩放函数动工具链不划算。 - **归一化在 provider 侧**(两个塔的图里都没归一化),检索按余弦。 - **指纹覆盖全部影响语义的产物**:两个 ONNX 图 + vocab.txt + embed_config.json, 读不到就写 MISSING(跳过等于对缺件不敏感)。 - 会话 Run 用 runMu 串行化(ORT 会话不保证并发安全),创建/销毁用 mu。 ## 模态范围 只声明 `text` 与 `image`;`audio`/`video` 明确返回 `ErrUnsupportedModality`, 绝不用别的模型向量冒充(这是「音频明确 unsupported」纪律的落地)。 ## 验证(实测) 导出侧:10 个用例(5 文本 + 5 图像)ONNX vs 官方 PyTorch 全部 `cos = 1.000000000`,覆盖度断言 10/10 通过。 Go 侧(`CHINESECLIP_MODEL_DIR=... go test -tags onnxruntime ./providers/chineseclip/ -v`): 11/11 通过,其中 - 文本 5 用例 `cos = 1.000000000000`(逐位一致) - 图像 4 纯色用例 `cos = 1.000000`(与官方预处理在 6 位小数内一致) - 跨模态判别:红图对「红色」文本高于「蓝色」文本 - 模态拒绝 / 空输入 / 指纹稳定 / 产物缺失报错 顺带修掉测试自身的一个假通过:参考向量是**未归一化**的原始输出(模长 10~36), 原先「点积当余弦 + 单侧下界」会让 13.6 也判过,已改为真余弦 + 双侧容差。 构建矩阵:`go build/vet ./...` 与 `-tags onnxruntime` 两种都过; `providers/... pkg/... internal/config/... internal/memory/vector/...` 回归通过 (qwen3vl 的 TestVideoModelInputMRope 需要 QWEN_ONNX_MODEL_DIR 指向含视频档的 v3 目录,缺该环境变量时用的是只有文本+图像的目录,与本改动无关)。 ## 未做(明确记录) - 发行版默认 provider 与构建标签变更:留下一提交(涉及打包与模型分发策略)。 - 模型产物(754MB)不进仓库,由导出脚本生成。
283 lines
7.8 KiB
Go
283 lines
7.8 KiB
Go
//go:build onnxruntime
|
||
|
||
package chineseclip
|
||
|
||
import (
|
||
"context"
|
||
"crypto/sha256"
|
||
"encoding/hex"
|
||
"fmt"
|
||
"math"
|
||
"os"
|
||
"path/filepath"
|
||
"strings"
|
||
"sync"
|
||
|
||
ort "github.com/yalue/onnxruntime_go"
|
||
|
||
"gitcode.com/JianFeeeee/HomeAgent/pkg/embedding"
|
||
)
|
||
|
||
func init() {
|
||
embedding.Register("chineseclip", func(cfg embedding.Config) (embedding.Provider, error) {
|
||
return New(cfg.Options["model_dir"])
|
||
})
|
||
}
|
||
|
||
// Embedder 是 Chinese-CLIP ViT-B/16 的进程内 provider。
|
||
//
|
||
// 并发:ONNX Runtime 的会话不保证多次 Run 可并发,故运行时用 runMu 串行化;
|
||
// 创建/销毁会话用 mu 保护。SPI 要求实现可安全并发调用,这里由我们自己保证。
|
||
type Embedder struct {
|
||
mu sync.RWMutex
|
||
runMu sync.Mutex
|
||
|
||
dir string
|
||
config embedConfig
|
||
tok *Tokenizer
|
||
|
||
text *ort.DynamicAdvancedSession
|
||
vision *ort.DynamicAdvancedSession
|
||
|
||
fp string
|
||
closeOnce sync.Once
|
||
}
|
||
|
||
// New 从产物目录构造 provider。
|
||
func New(modelDir string) (*Embedder, error) {
|
||
modelDir = strings.TrimSpace(modelDir)
|
||
if modelDir == "" {
|
||
return nil, fmt.Errorf("chineseclip: 未配置 model_dir(产物目录)")
|
||
}
|
||
cfg, err := loadConfig(modelDir)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
tok, err := LoadTokenizer(modelDir, cfg.MaxLength)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
if !ort.IsInitialized() {
|
||
if lib := findOnnxLib(); lib != "" {
|
||
ort.SetSharedLibraryPath(lib)
|
||
}
|
||
if err := ort.InitializeEnvironment(); err != nil {
|
||
return nil, fmt.Errorf("chineseclip: 初始化 onnx 环境: %w", err)
|
||
}
|
||
}
|
||
|
||
text, err := ort.NewDynamicAdvancedSession(
|
||
filepath.Join(modelDir, cfg.TextONNX),
|
||
[]string{"input_ids", "attention_mask"},
|
||
[]string{"text_features"}, nil,
|
||
)
|
||
if err != nil {
|
||
return nil, fmt.Errorf("chineseclip: 创建文本塔会话(%s): %w", cfg.TextONNX, err)
|
||
}
|
||
vision, err := ort.NewDynamicAdvancedSession(
|
||
filepath.Join(modelDir, cfg.VisionONNX),
|
||
[]string{"pixel_values"},
|
||
[]string{"image_features"}, nil,
|
||
)
|
||
if err != nil {
|
||
text.Destroy()
|
||
return nil, fmt.Errorf("chineseclip: 创建视觉塔会话(%s): %w", cfg.VisionONNX, err)
|
||
}
|
||
|
||
return &Embedder{
|
||
dir: modelDir,
|
||
config: cfg,
|
||
tok: tok,
|
||
text: text,
|
||
vision: vision,
|
||
fp: computeFingerprint(modelDir, cfg),
|
||
}, nil
|
||
}
|
||
|
||
// Embed 按模态分派。audio/video 一律返回 ErrUnsupportedModality——
|
||
// 本空间没有它们的原生编码器,用别的模型向量冒充会污染整个向量空间。
|
||
func (e *Embedder) Embed(ctx context.Context, in embedding.Input) ([]float64, error) {
|
||
if err := ctx.Err(); err != nil {
|
||
return nil, err
|
||
}
|
||
switch in.Modality {
|
||
case embedding.ModalityText:
|
||
if strings.TrimSpace(in.Text) == "" {
|
||
return nil, fmt.Errorf("chineseclip: 文本输入为空")
|
||
}
|
||
return e.embedText(in.Text)
|
||
case embedding.ModalityImage:
|
||
if len(in.Data) == 0 {
|
||
return nil, fmt.Errorf("chineseclip: 图像输入为空(modality=image 需要 Data)")
|
||
}
|
||
return e.embedImage(in.Data)
|
||
default:
|
||
return nil, fmt.Errorf("chineseclip: %w: %s", embedding.ErrUnsupportedModality, in.Modality)
|
||
}
|
||
}
|
||
|
||
func (e *Embedder) embedText(text string) ([]float64, error) {
|
||
e.mu.RLock()
|
||
sess, tok, dim := e.text, e.tok, e.config.Dimension
|
||
e.mu.RUnlock()
|
||
if sess == nil {
|
||
return nil, fmt.Errorf("chineseclip: provider 已关闭")
|
||
}
|
||
|
||
ids, mask := tok.Encode(text)
|
||
shape := ort.Shape{1, int64(len(ids))}
|
||
idTensor, err := ort.NewTensor(shape, ids)
|
||
if err != nil {
|
||
return nil, fmt.Errorf("chineseclip: 构造 input_ids 张量: %w", err)
|
||
}
|
||
defer idTensor.Destroy()
|
||
maskTensor, err := ort.NewTensor(shape, mask)
|
||
if err != nil {
|
||
return nil, fmt.Errorf("chineseclip: 构造 attention_mask 张量: %w", err)
|
||
}
|
||
defer maskTensor.Destroy()
|
||
|
||
outs := make([]ort.Value, 1)
|
||
e.runMu.Lock()
|
||
err = sess.Run([]ort.Value{idTensor, maskTensor}, outs)
|
||
e.runMu.Unlock()
|
||
if err != nil {
|
||
return nil, fmt.Errorf("chineseclip: 文本塔推理: %w", err)
|
||
}
|
||
if outs[0] == nil {
|
||
return nil, fmt.Errorf("chineseclip: 文本塔输出为空")
|
||
}
|
||
defer outs[0].Destroy()
|
||
return normalizeOutput(outs[0], 1, dim, "text_features")
|
||
}
|
||
|
||
func (e *Embedder) embedImage(data []byte) ([]float64, error) {
|
||
e.mu.RLock()
|
||
sess, cfg, dim := e.vision, e.config, e.config.Dimension
|
||
e.mu.RUnlock()
|
||
if sess == nil {
|
||
return nil, fmt.Errorf("chineseclip: provider 已关闭")
|
||
}
|
||
|
||
pixels, err := preprocessImage(data, cfg.ImageSize, cfg.ImageMean, cfg.ImageStd)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
shape := ort.Shape{1, 3, int64(cfg.ImageSize), int64(cfg.ImageSize)}
|
||
in, err := ort.NewTensor(shape, pixels)
|
||
if err != nil {
|
||
return nil, fmt.Errorf("chineseclip: 构造 pixel_values 张量: %w", err)
|
||
}
|
||
defer in.Destroy()
|
||
|
||
outs := make([]ort.Value, 1)
|
||
e.runMu.Lock()
|
||
err = sess.Run([]ort.Value{in}, outs)
|
||
e.runMu.Unlock()
|
||
if err != nil {
|
||
return nil, fmt.Errorf("chineseclip: 视觉塔推理: %w", err)
|
||
}
|
||
if outs[0] == nil {
|
||
return nil, fmt.Errorf("chineseclip: 视觉塔输出为空")
|
||
}
|
||
defer outs[0].Destroy()
|
||
return normalizeOutput(outs[0], 1, dim, "image_features")
|
||
}
|
||
|
||
// normalizeOutput 取出 [batch, dim] 输出并做 L2 归一化。
|
||
//
|
||
// 官方 Chinese-CLIP 的检索用法就是余弦相似度(归一化后点积),
|
||
// 归档前统一归一化可以避免下游反复判断。
|
||
func normalizeOutput(value ort.Value, batch, dim int, name string) ([]float64, error) {
|
||
tensor, ok := value.(*ort.Tensor[float32])
|
||
if !ok {
|
||
return nil, fmt.Errorf("chineseclip: %s 输出类型 %T,期望 float32 张量", name, value)
|
||
}
|
||
shape := tensor.GetShape()
|
||
if len(shape) != 2 || shape[0] != int64(batch) || shape[1] != int64(dim) {
|
||
return nil, fmt.Errorf("chineseclip: %s 形状 %v,期望 [%d %d]", name, shape, batch, dim)
|
||
}
|
||
raw := tensor.GetData()
|
||
if len(raw) < batch*dim {
|
||
return nil, fmt.Errorf("chineseclip: %s 数据长度 %d,期望 %d", name, len(raw), batch*dim)
|
||
}
|
||
out := make([]float64, dim)
|
||
var norm float64
|
||
for i := 0; i < dim; i++ {
|
||
v := float64(raw[i])
|
||
out[i] = v
|
||
norm += v * v
|
||
}
|
||
norm = math.Sqrt(norm)
|
||
if norm == 0 || math.IsNaN(norm) || math.IsInf(norm, 0) {
|
||
return nil, fmt.Errorf("chineseclip: %s 向量范数为 %v(模型输出异常)", name, norm)
|
||
}
|
||
for i := range out {
|
||
out[i] /= norm
|
||
}
|
||
return out, nil
|
||
}
|
||
|
||
// Info 只声明本 provider 能通过公共契约提供的模态。
|
||
//
|
||
// 契约要求 provider 自行解码 Data;这里没有视频/音频解码器,列进来只会让核心
|
||
// 据以创建输入、然后在运行时全部失败。audio/video 必须返回 ErrUnsupportedModality。
|
||
func (e *Embedder) Info() embedding.Info {
|
||
e.mu.RLock()
|
||
defer e.mu.RUnlock()
|
||
return embedding.Info{
|
||
Dimension: e.config.Dimension,
|
||
Fingerprint: e.fp,
|
||
Modalities: []embedding.Modality{
|
||
embedding.ModalityText,
|
||
embedding.ModalityImage,
|
||
},
|
||
}
|
||
}
|
||
|
||
func (e *Embedder) Close() {
|
||
e.closeOnce.Do(func() {
|
||
e.mu.Lock()
|
||
defer e.mu.Unlock()
|
||
if e.text != nil {
|
||
e.text.Destroy()
|
||
e.text = nil
|
||
}
|
||
if e.vision != nil {
|
||
e.vision.Destroy()
|
||
e.vision = nil
|
||
}
|
||
})
|
||
}
|
||
|
||
// computeFingerprint 覆盖**全部**影响向量语义的产物:两个 ONNX 图、词表与配置。
|
||
// 漏掉任何一个都会让「换了模型但指纹没变」,历史向量不会重算。
|
||
func computeFingerprint(modelDir string, cfg embedConfig) string {
|
||
h := sha256.New()
|
||
for _, name := range []string{cfg.TextONNX, cfg.VisionONNX, "vocab.txt", "embed_config.json"} {
|
||
data, err := os.ReadFile(filepath.Join(modelDir, name))
|
||
if err != nil {
|
||
// 读不到就写名字+错误,绝不跳过:跳过等于指纹对缺件不敏感。
|
||
fmt.Fprintf(h, "%s:MISSING:%v\n", name, err)
|
||
continue
|
||
}
|
||
fmt.Fprintf(h, "%s:%d\n", name, len(data))
|
||
h.Write(data)
|
||
h.Write([]byte{0})
|
||
}
|
||
return hex.EncodeToString(h.Sum(nil))
|
||
}
|
||
|
||
func findOnnxLib() string {
|
||
for _, p := range []string{
|
||
"/opt/onnxruntime/libonnxruntime.so",
|
||
"/usr/local/lib/libonnxruntime.so",
|
||
"/usr/lib/libonnxruntime.so",
|
||
} {
|
||
if _, err := os.Stat(p); err == nil {
|
||
return p
|
||
}
|
||
}
|
||
return ""
|
||
}
|