feat(memory): 新增 chineseclip provider —— text+image 的小体积可商用向量空间

## 为什么

用户决定「本轮不覆盖 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)不进仓库,由导出脚本生成。
This commit is contained in:
JianFeeeee
2026-09-11 23:58:53 +08:00
parent f868975c0e
commit 929fb94e7a
11 changed files with 1757 additions and 4 deletions

View File

@ -48,6 +48,7 @@ import (
// 空白导入内置 provider:它们各自在 init 里注册到 pkg/embedding。
// 想把核心换成自己的模型,只需替换这一行(或另建一个发行版 main)。
_ "gitcode.com/JianFeeeee/HomeAgent/providers/chineseclip"
_ "gitcode.com/JianFeeeee/HomeAgent/providers/qwen3vl"
)

View File

@ -1,6 +1,18 @@
# 统一多模态向量空间(Qwen3-VL-Embedding-2B)
# 统一多模态向量空间
文本、图像、**视频帧** 在同一模型、同一 2048 维、同一 fingerprint 空间里被编码。
核心不绑定任何具体模型:它按 provider 名从公共注册表(`pkg/embedding`)打开一个
向量空间。仓库内自带两个:
| provider | 模态 | 维度 | 实测常驻 | 许可 | 适用 |
|---|---|---|---|---|---|
| `chineseclip` | text + image | 512 | **1.15 GB** | Apache-2.0 | 默认(内存受限 / 中文图文) |
| `qwen3vl` | text + image(视频已实现未纳入契约) | 2048 | 9.4 GB | Apache-2.0 | 内存充足 / 需要更强文本语义或视频 |
| `http` | 由外部服务决定 | 由外部服务决定 | 由外部服务决定 | — | 侧车部署(如 jina-v5-omni-nano,注意其 CC BY-NC 许可) |
下面第一节是 Qwen3-VL(2048 维,最强但最重),第二节是 Chinese-CLIP(512 维,
默认推荐)。两者互斥启用,改配置后重启生效。
文本、图像、**视频帧** 在同一模型、同一维度、同一 fingerprint 空间里被编码。
记忆系统用它做三件事:多模态图记忆的跨模态召回、multimodal doc 的向量融合、
multimodal context 的相关性裁剪/淘汰。
@ -76,6 +88,94 @@ axis,实际却只能用导出的那个长度运行。
「能加载」不等于「算得对」:形状错、输入名错、池化位置错的图都能正常 load。
## 一·补、text+image 默认空间:Chinese-CLIP ViT-B/16
**为什么它是默认**:text+image 只需要一个向量空间时,同时满足「小、可商用、中文原生」
的选项只有一个。
| | Chinese-CLIP | jina-v5-omni-nano | Qwen3-VL-Emb-2B |
|---|---|---|---|
| 参数量 | 188M | 1.04B | 2B |
| 产物 / 实测常驻 | **721MB / 1.15GB** | ~2GB / 2.23GB | 8GB / 9.4GB |
| 维度 | 512 | 768 | 2048 |
| 许可 | **Apache-2.0** | CC BY-NC(不可商用) | Apache-2.0 |
| 中文 | 原生(~2 亿中文图文对) | 多语言 | 多语言 |
| 文本语义 | 弱(双塔对比) | 好 | 最好 |
| 视频 | 无 | 有 | 有 |
**要诚实记录的代价**:CLIP 是双塔对比学习,text↔image 是强项,但**纯文本语义
(text↔text)明显弱于 MLLM 型嵌入器**。文本检索仍由既有词向量/TF-IDF 路径兜底,
本空间主要用于跨模态召回与相关性裁剪。需要更强文本语义或视频时切回 `qwen3vl`。
### 产物与获取
产物约 754MB,**不进仓库**;用导出脚本从官方权重导出(脚本入库,保证可复现):
```bash
python3 scripts/export_chineseclip_onnx.py \
--model-dir /path/to/chinese-clip-vit-base-patch16 \
--out /home/newqqagent/models/chinese-clip-vit-b16-onnx
```
国内下载:本机 `huggingface.co` 走代理会被 reset,用 `hf-mirror.com` 且**不设代理**:
```bash
curl -4 -L --retry 3 -o vocab.txt \
https://hf-mirror.com/OFA-Sys/chinese-clip-vit-base-patch16/resolve/main/vocab.txt
```
### 产物契约(Go 侧按此读取)
| 文件 | 输入 | 输出 |
|---|---|---|
| `TextEncoder.onnx` | `input_ids` int64 `[B,52]`、`attention_mask` int64 `[B,52]` | `text_features` float `[B,512]` |
| `VisionEncoder.onnx` | `pixel_values` float `[B,3,224,224]` | `image_features` float `[B,512]` |
外加 `embed_config.json`(维度/预处理/分词超参/文件名——provider 的唯一权威)、
`vocab.txt`、`reference.json`(冻结参考:逐文本 token id + 逐样本向量)、`SHA256SUMS`。
图像预处理:缩放到 224×224(双三次,复刻 PIL 系数)→ `(x/255 - mean) / std`,
不裁剪。文本:BERT WordPiece,`max_length=52`,补 `[PAD]`,超长截断尾部。
两个塔的输出**都没有在图中归一化**,归一化由 provider 负责(检索按余弦)。
### 启用
```bash
core.memory.multimodal_space.provider = chineseclip
core.memory.multimodal_space.options.model_dir = /home/newqqagent/models/chinese-clip-vit-b16-onnx
```
同样要求 `homed` 带 `onnxruntime` build tag。
### 模态范围
只声明 `text` 与 `image`。`audio`/`video` **明确返回 `ErrUnsupportedModality`**——
本空间没有它们的原生编码器,用别的模型向量冒充会污染整个向量空间
(这正是「音频明确 unsupported」那条纪律的落地)。
### 验证
Go 侧回归对着官方 PyTorch 参考(`reference.json`),模型目录由
`CHINESECLIP_MODEL_DIR` 指定,缺失时 skip:
```bash
CHINESECLIP_MODEL_DIR=/home/newqqagent/models/chinese-clip-vit-b16-onnx \
go test -tags onnxruntime ./providers/chineseclip/ -v
```
实测结果:文本 5 个用例 `cos = 1.000000000000`(与官方逐位一致);
图像 4 个纯色用例 `cos = 1.000000`(自写 bicubic 与 PIL 在 6 位小数内一致);
另有跨模态判别、模态拒绝、指纹稳定性、产物缺失报错等用例。
### 两个已踩过的坑(都在测试里钉住了)
1. **分词器不能自己拼**。第一版探针用 `BertTokenizer(vocab_file=..., do_lower_case=True)`
手工分词,中文被整体切成 `[UNK]`,三个不同句子产出几乎相同的向量(余弦 0.98),
差点把「模型坏了」当成结论。官方配置是 `do_lower_case=true` + **删音标生效** +
**中文逐字切分**;Go 侧实现必须与官方**逐 token** 对齐(`TestTokenizerMatchesOfficialReference`)。
2. **参考向量是未归一化的原始输出**(模长 10~36)。用「点积当余弦 + 单侧下界」判定
会得到 13.6 而「通过」——测试里因此改成真余弦 + 双侧容差。
## 二、启用
核心不识别任何具体模型:它只按配置里的 **provider 名**从公共注册表

3
go.mod
View File

