diff --git a/cmd/homed/main.go b/cmd/homed/main.go index f8358ca..650d4dd 100644 --- a/cmd/homed/main.go +++ b/cmd/homed/main.go @@ -48,6 +48,7 @@ import ( // 空白导入内置 provider:它们各自在 init 里注册到 pkg/embedding。 // 想把核心换成自己的模型,只需替换这一行(或另建一个发行版 main)。 + _ "gitcode.com/JianFeeeee/HomeAgent/providers/chineseclip" _ "gitcode.com/JianFeeeee/HomeAgent/providers/qwen3vl" ) diff --git a/docs/zh/multimodal-space.md b/docs/zh/multimodal-space.md index d3ef043..b4b4d88 100644 --- a/docs/zh/multimodal-space.md +++ b/docs/zh/multimodal-space.md @@ -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 名**从公共注册表 diff --git a/go.mod b/go.mod index 330d03c..01a9c85 100644 --- a/go.mod +++ b/go.mod @@ -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 diff --git a/providers/chineseclip/config.go b/providers/chineseclip/config.go new file mode 100644 index 0000000..c3f21cb --- /dev/null +++ b/providers/chineseclip/config.go @@ -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 +} diff --git a/providers/chineseclip/embedder.go b/providers/chineseclip/embedder.go new file mode 100644 index 0000000..c8d64d2 --- /dev/null +++ b/providers/chineseclip/embedder.go @@ -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 "" +} diff --git a/providers/chineseclip/embedder_onnx_test.go b/providers/chineseclip/embedder_onnx_test.go new file mode 100644 index 0000000..0a4fabc --- /dev/null +++ b/providers/chineseclip/embedder_onnx_test.go @@ -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("产物不完整时应打开失败") + } +} diff --git a/providers/chineseclip/embedder_stub.go b/providers/chineseclip/embedder_stub.go new file mode 100644 index 0000000..07e12be --- /dev/null +++ b/providers/chineseclip/embedder_stub.go @@ -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() {} diff --git a/providers/chineseclip/image.go b/providers/chineseclip/image.go new file mode 100644 index 0000000..4030d59 --- /dev/null +++ b/providers/chineseclip/image.go @@ -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 +} diff --git a/providers/chineseclip/tokenizer.go b/providers/chineseclip/tokenizer.go new file mode 100644 index 0000000..1d48211 --- /dev/null +++ b/providers/chineseclip/tokenizer.go @@ -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 +} diff --git a/providers/chineseclip/tokenizer_test.go b/providers/chineseclip/tokenizer_test.go new file mode 100644 index 0000000..459aeae --- /dev/null +++ b/providers/chineseclip/tokenizer_test.go @@ -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 +} diff --git a/scripts/export_chineseclip_onnx.py b/scripts/export_chineseclip_onnx.py new file mode 100644 index 0000000..019df24 --- /dev/null +++ b/scripts/export_chineseclip_onnx.py @@ -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())