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 d1959cbe80
commit bfdb395731
11 changed files with 1757 additions and 4 deletions

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 侧自写 bicubicPIL 系数)与官方预处理不会逐位相同,故用余弦阈值;
// 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_filtera=-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_filtera = -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.15GBQwen 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
}
// isBERTWhitespaceHF _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)
}
// isControlHF _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)
}
// isBERTPunctuationASCII 标点区间 或 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)
}
// isCJKHF _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
}