@ -18,6 +18,7 @@ require (
github.com/charmbracelet/bubbletea v1.3.10
github.com/charmbracelet/lipgloss v1.1.0
golang.org/x/sys v0.38.0
golang.org/x/text v0.3.8
)
require (
@ -40,8 +41,6 @@ require (
github.com/muesli/termenv v0.16.0 // indirect
github.com/rivo/uniseg v0.4.7 // indirect
github.com/xo/terminfo v0.0.0-20220910002029-abceb7e1c41e // indirect
golang.org/x/text v0.3.8 // indirect
)
replace gitcode.com/JianFeeeee/homeagent-sdk => ./third_party/homeagent-sdk

View File

@ -0,0 +1,67 @@
package chineseclip
import (
"encoding/json"
"fmt"
"os"
"path/filepath"
)
// embedConfig 是产物目录里 embed_config.json 的映射:provider 的全部模型假设
// 都来自这个文件,不在代码里散落魔数。
type embedConfig struct {
Arch string `json:"arch"`
Dimension int `json:"dim"`
TextONNX string `json:"text_onnx"`
VisionONNX string `json:"vision_onnx"`
MaxLength int `json:"max_length"`
ImageSize int `json:"image_size"`
ImageMean []float64 `json:"image_mean"`
ImageStd []float64 `json:"image_std"`
Normalize bool `json:"normalize_vector"`
Modalities []string `json:"modalities"`
Unsupported []string `json:"unsupported_modalities"`
}
// loadConfig 读取并校验产物配置。任何不匹配都必须**明确报错**:
// 静默沿用默认值会在换错模型时产出「看起来正常、语义错误」的向量,
// 那类错误会污染整个图记忆且难以追查。
func loadConfig(dir string) (embedConfig, error) {
path := filepath.Join(dir, "embed_config.json")
data, err := os.ReadFile(path)
if err != nil {
return embedConfig{}, fmt.Errorf("chineseclip: 读取 %s: %w", path, err)
}
var cfg embedConfig
if err := json.Unmarshal(data, &cfg); err != nil {
return embedConfig{}, fmt.Errorf("chineseclip: 解析 %s: %w", path, err)
}
if cfg.Dimension != 512 {
return embedConfig{}, fmt.Errorf("chineseclip: 维度不匹配 dim=%d(期望 512)", cfg.Dimension)
}
if cfg.MaxLength <= 0 || cfg.MaxLength > 512 {
return embedConfig{}, fmt.Errorf("chineseclip: max_length 非法: %d", cfg.MaxLength)
}
if cfg.ImageSize != 224 {
return embedConfig{}, fmt.Errorf("chineseclip: image_size 不匹配 %d(期望 224)", cfg.ImageSize)
}
if len(cfg.ImageMean) != 3 || len(cfg.ImageStd) != 3 {
return embedConfig{}, fmt.Errorf("chineseclip: image_mean/std 必须各 3 个分量,得到 %d/%d",
len(cfg.ImageMean), len(cfg.ImageStd))
}
for i := range cfg.ImageStd {
if cfg.ImageStd[i] == 0 {
return embedConfig{}, fmt.Errorf("chineseclip: image_std[%d] 为 0", i)
}
}
if cfg.TextONNX == "" || cfg.VisionONNX == "" {
return embedConfig{}, fmt.Errorf("chineseclip: 未声明 onnx 文件名(text=%q vision=%q)",
cfg.TextONNX, cfg.VisionONNX)
}
for _, name := range []string{cfg.TextONNX, cfg.VisionONNX, "vocab.txt"} {
if _, err := os.Stat(filepath.Join(dir, name)); err != nil {
return embedConfig{}, fmt.Errorf("chineseclip: 产物缺少 %s: %w", name, err)
}
}
return cfg, nil
}

View File

@ -0,0 +1,282 @@
//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 ""
}

View File

@ -0,0 +1,273 @@
//go:build onnxruntime
package chineseclip
import (
"bytes"
"context"
"errors"
"image"
"image/color"
"image/png"
"math"
"os"
"testing"
"gitcode.com/JianFeeeee/HomeAgent/pkg/embedding"
)
// cosine 计算两个向量的**真余弦**:两边都先归一化。
//
// 参考向量存的是 ONNX 的原始输出(未归一化,模长 10~36),而 provider 的输出是
// L2 归一化后的。直接点积会得到参考向量的模长(例如 13.6),既不是余弦,
// 也会让单侧阈值判定变成假通过。
func cosine(a, b []float64) float64 {
if len(a) != len(b) || len(a) == 0 {
return math.NaN()
}
var dot, na, nb float64
for i := range a {
dot += a[i] * b[i]
na += a[i] * a[i]
nb += b[i] * b[i]
}
if na == 0 || nb == 0 {
return math.NaN()
}
return dot / (math.Sqrt(na) * math.Sqrt(nb))
}
// solidPNG 生成与导出脚本 FIXTURE_IMAGES 一致的纯色图(320×320)。
func solidPNG(t *testing.T, c color.RGBA) []byte {
t.Helper()
img := image.NewRGBA(image.Rect(0, 0, 320, 320))
for y := 0; y < 320; y++ {
for x := 0; x < 320; x++ {
img.Set(x, y, c)
}
}
var buf bytes.Buffer
if err := png.Encode(&buf, img); err != nil {
t.Fatalf("生成测试图: %v", err)
}
return buf.Bytes()
}
func openTestProvider(t *testing.T) embedding.Provider {
t.Helper()
dir := modelDir(t)
p, err := embedding.Open("chineseclip", embedding.Config{
Options: map[string]string{"model_dir": dir},
})
if err != nil {
t.Fatalf("embedding.Open(chineseclip): %v", err)
}
t.Cleanup(p.Close)
return p
}
// provider 必须通过公共 SPI 可见,并声明正确的向量空间身份。
func TestProviderInfo(t *testing.T) {
p := openTestProvider(t)
info := p.Info()
if info.Dimension != 512 {
t.Errorf("维度应为 512,得到 %d", info.Dimension)
}
if len(info.Fingerprint) != 64 {
t.Errorf("指纹应为 64 位十六进制,得到 %q", info.Fingerprint)
}
want := map[embedding.Modality]bool{embedding.ModalityText: true, embedding.ModalityImage: true}
if len(info.Modalities) != len(want) {
t.Fatalf("模态应为 %v,得到 %v", want, info.Modalities)
}
for _, m := range info.Modalities {
if !want[m] {
t.Errorf("声明了未支持的模态 %q", m)
}
}
}
// 文本向量必须与官方 PyTorch 参考一致。
//
// 文本侧没有预处理歧义(分词器已逐 token 对齐),所以要求非常严:
// 余弦与参考的偏差应小于 1e-9。
func TestEmbedTextMatchesReference(t *testing.T) {
dir := modelDir(t)
p := openTestProvider(t)
ref := loadReference(t, dir)
if len(ref.Texts) == 0 {
t.Fatal("参考里没有文本用例")
}
ctx := context.Background()
for _, c := range ref.Texts {
got, err := p.Embed(ctx, embedding.Input{
Modality: embedding.ModalityText,
Purpose: embedding.PurposeQuery,
Text: c.Text,
})
if err != nil {
t.Fatalf("Embed(text=%q): %v", c.Text, err)
}
if err := embedding.ValidateVector(got, 512); err != nil {
t.Fatalf("向量不合法 %q: %v", c.Text, err)
}
cos := cosine(got, c.Vector)
if math.Abs(cos-1.0) > 1e-9 {
t.Errorf("文本向量与参考不一致 %q: cos=%.12f", c.Text, cos)
}
t.Logf("文本 cos=%.12f %s", cos, c.Text[:min(len(c.Text), 24)])
}
}
// 图像向量与官方参考一致(容忍缩放实现差异)。
//
// Go 侧自写 bicubic(PIL 系数)与官方预处理不会逐位相同,故用余弦阈值;
// 0.999 足以证明「同一条管线」,同时不会掩盖把像素顺序或归一化写错这类错误
// (那类错误会直接掉到 0.9 以下)。
func TestEmbedImageMatchesReference(t *testing.T) {
dir := modelDir(t)
p := openTestProvider(t)
ref := loadReference(t, dir)
ctx := context.Background()
colors := map[string]color.RGBA{
"red": {R: 220, G: 30, B: 30, A: 255},
"green": {R: 60, G: 120, B: 60, A: 255},
"blue": {R: 30, G: 30, B: 220, A: 255},
"gray": {R: 128, G: 128, B: 128, A: 255},
}
checked := 0
for _, c := range ref.Images {
rgba, ok := colors[c.Name]
if !ok {
continue // gradient 在 Go 侧不便逐位复刻,跳过(仍由导出脚本覆盖)
}
got, err := p.Embed(ctx, embedding.Input{
Modality: embedding.ModalityImage,
Purpose: embedding.PurposeDocument,
Data: solidPNG(t, rgba),
MIME: "image/png",
})
if err != nil {
t.Fatalf("Embed(image=%s): %v", c.Name, err)
}
cos := cosine(got, c.Vector)
// 双侧判定:单侧下界挡不住「模长缩放」这类错误(未归一化的参考向量
// 会让点积恰好远大于 1 而“通过”)。
if math.Abs(cos-1.0) > 1e-3 {
t.Errorf("图像向量与参考不一致 %s: cos=%.6f", c.Name, cos)
}
t.Logf("图像 cos=%.6f %s", cos, c.Name)
checked++
}
if checked == 0 {
t.Fatal("没有比对任何图像用例(参考里缺少纯色样例)")
}
}
// 跨模态必须真的有区分度:红图对"红色"文本应高于"蓝色"文本。
// 这条防的是「向量塌缩但 cos 检查全过」那类假通过。
func TestCrossModalDiscrimination(t *testing.T) {
dir := modelDir(t)
p := openTestProvider(t)
ref := loadReference(t, dir)
ctx := context.Background()
var redText, blueText string
for _, c := range ref.Texts {
switch c.Text {
case "一张红色方块的图片":
redText = c.Text
case "蓝色的天空":
blueText = c.Text
}
}
if redText == "" || blueText == "" {
t.Skip("参考里缺少用于跨模态判别的文本")
}
imgVec, err := p.Embed(ctx, embedding.Input{
Modality: embedding.ModalityImage, Data: solidPNG(t, color.RGBA{R: 220, G: 30, B: 30, A: 255}),
MIME: "image/png",
})
if err != nil {
t.Fatalf("Embed(image): %v", err)
}
redVec, err := p.Embed(ctx, embedding.Input{Modality: embedding.ModalityText, Text: redText})
if err != nil {
t.Fatalf("Embed(red text): %v", err)
}
blueVec, err := p.Embed(ctx, embedding.Input{Modality: embedding.ModalityText, Text: blueText})
if err != nil {
t.Fatalf("Embed(blue text): %v", err)
}
if cosine(imgVec, redVec) <= cosine(imgVec, blueVec) {
t.Errorf("跨模态判别失败:红图-红文本 %.4f 应高于 红图-蓝文本 %.4f",
cosine(imgVec, redVec), cosine(imgVec, blueVec))
}
}
// 音频/视频必须明确拒绝,绝不用别的模型向量冒充。
func TestEmbedRejectsUnsupportedModalities(t *testing.T) {
p := openTestProvider(t)
ctx := context.Background()
for _, m := range []embedding.Modality{embedding.ModalityAudio, embedding.ModalityVideo} {
_, err := p.Embed(ctx, embedding.Input{Modality: m, Data: []byte("x"), MIME: "application/octet-stream"})
if err == nil {
t.Fatalf("模态 %s 应被拒绝", m)
}
if !errors.Is(err, embedding.ErrUnsupportedModality) {
t.Errorf("模态 %s 的错误应可判定为 ErrUnsupportedModality,得到: %v", m, err)
}
}
}
// 空输入必须报错而不是产出垃圾向量。
func TestEmbedRejectsEmptyInput(t *testing.T) {
p := openTestProvider(t)
ctx := context.Background()
if _, err := p.Embed(ctx, embedding.Input{Modality: embedding.ModalityText, Text: " "}); err == nil {
t.Error("空文本应报错")
}
if _, err := p.Embed(ctx, embedding.Input{Modality: embedding.ModalityImage}); err == nil {
t.Error("空图像应报错")
}
}
// 指纹必须稳定且对产物内容敏感(同目录两次打开一致)。
func TestFingerprintStable(t *testing.T) {
dir := modelDir(t)
first, err := embedding.Open("chineseclip", embedding.Config{Options: map[string]string{"model_dir": dir}})
if err != nil {
t.Fatalf("首次打开: %v", err)
}
defer first.Close()
second, err := embedding.Open("chineseclip", embedding.Config{Options: map[string]string{"model_dir": dir}})
if err != nil {
t.Fatalf("二次打开: %v", err)
}
defer second.Close()
if first.Info().Fingerprint != second.Info().Fingerprint {
t.Errorf("同一产物两次打开指纹不一致: %s vs %s",
first.Info().Fingerprint, second.Info().Fingerprint)
}
}
// 缺 model_dir 必须明确报错(便于区分「没配置」与「模型坏了」)。
func TestOpenRejectsMissingModelDir(t *testing.T) {
if _, err := embedding.Open("chineseclip", embedding.Config{}); err == nil {
t.Fatal("缺 model_dir 时应打开失败")
}
}
// 目录存在但不是本模型产物时,必须报出缺哪个文件,而不是静默用默认值。
func TestOpenRejectsIncompleteArtifacts(t *testing.T) {
dir := t.TempDir()
if err := os.WriteFile(dir+"/embed_config.json", []byte(`{"dim":512,"max_length":52,"image_size":224,"image_mean":[0.5,0.5,0.5],"image_std":[0.5,0.5,0.5],"text_onnx":"TextEncoder.onnx","vision_onnx":"VisionEncoder.onnx"}`), 0o644); err != nil {
t.Fatal(err)
}
_, err := embedding.Open("chineseclip", embedding.Config{Options: map[string]string{"model_dir": dir}})
if err == nil {
t.Fatal("产物不完整时应打开失败")
}
}

View File

@ -0,0 +1,33 @@
//go:build !onnxruntime
package chineseclip
import (
"context"
"errors"
"fmt"
"gitcode.com/JianFeeeee/HomeAgent/pkg/embedding"
)
// 未启用 onnxruntime 构建标签时,chineseclip 仍注册到名字表,但打开即报错:
// 这样「provider 名写错」与「本次构建没带 ONNX」是两种可区分的失败,
// 而不是一句含糊的 unknown provider。
func init() {
embedding.Register("chineseclip", func(embedding.Config) (embedding.Provider, error) {
return nil, fmt.Errorf("chineseclip provider requires build tag 'onnxruntime' " +
"(go build -tags onnxruntime)")
})
}
// Embedder 在未启用 onnxruntime 时不可用;保留类型是为了让引用它的代码在
// 默认构建下也能编译。真正的 ONNX 实现见 embedder.go。
type Embedder struct{}
func (e *Embedder) Embed(context.Context, embedding.Input) ([]float64, error) {
return nil, errors.New("chineseclip provider not available in this build")
}
func (e *Embedder) Info() embedding.Info { return embedding.Info{} }
func (e *Embedder) Close() {}

View File

@ -0,0 +1,171 @@
package chineseclip
import (
"bytes"
"fmt"
"image"
// 契约要求 provider 自行解码 Data,所以这里注册常见图像格式。
_ "image/gif"
_ "image/jpeg"
_ "image/png"
)
// plane 是单通道浮点平面。
type plane struct {
w, h int
data []float32
}
// preprocessImage 把原始图像字节变成 ONNX 需要的 NCHW 张量:
// 缩放到 size×size(双三次,复刻 PIL 的系数)→ 归一化(x/255 - mean)/ std。
//
// 缩放在 RGB 三个通道上分别进行,与官方 ChineseCLIPFeatureExtractor 一致
// (do_resize=true、do_center_crop=false、resample=BICUBIC、rescale 1/255)。
func preprocessImage(data []byte, size int, mean, std []float64) ([]float32, error) {
img, _, err := image.Decode(bytes.NewReader(data))
if err != nil {
return nil, fmt.Errorf("chineseclip: 解码图像: %w", err)
}
bounds := img.Bounds()
if bounds.Dx() <= 0 || bounds.Dy() <= 0 {
return nil, fmt.Errorf("chineseclip: 图像尺寸非法 %dx%d", bounds.Dx(), bounds.Dy())
}
planes := [3]plane{}
for c := range planes {
planes[c] = plane{w: bounds.Dx(), h: bounds.Dy(), data: make([]float32, bounds.Dx()*bounds.Dy())}
}
for y := bounds.Min.Y; y < bounds.Max.Y; y++ {
for x := bounds.Min.X; x < bounds.Max.X; x++ {
r, g, b, _ := img.At(x, y).RGBA()
idx := (y-bounds.Min.Y)*bounds.Dx() + (x - bounds.Min.X)
// RGBA() 返回的是 16 位预乘值;不透明图像下右移 8 位即得 8 位分量。
planes[0].data[idx] = float32(r >> 8)
planes[1].data[idx] = float32(g >> 8)
planes[2].data[idx] = float32(b >> 8)
}
}
out := make([]float32, 3*size*size)
for c := range planes {
resized := resizeBicubic(planes[c], size, size)
for i, v := range resized.data {
scaled := float64(v) / 255.0
out[c*size*size+i] = float32((scaled - mean[c]) / std[c])
}
}
return out, nil
}
// resizeBicubic 复刻 PIL 的可分离双三次缩放(系数来自 PIL 的
// precompute_coeffs + bicubic_filter,a=-0.5)。
//
// 为什么不引第三方 resize:本机 x/image 未进模块缓存,而它最新版还要求把整个
// 工具链升到 Go 1.26;为一个缩放函数动工具链不划算。PIL 的算法只有几十行,
// 照抄系数能保证与官方预处理足够接近(已用端到端 cos 验证)。
func resizeBicubic(src plane, dstW, dstH int) plane {
if src.w == dstW && src.h == dstH {
return src
}
wsX := buildWeights(src.w, dstW)
tmp := plane{w: dstW, h: src.h, data: make([]float32, dstW*src.h)}
for y := 0; y < src.h; y++ {
row := y * src.w
for dx := 0; dx < dstW; dx++ {
var sum float32
for _, t := range wsX[dx] {
sum += t.w * src.data[row+t.i]
}
tmp.data[y*dstW+dx] = sum
}
}
wsY := buildWeights(src.h, dstH)
dst := plane{w: dstW, h: dstH, data: make([]float32, dstW*dstH)}
for dy := 0; dy < dstH; dy++ {
for x := 0; x < dstW; x++ {
var sum float32
for _, t := range wsY[dy] {
sum += t.w * tmp.data[t.i*dstW+x]
}
dst.data[dy*dstW+x] = sum
}
}
return dst
}
type weightTerm struct {
i int
w float32
}
// buildWeights 按 PIL 的 precompute_coeffs 计算每个目标像素的源像素权重。
func buildWeights(srcLen, dstLen int) [][]weightTerm {
const support = 2.0 // BICUBIC 的支撑半径
filterScale := float64(srcLen) / float64(dstLen)
if filterScale < 1.0 {
filterScale = 1.0
}
scale := filterScale
filterSupport := support * filterScale
invScale := 1.0 / filterScale
out := make([][]weightTerm, dstLen)
for d := 0; d < dstLen; d++ {
center := (float64(d) + 0.5) * scale
xmin := int(center - filterSupport + 0.5)
if xmin < 0 {
xmin = 0
}
xmax := int(center + filterSupport + 0.5)
if xmax > srcLen {
xmax = srcLen
}
if xmax <= xmin {
// 极端缩放下的兜底:退化为最近邻,避免空权重导致除零。
idx := int(center)
if idx < 0 {
idx = 0
}
if idx >= srcLen {
idx = srcLen - 1
}
out[d] = []weightTerm{{i: idx, w: 1}}
continue
}
terms := make([]weightTerm, 0, xmax-xmin)
var total float64
for x := xmin; x < xmax; x++ {
w := bicubicKernel((float64(x) - center + 0.5) * invScale)
if w == 0 {
continue
}
terms = append(terms, weightTerm{i: x, w: float32(w)})
total += w
}
if total != 0 {
for i := range terms {
terms[i].w = float32(float64(terms[i].w) / total)
}
}
out[d] = terms
}
return out
}
// bicubicKernel 是 PIL 的 bicubic_filter(a = -0.5)。
func bicubicKernel(x float64) float64 {
const a = -0.5
if x < 0 {
x = -x
}
switch {
case x < 1.0:
return ((a+2.0)*x-(a+3.0))*x*x + 1.0
case x < 2.0:
return (((x-5.0)*x+8.0)*x - 4.0) * a
}
return 0
}

View File

@ -0,0 +1,291 @@
// Package chineseclip 提供 Chinese-CLIP ViT-B/16 的 text+image 向量空间 provider。
//
// 为什么是它(而不是 Qwen3-VL-Embedding-2B / jina-v5-omni-nano):
// - 体积:721MB ONNX、实测常驻 1.15GB;Qwen 2B 需要 9.4GB,本机可用内存只有 5.3GB。
// - 许可:Apache-2.0,可随发行版分发;jina-v5-omni-nano 是 CC BY-NC(不可商用)。
// - 中文:原生在 ~2 亿中文图文对上训练。
//
// 代价(明确记录):CLIP 是双塔对比学习,text↔image 是强项,但纯文本语义
// (text↔text)明显弱于 MLLM 型嵌入器。文本检索仍由既有词向量/TF-IDF 路径兜底,
// 本空间主要用于跨模态召回与相关性裁剪。需要视频或更强文本语义时应切回
// providers/qwen3vl(内存允许时)。
//
// 模态范围:仅 text 与 image。audio / video 返回 embedding.ErrUnsupportedModality,
// 绝不用别的模型向量冒充。
package chineseclip
import (
"fmt"
"os"
"path/filepath"
"strings"
"unicode"
"golang.org/x/text/unicode/norm"
)
// BERT 的固定特殊 token(与官方 Chinese-CLIP 的 vocab.txt 一致)。
const (
tokenCLS = "[CLS]"
tokenSEP = "[SEP]"
tokenPAD = "[PAD]"
tokenUNK = "[UNK]"
// maxInputCharsPerWord 与 HF BertTokenizer 一致:超过就整词判 UNK。
maxInputCharsPerWord = 100
)
// Tokenizer 是 BERT WordPiece 分词器(Chinese-CLIP 官方配置:do_lower_case=true、
// strip_accents 生效、tokenize_chinese_chars=true)。
type Tokenizer struct {
vocab map[string]int32
maxLength int
}
// LoadTokenizer 从模型目录读取 vocab.txt。目录里那份词表是产物的组成部分,
// provider 只依赖这个目录,不去猜任何外部路径。
func LoadTokenizer(dir string, maxLength int) (*Tokenizer, error) {
path := filepath.Join(dir, "vocab.txt")
data, err := os.ReadFile(path)
if err != nil {
return nil, fmt.Errorf("chineseclip: 读取词表 %s: %w", path, err)
}
if maxLength <= 0 {
return nil, fmt.Errorf("chineseclip: max_length 必须为正,得到 %d", maxLength)
}
vocab := make(map[string]int32, 32768)
for i, line := range strings.Split(string(data), "\n") {
piece := strings.TrimRight(line, "\r")
if piece == "" {
continue
}
if _, dup := vocab[piece]; dup {
// 词表出现重复行说明文件被写坏;不静默用后者覆盖前者。
return nil, fmt.Errorf("chineseclip: 词表第 %d 行重复: %q", i+1, piece)
}
vocab[piece] = int32(len(vocab))
}
for _, special := range []string{tokenCLS, tokenSEP, tokenPAD, tokenUNK} {
if _, ok := vocab[special]; !ok {
return nil, fmt.Errorf("chineseclip: 词表缺少特殊 token %s", special)
}
}
return &Tokenizer{vocab: vocab, maxLength: maxLength}, nil
}
// MaxLength 返回文本侧的最大 token 数(含特殊 token)。
func (t *Tokenizer) MaxLength() int { return t.maxLength }
// Encode 返回补齐到 maxLength 的 input_ids 与 attention_mask。
// attention_mask 与官方 tokenizer 的 padding='max_length' 行为一致:真实 token 为 1,
// padding 为 0。
func (t *Tokenizer) Encode(text string) ([]int64, []int64) {
pieces := t.tokenize(text)
// 预留 [CLS] 与 [SEP];超长直接截断尾部(官方 truncation=True 的默认方向)。
if limit := t.maxLength - 2; len(pieces) > limit {
pieces = pieces[:limit]
}
ids := make([]int64, 0, t.maxLength)
mask := make([]int64, 0, t.maxLength)
ids = append(ids, int64(t.vocab[tokenCLS]))
mask = append(mask, 1)
for _, p := range pieces {
ids = append(ids, int64(t.vocab[p]))
mask = append(mask, 1)
}
ids = append(ids, int64(t.vocab[tokenSEP]))
mask = append(mask, 1)
for len(ids) < t.maxLength {
ids = append(ids, int64(t.vocab[tokenPAD]))
mask = append(mask, 0)
}
return ids, mask
}
// tokenize 复刻 HF BasicTokenizer + WordPieceTokenizer 的完整流水线。
func (t *Tokenizer) tokenize(text string) []string {
var pieces []string
for _, basic := range basicTokenize(text) {
pieces = append(pieces, t.wordpiece(basic)...)
}
return pieces
}
// basicTokenize 实现 BasicTokenizer(空模型版):清洗 → 中文逐字加空格 →
// 按空白切分 → 删音标 + 转小写 → 按标点再次切分。
func basicTokenize(text string) []string {
cleaned := cleanText(text)
var out []string
for _, token := range strings.Fields(tokenizeChineseChars(cleaned)) {
if len([]rune(token)) > maxInputCharsPerWord {
// 与 HF 一致:超长基本 token 直接丢弃(后续不会产出 UNK)。
continue
}
stripped := stripAccents(strings.ToLower(token))
out = append(out, splitOnPunctuation(stripped)...)
}
return out
}
// cleanText 与 HF _clean_text 一致:丢弃 NUL/替换符与控制符,空白统一为空格。
func cleanText(text string) string {
var b strings.Builder
b.Grow(len(text))
for _, r := range text {
switch {
case r == 0 || r == 0xFFFD:
continue
case isControl(r):
continue
case isBERTWhitespace(r):
b.WriteRune(' ')
default:
b.WriteRune(r)
}
}
return b.String()
}
// tokenizeChineseChars 在 CJK 字符两侧插入空格,使每个汉字成为独立基本 token。
func tokenizeChineseChars(text string) string {
var b strings.Builder
b.Grow(len(text) + 16)
for _, r := range text {
if isCJK(r) {
b.WriteRune(' ')
b.WriteRune(r)
b.WriteRune(' ')
continue
}
b.WriteRune(r)
}
return b.String()
}
// stripAccents 与 HF _run_strip_accents 一致:NFD 分解后丢弃 Mn 组合记号
// ("café" → "cafe")。
func stripAccents(text string) string {
if isASCII(text) {
return text
}
var b strings.Builder
b.Grow(len(text))
for _, r := range norm.NFD.String(text) {
if unicode.Is(unicode.Mn, r) {
continue
}
b.WriteRune(r)
}
return b.String()
}
// splitOnPunctuation 与 HF _run_split_on_punc 一致:标点自成一段。
//
// 注意 ASCII 段必须显式列出:'$' '+' '=' '^' '`' '|' '~' 属于 Sc/Sm/Sk,
// 不是 Unicode P*,但它们也是标点(HF 用的是 ASCII 码点区间)。
func splitOnPunctuation(text string) []string {
runes := []rune(text)
var out []string
var cur []rune
flush := func() {
if len(cur) > 0 {
out = append(out, string(cur))
cur = cur[:0]
}
}
for _, r := range runes {
if isBERTPunctuation(r) {
flush()
out = append(out, string(r))
continue
}
cur = append(cur, r)
}
flush()
return out
}
// wordpiece 贪心最长匹配;整词任一段无法匹配则该词整体退化为 [UNK]。
func (t *Tokenizer) wordpiece(token string) []string {
runes := []rune(token)
if len(runes) > maxInputCharsPerWord {
return []string{tokenUNK}
}
var out []string
start := 0
for start < len(runes) {
end := len(runes)
var cur string
found := false
for end > start {
piece := string(runes[start:end])
if start > 0 {
piece = "##" + piece
}
if _, ok := t.vocab[piece]; ok {
cur = piece
found = true
break
}
end--
}
if !found {
return []string{tokenUNK}
}
out = append(out, cur)
start = end
}
return out
}
func isASCII(s string) bool {
for i := 0; i < len(s); i++ {
if s[i] >= 0x80 {
return false
}
}
return true
}
// isBERTWhitespace:HF _is_whitespace = 空格/制表/换行/回车 或 Unicode Zs。
func isBERTWhitespace(r rune) bool {
if r == ' ' || r == '\t' || r == '\n' || r == '\r' {
return true
}
return unicode.Is(unicode.Zs, r)
}
// isControl:HF _is_control = Cc/Cf,但制表/换行/回车不算。
func isControl(r rune) bool {
if r == '\t' || r == '\n' || r == '\r' {
return false
}
return unicode.Is(unicode.Cc, r) || unicode.Is(unicode.Cf, r)
}
// isBERTPunctuation:ASCII 标点区间 或 Unicode P*。
func isBERTPunctuation(r rune) bool {
if (r >= 33 && r <= 47) || (r >= 58 && r <= 64) || (r >= 91 && r <= 96) || (r >= 123 && r <= 126) {
return true
}
return unicode.IsPunct(r)
}
// isCJK:HF _tokenize_chinese_chars 使用的区间表。
func isCJK(r rune) bool {
switch {
case r >= 0x4E00 && r <= 0x9FFF,
r >= 0x3400 && r <= 0x4DBF,
r >= 0x20000 && r <= 0x2A6DF,
r >= 0x2A700 && r <= 0x2B73F,
r >= 0x2B740 && r <= 0x2B81F,
r >= 0x2B820 && r <= 0x2CEAF,
r >= 0xF900 && r <= 0xFAFF,
r >= 0x2F800 && r <= 0x2FA1F:
return true
}
return false
}

View File

@ -0,0 +1,165 @@
package chineseclip
import (
"encoding/json"
"os"
"path/filepath"
"reflect"
"testing"
)
// referenceText 是官方导出时冻结的逐文本 token 参考。
type referenceText struct {
Text string `json:"text"`
InputIDs []int64 `json:"input_ids"`
Attention []int64 `json:"attention_mask"`
Vector []float64 `json:"vector"`
PixelsSHA256 string `json:"pixels_sha256"`
Name string `json:"name"`
}
type reference struct {
Texts []referenceText `json:"texts"`
Images []referenceText `json:"images"`
}
func loadReference(t *testing.T, dir string) reference {
t.Helper()
path := filepath.Join(dir, "reference.json")
data, err := os.ReadFile(path)
if err != nil {
t.Skipf("缺少冻结参考 %s(由 scripts/export_chineseclip_onnx.py 生成): %v", path, err)
}
var ref reference
if err := json.Unmarshal(data, &ref); err != nil {
t.Fatalf("解析参考 %s: %v", path, err)
}
return ref
}
func modelDir(t *testing.T) string {
t.Helper()
dir := os.Getenv("CHINESECLIP_MODEL_DIR")
if dir == "" {
t.Skip("未设置 CHINESECLIP_MODEL_DIR,跳过需要真实产物的用例")
}
if _, err := os.Stat(filepath.Join(dir, "embed_config.json")); err != nil {
t.Skipf("模型目录 %s 缺 embed_config.json: %v", dir, err)
}
return dir
}
// 分词器必须与官方 Chinese-CLIP 逐 token 一致。
//
// 这条测试是有来历的:第一版探针自己拼 BertTokenizer(只给 vocab.txt、没删音标、
// 中文没逐字切),中文全被切成 [UNK],三个不同句子产出几乎相同的向量
// (余弦 0.98)——差点把「模型坏了」当成结论。分词不一致会静默毁掉整个向量空间。
func TestTokenizerMatchesOfficialReference(t *testing.T) {
dir := modelDir(t)
cfg, err := loadConfig(dir)
if err != nil {
t.Fatalf("读取 embed_config.json: %v", err)
}
tok, err := LoadTokenizer(dir, cfg.MaxLength)
if err != nil {
t.Fatalf("加载分词器: %v", err)
}
ref := loadReference(t, dir)
if len(ref.Texts) == 0 {
t.Fatal("参考里没有文本用例")
}
for _, c := range ref.Texts {
ids, mask := tok.Encode(c.Text)
if !reflect.DeepEqual(ids, c.InputIDs) {
t.Errorf("input_ids 不一致 %q\n got %v\n want %v", c.Text, ids, c.InputIDs)
}
if !reflect.DeepEqual(mask, c.Attention) {
t.Errorf("attention_mask 不一致 %q\n got %v\n want %v", c.Text, mask, c.Attention)
}
}
}
// 不依赖真实产物的纯逻辑用例:覆盖中文逐字、删音标、标点切分、UNK、补齐。
func TestTokenizerUnitCases(t *testing.T) {
dir := t.TempDir()
vocab := []string{
"[PAD]", "[UNK]", "[CLS]", "[SEP]", "[MASK]",
"红", "色", "一", "张", "方", "块", "的", "图", "片",
"cafe", "hello", "world", "##ive", "na", "a", "-", "b", "(", ")", "括", "号", "12345",
}
if err := os.WriteFile(filepath.Join(dir, "vocab.txt"),
[]byte(joinLines(vocab)), 0o644); err != nil {
t.Fatal(err)
}
tok, err := LoadTokenizer(dir, 8)
if err != nil {
t.Fatalf("加载分词器: %v", err)
}
cases := []struct {
name string
text string
want []string // 期望的 token 文本(不含特殊 token,便于阅读)
}{
{"中文逐字", "红色", []string{"红", "色"}},
{"删音标+小写", "CAFÉ", []string{"cafe"}},
{"删音标词内组合", "naïve", []string{"na", "##ive"}},
{"标点切分", "a-b", []string{"a", "-", "b"}},
{"全角括号按标点处理", "(括号)", []string{"(", "括", "号", ")"}},
{"纯数字", "12345", []string{"12345"}},
}
for _, c := range cases {
got := tok.tokenize(c.text)
if !reflect.DeepEqual(got, c.want) {
t.Errorf("%s: %q\n got %v\n want %v", c.name, c.text, got, c.want)
}
}
// 未登录词整体退化为 [UNK](与 HF 一致)。
if got := tok.tokenize("zzz"); !reflect.DeepEqual(got, []string{"[UNK]"}) {
t.Errorf("未登录词应退化为 [UNK],得到 %v", got)
}
// 补齐:maxLength=8,"红色" 只占 2 个位置,其余补 [PAD]。
ids, mask := tok.Encode("红色")
if len(ids) != 8 || len(mask) != 8 {
t.Fatalf("补齐长度应为 8,得到 ids=%d mask=%d", len(ids), len(mask))
}
if ids[0] != 2 || ids[1] != 5 || ids[2] != 6 || ids[3] != 3 {
t.Errorf("应为 [CLS] 红 色 [SEP],得到 %v", ids[:4])
}
for i := 4; i < 8; i++ {
if ids[i] != 0 {
t.Errorf("位置 %d 应为 [PAD],得到 %v", i, ids)
}
}
wantMask := []int64{1, 1, 1, 1, 0, 0, 0, 0}
if !reflect.DeepEqual(mask, wantMask) {
t.Errorf("attention_mask 应为 %v,得到 %v", wantMask, mask)
}
// 截断:9 个汉字在 maxLength=8 下只保留 6 个,正好填满,不应出现 [PAD]。
truncIDs, truncMask := tok.Encode("红色一张方块的图片")
if len(truncIDs) != 8 {
t.Fatalf("截断后长度应为 8,得到 %d", len(truncIDs))
}
if truncIDs[0] != 2 || truncIDs[7] != 3 {
t.Errorf("截断后首尾应为 [CLS]/[SEP],得到 %v", truncIDs)
}
for i, id := range truncIDs {
if id == 0 {
t.Errorf("截断后位置 %d 不应是 [PAD]: %v", i, truncIDs)
}
}
if !reflect.DeepEqual(truncMask, []int64{1, 1, 1, 1, 1, 1, 1, 1}) {
t.Errorf("截断后 attention_mask 应全为 1,得到 %v", truncMask)
}
}
func joinLines(items []string) string {
out := ""
for _, item := range items {
out += item + "\n"
}
return out
}

View File

@ -0,0 +1,371 @@
#!/usr/bin/env python3
"""把 Chinese-CLIP ViT-B/16 导出成 HomeAgent 的规范 ONNX 产物。
为什么是 Chinese-CLIP:
text+image 的默认向量空间要同时满足「小、可商用、中文原生」。
Chinese-CLIP ViT-B/16 = 188M 参数 / 721MB ONNX / 实测常驻 1.15GB,
许可是 Apache-2.0(可随发行版分发),且原生在 2 亿中文图文对上训练。
对比:jina-v5-omni-nano 2.23GB 但 CC BY-NC(不可商用);Qwen3-VL-Emb-2B
9.4GB(质量最好,保留为可选 provider)。
产物(--out 目录,会被清空重建):
TextEncoder.onnx input_ids[·,52] + attention_mask[·,52] → text_features[·,512]
VisionEncoder.onnx pixel_values[·,3,224,224] → image_features[·,512]
embed_config.json 维度/预处理/分词超参/文件名(provider 侧的唯一权威)
vocab.txt 分词器词表(来自官方模型目录)
reference.json 冻结参考:逐文本 token id + 逐样本参考向量(Go 侧回归用)
SHA256SUMS
自检纪律(对齐 export_qwen3vl_embedding_onnx.py):
1. 每个用例都真跑一次 ONNX 并与 PyTorch 对比,打印逐用例 cos;
2. 计划用例集合与实际执行集合必须相等,否则非零退出(防「先跳过再校验」的假通过);
3. 不采信退出码,失败一律非零退出并说明原因。
用法:
scripts/export_chineseclip_onnx.py --out DIR [--model-dir DIR|--model-id REPO]
[--verify-only] [--no-reference] [--skip-verify]
"""
from __future__ import annotations
import argparse
import hashlib
import json
import os
import shutil
import sys
import time
import numpy as np
DEFAULT_MODEL_ID = "OFA-Sys/chinese-clip-vit-base-patch16"
TEXT_FILE = "TextEncoder.onnx"
VISION_FILE = "VisionEncoder.onnx"
CONFIG_FILE = "embed_config.json"
REFERENCE_FILE = "reference.json"
OPSET = 17
MAX_LENGTH = 52
IMAGE_SIZE = 224
# 冻结用例:文本覆盖纯中文/中英混/长文本/标点,图像覆盖纯色与渐变。
FIXTURE_TEXTS = [
"一张红色方块的图片",
"一只猫在草地上",
"蓝色的天空",
"HomeAgent 是一个本地 AI 管家",
"这是一段比较长的中文文本,用来验证分词器在超过五十个 token 时的截断行为是否正确,"
"同时检查标点符号、数字 12345 和英文单词 embedding 的处理。",
]
FIXTURE_IMAGES = [
("red", (220, 30, 30)),
("green", (60, 120, 60)),
("blue", (30, 30, 220)),
("gray", (128, 128, 128)),
]
def solid(color: tuple[int, int, int]) -> np.ndarray:
from PIL import Image
img = Image.new("RGB", (320, 320), color)
return np.asarray(img, dtype=np.uint8)
def gradient() -> np.ndarray:
"""确定性的横向渐变,避免只有纯色导致区分度不足。"""
row = np.linspace(0, 255, 320, dtype=np.uint8)
img = np.zeros((320, 320, 3), dtype=np.uint8)
img[:, :, 0] = row[None, :]
img[:, :, 1] = row[:, None]
img[:, :, 2] = 64
return img
def norm_cos(a: np.ndarray, b: np.ndarray) -> float:
a = a.reshape(-1).astype(np.float64)
b = b.reshape(-1).astype(np.float64)
na, nb = np.linalg.norm(a), np.linalg.norm(b)
if na == 0 or nb == 0:
return 0.0
return float(np.dot(a, b) / (na * nb))
def sha256_file(path: str) -> str:
h = hashlib.sha256()
with open(path, "rb") as f:
for chunk in iter(lambda: f.read(1 << 20), b""):
h.update(chunk)
return h.hexdigest()
def load_model(model_dir: str | None, model_id: str, local_only: bool):
import torch
from transformers import ChineseCLIPModel, ChineseCLIPProcessor
src = model_dir or model_id
kw = {"local_files_only": True} if local_only else {}
print(f"[load] {src}")
t0 = time.time()
processor = ChineseCLIPProcessor.from_pretrained(src, **kw)
model = ChineseCLIPModel.from_pretrained(src, **kw).eval()
print(f"[load] 用时 {time.time() - t0:.1f}s")
return model, processor
def tower_forward_text(model, input_ids, attention_mask):
import torch
with torch.inference_mode():
out = model.text_model(input_ids=input_ids, attention_mask=attention_mask)
pooled = out.pooler_output
if pooled is None:
pooled = out.last_hidden_state[:, 0]
return model.text_projection(pooled)
def tower_forward_vision(model, pixel_values):
import torch
with torch.inference_mode():
out = model.vision_model(pixel_values=pixel_values)
pooled = out.pooler_output
if pooled is None:
pooled = out.last_hidden_state[:, 0]
return model.visual_projection(pooled)
class TextTowerWrapper:
"""torch.onnx.export 需要 nn.Module,这里在函数内构造以避免顶层 import torch。"""
def make_wrappers(model):
import torch
class TextTower(torch.nn.Module):
def __init__(self, m):
super().__init__()
self.m = m
def forward(self, input_ids, attention_mask):
out = self.m.text_model(input_ids=input_ids, attention_mask=attention_mask)
pooled = out.pooler_output
if pooled is None:
pooled = out.last_hidden_state[:, 0]
return self.m.text_projection(pooled)
class VisionTower(torch.nn.Module):
def __init__(self, m):
super().__init__()
self.m = m
def forward(self, pixel_values):
out = self.m.vision_model(pixel_values=pixel_values)
pooled = out.pooler_output
if pooled is None:
pooled = out.last_hidden_state[:, 0]
return self.m.visual_projection(pooled)
return TextTower(model).eval(), VisionTower(model).eval()
def export_onnx(model, processor, out_dir: str) -> None:
import torch
text_tower, vision_tower = make_wrappers(model)
tok = processor.tokenizer
enc = tok(["占位"], padding="max_length", truncation=True,
max_length=MAX_LENGTH, return_tensors="pt")
pixel = torch.zeros(1, 3, IMAGE_SIZE, IMAGE_SIZE, dtype=torch.float32)
print(f"[export] {TEXT_FILE}")
torch.onnx.export(
text_tower,
(enc["input_ids"], enc["attention_mask"]),
os.path.join(out_dir, TEXT_FILE),
input_names=["input_ids", "attention_mask"],
output_names=["text_features"],
dynamic_axes={"input_ids": {0: "batch"}, "attention_mask": {0: "batch"},
"text_features": {0: "batch"}},
opset_version=OPSET,
do_constant_folding=True,
dynamo=False,
)
print(f"[export] {VISION_FILE}")
torch.onnx.export(
vision_tower,
(pixel,),
os.path.join(out_dir, VISION_FILE),
input_names=["pixel_values"],
output_names=["image_features"],
dynamic_axes={"pixel_values": {0: "batch"}, "image_features": {0: "batch"}},
opset_version=OPSET,
do_constant_folding=True,
dynamo=False,
)
def preprocess_images(processor, images: list[np.ndarray]):
"""用官方 processor 做图像预处理,得到与 PyTorch 完全一致的像素张量。"""
from PIL import Image
pil = [Image.fromarray(a) for a in images]
enc = processor(images=pil, return_tensors="pt")
return enc["pixel_values"]
def run_verification(model, processor, out_dir: str, plan: list[str]) -> dict:
"""逐个用例真跑 ONNX 并与 PyTorch 比对;返回参考数据。"""
import onnxruntime as ort
tok = processor.tokenizer
text_sess = ort.InferenceSession(os.path.join(out_dir, TEXT_FILE),
providers=["CPUExecutionProvider"])
vision_sess = ort.InferenceSession(os.path.join(out_dir, VISION_FILE),
providers=["CPUExecutionProvider"])
executed: list[str] = []
reference: dict = {"texts": [], "images": []}
print("\n[verify] 文本塔")
for text in FIXTURE_TEXTS:
enc = tok([text], padding="max_length", truncation=True,
max_length=MAX_LENGTH, return_tensors="pt")
ids = enc["input_ids"].numpy().astype(np.int64)
mask = enc["attention_mask"].numpy().astype(np.int64)
pt = tower_forward_text(model, enc["input_ids"], enc["attention_mask"]).numpy()
ox = text_sess.run(["text_features"], {"input_ids": ids, "attention_mask": mask})[0]
cos = norm_cos(pt, ox)
name = f"text:{text[:24]}"
executed.append(name)
print(f" cos={cos:.9f} ids[:8]={ids[0][:8].tolist()} {text[:28]}")
if cos < 0.9999:
raise SystemExit(f"文本塔导出不一致: {name} cos={cos}")
reference["texts"].append({"text": text, "input_ids": ids[0].tolist(),
"attention_mask": mask[0].tolist(),
"vector": [float(v) for v in ox.reshape(-1)]})
print("\n[verify] 视觉塔")
images = [solid(c) for _, c in FIXTURE_IMAGES] + [gradient()]
names = [n for n, _ in FIXTURE_IMAGES] + ["gradient"]
pixel = preprocess_images(processor, images)
import torch
for i, (nm, _) in enumerate(zip(names, images)):
px = pixel[i:i + 1]
pt = tower_forward_vision(model, px).numpy()
ox = vision_sess.run(["image_features"],
{"pixel_values": px.numpy().astype(np.float32)})[0]
cos = norm_cos(pt, ox)
executed.append(f"image:{nm}")
print(f" cos={cos:.9f} {nm}")
if cos < 0.9999:
raise SystemExit(f"视觉塔导出不一致: {nm} cos={cos}")
# 参考向量直接存像素张量的 sha256,Go 侧用同一预处理即可复算
reference["images"].append({
"name": nm,
"pixels_sha256": hashlib.sha256(px.numpy().astype(np.float32).tobytes()).hexdigest(),
"vector": [float(v) for v in ox.reshape(-1)],
})
missing = [c for c in plan if c not in executed]
extra = [c for c in executed if c not in plan]
if missing or extra:
raise SystemExit(f"用例覆盖不一致: 缺 {missing} 多 {extra}")
print(f"\n[verify] 覆盖度 OK({len(executed)} 个用例,计划 {len(plan)})")
return reference
def main() -> int:
ap = argparse.ArgumentParser()
ap.add_argument("--out", required=True)
ap.add_argument("--model-dir", default="")
ap.add_argument("--model-id", default=DEFAULT_MODEL_ID)
ap.add_argument("--verify-only", action="store_true")
ap.add_argument("--no-reference", action="store_true")
ap.add_argument("--skip-verify", action="store_true")
args = ap.parse_args()
plan = [f"text:{t[:24]}" for t in FIXTURE_TEXTS] + \
[f"image:{n}" for n, _ in FIXTURE_IMAGES] + ["image:gradient"]
if not args.verify_only:
# 清空重建,避免旧产物被当成这次的成果
if os.path.isdir(args.out):
shutil.rmtree(args.out)
os.makedirs(args.out, exist_ok=True)
model, processor = load_model(args.model_dir or None, args.model_id,
local_only=bool(args.model_dir))
if not os.path.isdir(args.out):
os.makedirs(args.out, exist_ok=True)
if not args.verify_only:
export_onnx(model, processor, args.out)
# 词表随产物一起放:provider 只依赖这个目录
src_vocab = os.path.join(args.model_dir, "vocab.txt") if args.model_dir else None
if src_vocab and os.path.exists(src_vocab):
shutil.copy2(src_vocab, os.path.join(args.out, "vocab.txt"))
reference = None
if not args.skip_verify:
reference = run_verification(model, processor, args.out, plan)
else:
print("[verify] 已按 --skip-verify 跳过(不据此宣布成功)")
if not args.verify_only:
cfg = {
"arch": "chinese-clip-vit-base-patch16",
"dim": 512,
"text_onnx": TEXT_FILE,
"vision_onnx": VISION_FILE,
"max_length": MAX_LENGTH,
"image_size": IMAGE_SIZE,
"resample": "bicubic",
"rescale": 1.0 / 255.0,
"image_mean": [0.48145466, 0.4578275, 0.40821073],
"image_std": [0.26862954, 0.26130258, 0.27577711],
"normalize_vector": True, # provider 必须 L2 归一化后再入库
"tokenizer": {
"type": "bert-wordpiece",
"vocab": "vocab.txt",
"do_lower_case": True,
"tokenize_chinese_chars": True,
"cls_id": 101, "sep_id": 102, "pad_id": 0, "unk_id": 100,
},
"modalities": ["text", "image"],
"unsupported_modalities": ["audio", "video"],
"notes": "Chinese-CLIP ViT-B/16:视觉 ViT-B/16 + 文本 RoBERTa-wwm-base,"
"输出 512 维共享空间。文本塔取 CLS(pooler)后过 text_projection,"
"视觉塔取 CLS 后过 visual_projection;两者均未在图中归一化,"
"归一化由 provider 负责。",
}
with open(os.path.join(args.out, CONFIG_FILE), "w") as f:
json.dump(cfg, f, ensure_ascii=False, indent=2)
if reference is not None and not args.no_reference:
reference["source"] = {"model_id": args.model_id,
"model_dir": args.model_dir or "(hub)"}
reference["artifacts"] = {n: sha256_file(os.path.join(args.out, n))
for n in (TEXT_FILE, VISION_FILE)}
with open(os.path.join(args.out, REFERENCE_FILE), "w") as f:
json.dump(reference, f, ensure_ascii=False, indent=2)
with open(os.path.join(args.out, "SHA256SUMS"), "w") as f:
for n in sorted(os.listdir(args.out)):
if n == "SHA256SUMS":
continue
p = os.path.join(args.out, n)
if os.path.isfile(p):
f.write(f"{sha256_file(p)} {n}\n")
print(f"\n[out] {args.out}")
for n in sorted(os.listdir(args.out)):
p = os.path.join(args.out, n)
if os.path.isfile(p):
print(f" {n:<20} {os.path.getsize(p) / 1e6:9.1f} MB")
return 0
if __name__ == "__main__":
sys.exit(main())