mirror of
https://gitcode.com/JianFeeeee/HomeAgent.git
synced 2026-10-04 00:03:59 +00:00
refactor(memory): 核心不再适配具体模型——公共 embedding provider SPI + 注册表
问题:cmd/homed 里 `case "onnx": qwen.New(modelDir)` 把模型适配写进了核心, `type=onnx` 名义上是格式、实际写死了一个模型家族;2117 行 Qwen 专属代码 (BPE、chat template、M-RoPE、Vision_gN 命名)住在内核树里,还带着一对 `//go:build onnxruntime` 的 stub。加任何新模型都要改内核。 现在核心只认一个模型无关的公共契约(pkg/embedding): - 输入是不透明的 Data+MIME,解码/预处理/时序分组全归 provider - 能力是数据(Info.Modalities),不是接口方法——新增模态无需改核心接口 - 不支持的模态返回 embedding.ErrUnsupportedModality(可 errors.Is 识别) - 按名字注册,重复注册 panic;Options 是 provider 私有命名空间,核心不解释 改动: - 新增 pkg/embedding:Modality/Purpose/Input/Info/Provider/Config + 注册表 (Open 校验 Info,ValidateVector 在入库前拦下维度错与非有限值) - providers/qwen3vl:Qwen 实现整体移出内核(git mv),实现公共 SPI 并自注册 - internal/memory/vector:新增 ProviderAdapter(公共 SPI → 内部小接口); ErrModalityUnsupported 改为公共哨兵别名;删除 VideoEmbedder 可选接口 (那正是「核心为每个新模态长方法」的坏味道) - http embedder 也变成普通 provider(注册名 http) - cmd/homed:删除 qwen import 与 onnx/http 分支,改为按 provider 名打开 + 透传 options.*;provider 打开失败只警告并禁用多模态检索,不影响启动 - config:multimodal_space.type/onnx./http.* → provider + options.* - 删除 internal/memory/qwen(整体搬迁) 测试: - pkg/embedding:注册表隔离/未知名字/非法 Info 自动关闭/ValidateVector - vector:适配器原样透传字节与 MIME、维度错被拦、Close 幂等且停止使用、 两个哨兵 errors.Is 互通 - providers/qwen3vl:新增公共 SPI 全链路集成测试(Open→Info→Embed→ 未知模态哨兵),并明确断言 Info 不声明 video 已知未完成(不得当作已验证): - 视频冻结回归 TestEmbedderVideoMatchesONNXReference **显式跳过**:Go 侧 video 模板缺少 processor 按时间组插入的字面时间戳文本 (<0.0 seconds>/<1.0 seconds>),同一输入 Python seq=1190(1152+38)、 Go 只有 22 个文本 token。时间戳也占 M-RoPE 位置,故现有 M-RoPE 自洽断言 通过不能证明与官方实现一致。修复属 provider 内部工作。 - 视觉侧三档已导出并逐档校验通过(cos 1.000000119/1.000000119/1.000000000) 验证:go build ./... ;go vet -tags onnxruntime ./... ; go test -short ./internal/memory/... ./internal/agent/core/... ./internal/sdk/... ./pkg/... ;onnxruntime 下 providers/qwen3vl 全绿(视频为显式 skip)
This commit is contained in:
625
providers/qwen3vl/embedder.go
Normal file
625
providers/qwen3vl/embedder.go
Normal file
@ -0,0 +1,625 @@
|
||||
//go:build onnxruntime
|
||||
|
||||
// Package qwen3vl provides the optional Qwen3-VL-Embedding ONNX provider.
|
||||
// Model-specific tokenization, preprocessing, graph layout, and runtime code
|
||||
// live here rather than in the HomeAgent core.
|
||||
package qwen3vl
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"math"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
ort "github.com/yalue/onnxruntime_go"
|
||||
|
||||
"gitcode.com/JianFeeeee/HomeAgent/pkg/embedding"
|
||||
)
|
||||
|
||||
func init() {
|
||||
embedding.Register("qwen3vl", func(cfg embedding.Config) (embedding.Provider, error) {
|
||||
modelDir := cfg.Options["model_dir"]
|
||||
e, err := New(modelDir)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return e, nil
|
||||
})
|
||||
}
|
||||
|
||||
type embedConfig struct {
|
||||
Arch string `json:"arch"`
|
||||
Dimension int `json:"dim"`
|
||||
MaxLength int `json:"max_length"`
|
||||
Instruction string `json:"instruction"`
|
||||
Pooling string `json:"pooling"`
|
||||
ImageSize int `json:"image_size"`
|
||||
PatchSize int `json:"patch_size"`
|
||||
TemporalPatch int `json:"temporal_patch_size"`
|
||||
SpatialMerge int `json:"spatial_merge_size"`
|
||||
ImageMean []float64 `json:"image_mean"`
|
||||
ImageStd []float64 `json:"image_std"`
|
||||
RopeTheta float64 `json:"rope_theta"`
|
||||
MRopeSection []int `json:"mrope_section"`
|
||||
}
|
||||
|
||||
type Embedder struct {
|
||||
mu sync.RWMutex
|
||||
|
||||
loaded bool
|
||||
dir string
|
||||
config embedConfig
|
||||
tok *Tokenizer
|
||||
token *ort.DynamicAdvancedSession
|
||||
transform *ort.DynamicAdvancedSession
|
||||
|
||||
// vision 按时间组数缓存视觉图会话:1 = 单图(Vision.onnx),
|
||||
// G>1 = 视频(Vision_g{G}.onnx)。每张图约 1.6GB,因此按需加载而不是
|
||||
// 启动时全开;未用到的档位不占内存。
|
||||
visionMu sync.Mutex
|
||||
vision map[int]*ort.DynamicAdvancedSession
|
||||
|
||||
fp string
|
||||
close sync.Once
|
||||
}
|
||||
|
||||
func New(modelDir string) (*Embedder, error) {
|
||||
if modelDir == "" {
|
||||
return nil, fmt.Errorf("qwen model dir not specified")
|
||||
}
|
||||
cfgRaw, err := os.ReadFile(filepath.Join(modelDir, "embed_config.json"))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("read embed_config.json: %w", err)
|
||||
}
|
||||
var cfg embedConfig
|
||||
if err := json.Unmarshal(cfgRaw, &cfg); err != nil {
|
||||
return nil, fmt.Errorf("parse embed_config.json: %w", err)
|
||||
}
|
||||
if cfg.Dimension != 2048 || cfg.MaxLength < 598 || cfg.Pooling != "last_token" {
|
||||
return nil, fmt.Errorf("qwen: incompatible config dim=%d max_length=%d pooling=%q", cfg.Dimension, cfg.MaxLength, cfg.Pooling)
|
||||
}
|
||||
if cfg.ImageSize != qwenImageSize || cfg.PatchSize != qwenPatchSize || cfg.TemporalPatch != qwenTemporalPatch || cfg.SpatialMerge != qwenSpatialMerge {
|
||||
return nil, fmt.Errorf("qwen: incompatible vision layout image=%d patch=%d temporal=%d merge=%d", cfg.ImageSize, cfg.PatchSize, cfg.TemporalPatch, cfg.SpatialMerge)
|
||||
}
|
||||
if cfg.RopeTheta <= 0 || len(cfg.MRopeSection) != 3 || cfg.MRopeSection[0]+cfg.MRopeSection[1]+cfg.MRopeSection[2] != qwenRotaryHalfDim {
|
||||
return nil, fmt.Errorf("qwen: incompatible rope theta=%g section=%v", cfg.RopeTheta, cfg.MRopeSection)
|
||||
}
|
||||
|
||||
tok, err := LoadTokenizer(modelDir)
|
||||
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("init onnx env: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
token, err := ort.NewDynamicAdvancedSession(
|
||||
filepath.Join(modelDir, "TokenEmbedding.onnx"),
|
||||
[]string{"input_ids"}, []string{"hidden"}, nil,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("create qwen token embedding session: %w", err)
|
||||
}
|
||||
transform, err := ort.NewDynamicAdvancedSession(
|
||||
filepath.Join(modelDir, "Transformer.onnx"),
|
||||
[]string{"hidden", "deepstack_0", "deepstack_1", "deepstack_2", "rotary_cos", "rotary_sin", "causal_mask"},
|
||||
[]string{"embedding"}, nil,
|
||||
)
|
||||
if err != nil {
|
||||
token.Destroy()
|
||||
return nil, fmt.Errorf("create qwen transformer session: %w", err)
|
||||
}
|
||||
vision, err := ort.NewDynamicAdvancedSession(
|
||||
filepath.Join(modelDir, "Vision.onnx"), []string{"pixel_values"},
|
||||
[]string{"deepstack_feature_0", "deepstack_feature_1", "deepstack_feature_2", "vision_hidden_states"}, nil,
|
||||
)
|
||||
if err != nil {
|
||||
token.Destroy()
|
||||
transform.Destroy()
|
||||
return nil, fmt.Errorf("create qwen vision session: %w", err)
|
||||
}
|
||||
|
||||
return &Embedder{
|
||||
loaded: true, dir: modelDir, config: cfg, tok: tok,
|
||||
token: token, transform: transform,
|
||||
vision: map[int]*ort.DynamicAdvancedSession{1: vision},
|
||||
fp: computeFingerprint(modelDir),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (e *Embedder) VectorizeDense(text string) ([]float64, error) {
|
||||
e.mu.RLock()
|
||||
loaded, cfg, tok := e.loaded, e.config, e.tok
|
||||
e.mu.RUnlock()
|
||||
if !loaded {
|
||||
return nil, fmt.Errorf("qwen embedder not loaded")
|
||||
}
|
||||
ids, _, position, _, err := tok.textModelInput(cfg.Instruction, text, cfg.MaxLength)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
hidden, err := e.runTokenEmbedding(ids)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
deep := make([][]float32, 3)
|
||||
for i := range deep {
|
||||
deep[i] = make([]float32, len(hidden))
|
||||
}
|
||||
return e.runTransformer(hidden, deep, position, len(ids))
|
||||
}
|
||||
|
||||
// EmbedImageDense 把一张图片编码到统一空间。
|
||||
//
|
||||
// mime 决定这个模态是否在本空间的原生覆盖范围内:Qwen3-VL 能原生编码文本、
|
||||
// 图像与视频,但**不原生支持音频**(模型卡与 config 双重确认:没有
|
||||
// audio_token_id / audio_config)。音频必须返回 ErrModalityUnsupported,
|
||||
// 而不是拿视觉塔去编码——那会往统一空间里灌入语义错误的坐标,而错误是静默的。
|
||||
//
|
||||
// video/*(视频文件)也在这里拒绝:本函数的入参是**单帧字节**,Go 侧没有
|
||||
// 视频解码器;多帧请走 EmbedVideoDense。
|
||||
func (e *Embedder) EmbedImageDense(raw []byte, mime string) ([]float64, error) {
|
||||
if err := checkImageMime(mime); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
pixels, err := preprocessImage(raw)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return e.embedVision(pixels, 1)
|
||||
}
|
||||
|
||||
// EmbedVideoDense 把已按时间排序的视频帧编码到统一空间(原生跨帧时序)。
|
||||
//
|
||||
// frames 是已解码的帧(PNG/JPEG 字节),相邻两帧构成一个时间组;
|
||||
// groups = len(frames)/2 必须恰好是导出时固定的某一档(Vision_g{G}.onnx),
|
||||
// 否则本方法明确报错并告知已导出哪些档位。
|
||||
//
|
||||
// 为何必须精确匹配而不能“差不多就行”:视觉塔的注意力按 grid 划分,
|
||||
// 用 G=2 的图喂 G=3 的数据是未定义行为。实测(onnxruntime)会因维度不符
|
||||
// 报 InvalidArgument,因此不会静默算错——但也不该依赖那次报错来兜底。
|
||||
//
|
||||
// 帧数为奇数时只用得上前 2×floor(n/2) 帧,多余一帧被丢弃(不补重复帧:
|
||||
// 那会改变跨帧注意力看到的运动)。
|
||||
func (e *Embedder) EmbedVideoDense(frames [][]byte, mime string) ([]float64, error) {
|
||||
if err := checkVideoMime(mime); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(frames) < qwenTemporalPatch {
|
||||
return nil, fmt.Errorf("qwen: video needs at least %d frames, got %d", qwenTemporalPatch, len(frames))
|
||||
}
|
||||
groups := len(frames) / qwenTemporalPatch
|
||||
// 先确认这一档的视觉图确实已导出,再去做昂贵的预处理:
|
||||
// 不然一个未导出档位会先白算一遍(每组 2304×1536 浮点)才报错。
|
||||
if _, err := e.visionFor(groups); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
pixels, groups, err := preprocessVideoFrames(frames)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return e.embedVision(pixels, groups)
|
||||
}
|
||||
|
||||
// checkImageMime 把「不在本空间覆盖范围内」与「参数用错」分开报。
|
||||
func checkImageMime(mime string) error {
|
||||
switch {
|
||||
case strings.HasPrefix(mime, "audio/"):
|
||||
return fmt.Errorf("%w: audio (%s) 不在 Qwen3-VL 原生模态内(无 audio_token_id),需真正的统一音频模型",
|
||||
embedding.ErrUnsupportedModality, mime)
|
||||
case strings.HasPrefix(mime, "video/"):
|
||||
return fmt.Errorf("%w: 视频文件 (%s) 无法解码;多帧请用 EmbedVideoDense",
|
||||
embedding.ErrUnsupportedModality, mime)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func checkVideoMime(mime string) error {
|
||||
if strings.HasPrefix(mime, "audio/") {
|
||||
return fmt.Errorf("%w: audio (%s)", embedding.ErrUnsupportedModality, mime)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// visionFor 返回指定时间组数的视觉图会话,按需创建。
|
||||
func (e *Embedder) visionFor(groups int) (*ort.DynamicAdvancedSession, error) {
|
||||
e.visionMu.Lock()
|
||||
defer e.visionMu.Unlock()
|
||||
if s, ok := e.vision[groups]; ok {
|
||||
return s, nil
|
||||
}
|
||||
name := "Vision.onnx"
|
||||
if groups > 1 {
|
||||
name = fmt.Sprintf("Vision_g%d.onnx", groups)
|
||||
}
|
||||
path := filepath.Join(e.dir, name)
|
||||
if _, err := os.Stat(path); err != nil {
|
||||
return nil, fmt.Errorf("qwen: 缺少 %s(时间组数 %d 未导出,已加载档位 %v);"+
|
||||
"用 scripts/export_qwen3vl_embedding_onnx.py 按所需 --video-groups 重新导出: %w",
|
||||
name, groups, e.visionGroupsLocked(), err)
|
||||
}
|
||||
s, err := ort.NewDynamicAdvancedSession(path, []string{"pixel_values"},
|
||||
[]string{"deepstack_feature_0", "deepstack_feature_1", "deepstack_feature_2", "vision_hidden_states"}, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("create qwen vision session %s: %w", name, err)
|
||||
}
|
||||
e.vision[groups] = s
|
||||
return s, nil
|
||||
}
|
||||
|
||||
// visionGroupsLocked 返回已加载的档位(调用方需持 visionMu)。
|
||||
func (e *Embedder) visionGroupsLocked() []int {
|
||||
out := make([]int, 0, len(e.vision))
|
||||
for g := range e.vision {
|
||||
out = append(out, g)
|
||||
}
|
||||
sort.Ints(out)
|
||||
return out
|
||||
}
|
||||
|
||||
// embedVision 是图像与视频共用的后半段:视觉塔 → 散射到 hidden → Transformer。
|
||||
func (e *Embedder) embedVision(pixels []float32, groups int) ([]float64, error) {
|
||||
e.mu.RLock()
|
||||
loaded, cfg, tok := e.loaded, e.config, e.tok
|
||||
e.mu.RUnlock()
|
||||
if !loaded {
|
||||
return nil, fmt.Errorf("qwen embedder not loaded")
|
||||
}
|
||||
|
||||
vision, err := e.visionFor(groups)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
wantTokens := groups * qwenVisualTokens
|
||||
features, err := e.runVision(vision, pixels, wantTokens, cfg.Dimension)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var ids []int
|
||||
var position []int64
|
||||
var visual []bool
|
||||
if groups == 1 {
|
||||
ids, _, position, visual, err = tok.imageModelInput(cfg.Instruction, cfg.MaxLength)
|
||||
} else {
|
||||
ids, _, position, visual, err = tok.videoModelInput(cfg.Instruction, groups, cfg.MaxLength)
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
hidden, err := e.runTokenEmbedding(ids)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
deep := make([][]float32, 3)
|
||||
for i := range deep {
|
||||
deep[i] = make([]float32, len(hidden))
|
||||
}
|
||||
// 视觉特征按占位符出现顺序就地替换 token embedding:
|
||||
// 图像一个组、视频 G 个组,展开后长度都是 groups×576,与视觉塔输出一致。
|
||||
visualIndex := 0
|
||||
for tokenIndex, isVisual := range visual {
|
||||
if !isVisual {
|
||||
continue
|
||||
}
|
||||
dst := tokenIndex * cfg.Dimension
|
||||
src := visualIndex * cfg.Dimension
|
||||
copy(hidden[dst:dst+cfg.Dimension], features[3][src:src+cfg.Dimension])
|
||||
for layer := range deep {
|
||||
copy(deep[layer][dst:dst+cfg.Dimension], features[layer][src:src+cfg.Dimension])
|
||||
}
|
||||
visualIndex++
|
||||
}
|
||||
if visualIndex != wantTokens {
|
||||
return nil, fmt.Errorf("qwen: injected visual tokens=%d, want %d (groups=%d)", visualIndex, wantTokens, groups)
|
||||
}
|
||||
return e.runTransformer(hidden, deep, position, len(ids))
|
||||
}
|
||||
|
||||
func (e *Embedder) runTokenEmbedding(ids []int) ([]float32, error) {
|
||||
inputIDs := make([]int64, len(ids))
|
||||
for i, id := range ids {
|
||||
inputIDs[i] = int64(id)
|
||||
}
|
||||
in, err := ort.NewTensor(ort.Shape{1, int64(len(ids))}, inputIDs)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("qwen token input: %w", err)
|
||||
}
|
||||
defer in.Destroy()
|
||||
outs := make([]ort.Value, 1)
|
||||
if err := e.token.Run([]ort.Value{in}, outs); err != nil {
|
||||
return nil, fmt.Errorf("qwen token embedding run: %w", err)
|
||||
}
|
||||
if outs[0] == nil {
|
||||
return nil, fmt.Errorf("qwen token embedding output is nil")
|
||||
}
|
||||
defer outs[0].Destroy()
|
||||
tensor, ok := outs[0].(*ort.Tensor[float32])
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("qwen token embedding output type %T", outs[0])
|
||||
}
|
||||
shape := tensor.GetShape()
|
||||
if len(shape) != 3 || shape[0] != 1 || shape[1] != int64(len(ids)) || shape[2] != int64(e.config.Dimension) {
|
||||
return nil, fmt.Errorf("qwen token embedding shape=%v", shape)
|
||||
}
|
||||
return append([]float32(nil), tensor.GetData()...), nil
|
||||
}
|
||||
|
||||
func (e *Embedder) runVision(sess *ort.DynamicAdvancedSession, pixels []float32, wantTokens, dim int) ([][]float32, error) {
|
||||
wantPatches := int64(wantTokens) * qwenSpatialMerge * qwenSpatialMerge
|
||||
if int64(len(pixels)) != wantPatches*qwenPatchVectorSize {
|
||||
return nil, fmt.Errorf("qwen vision input: %d floats, want %d patches × %d",
|
||||
len(pixels), wantPatches, qwenPatchVectorSize)
|
||||
}
|
||||
in, err := ort.NewTensor(ort.Shape{wantPatches, qwenPatchVectorSize}, pixels)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("qwen vision input: %w", err)
|
||||
}
|
||||
defer in.Destroy()
|
||||
outs := make([]ort.Value, 4)
|
||||
if err := sess.Run([]ort.Value{in}, outs); err != nil {
|
||||
return nil, fmt.Errorf("qwen vision run: %w", err)
|
||||
}
|
||||
features := make([][]float32, 4)
|
||||
for i, value := range outs {
|
||||
if value == nil {
|
||||
return nil, fmt.Errorf("qwen vision output %d is nil", i)
|
||||
}
|
||||
defer value.Destroy()
|
||||
tensor, ok := value.(*ort.Tensor[float32])
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("qwen vision output %d type %T", i, value)
|
||||
}
|
||||
shape := tensor.GetShape()
|
||||
if len(shape) != 2 || shape[0] != int64(wantTokens) || shape[1] != int64(dim) {
|
||||
return nil, fmt.Errorf("qwen vision output %d shape=%v, want [%d %d]", i, shape, wantTokens, dim)
|
||||
}
|
||||
features[i] = append([]float32(nil), tensor.GetData()...)
|
||||
}
|
||||
return features, nil
|
||||
}
|
||||
|
||||
func (e *Embedder) runTransformer(hidden []float32, deep [][]float32, position []int64, seq int) ([]float64, error) {
|
||||
if len(hidden) != seq*e.config.Dimension || len(deep) != 3 || len(position) != 3*seq {
|
||||
return nil, fmt.Errorf("qwen: invalid transformer inputs hidden=%d deep=%d position=%d seq=%d", len(hidden), len(deep), len(position), seq)
|
||||
}
|
||||
cos, sin := e.rotary(position, seq)
|
||||
causal := causalMask(seq)
|
||||
|
||||
hiddenTensor, err := ort.NewTensor(ort.Shape{1, int64(seq), int64(e.config.Dimension)}, hidden)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("qwen hidden tensor: %w", err)
|
||||
}
|
||||
defer hiddenTensor.Destroy()
|
||||
inputs := []ort.Value{hiddenTensor}
|
||||
var deepTensors []*ort.Tensor[float32]
|
||||
for i, data := range deep {
|
||||
if len(data) != len(hidden) {
|
||||
return nil, fmt.Errorf("qwen deepstack %d length=%d, want %d", i, len(data), len(hidden))
|
||||
}
|
||||
t, err := ort.NewTensor(ort.Shape{1, int64(seq), int64(e.config.Dimension)}, data)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("qwen deepstack %d tensor: %w", i, err)
|
||||
}
|
||||
deepTensors = append(deepTensors, t)
|
||||
inputs = append(inputs, t)
|
||||
}
|
||||
defer func() {
|
||||
for _, t := range deepTensors {
|
||||
t.Destroy()
|
||||
}
|
||||
}()
|
||||
cosTensor, err := ort.NewTensor(ort.Shape{1, int64(seq), qwenRotaryDim}, cos)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("qwen rotary cos: %w", err)
|
||||
}
|
||||
defer cosTensor.Destroy()
|
||||
sinTensor, err := ort.NewTensor(ort.Shape{1, int64(seq), qwenRotaryDim}, sin)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("qwen rotary sin: %w", err)
|
||||
}
|
||||
defer sinTensor.Destroy()
|
||||
causalTensor, err := ort.NewTensor(ort.Shape{1, 1, int64(seq), int64(seq)}, causal)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("qwen causal mask: %w", err)
|
||||
}
|
||||
defer causalTensor.Destroy()
|
||||
inputs = append(inputs, cosTensor, sinTensor, causalTensor)
|
||||
|
||||
out, err := ort.NewEmptyTensor[float32](ort.Shape{1, int64(e.config.Dimension)})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("qwen output tensor: %w", err)
|
||||
}
|
||||
defer out.Destroy()
|
||||
if err := e.transform.Run(inputs, []ort.Value{out}); err != nil {
|
||||
return nil, fmt.Errorf("qwen transformer run: %w", err)
|
||||
}
|
||||
return normalize(out.GetData()), nil
|
||||
}
|
||||
|
||||
const (
|
||||
qwenRotaryHalfDim = 64
|
||||
qwenRotaryDim = 128
|
||||
)
|
||||
|
||||
func (e *Embedder) rotary(position []int64, seq int) ([]float32, []float32) {
|
||||
cos := make([]float32, seq*qwenRotaryDim)
|
||||
sin := make([]float32, seq*qwenRotaryDim)
|
||||
inv := make([]float64, qwenRotaryHalfDim)
|
||||
for i := range inv {
|
||||
inv[i] = 1 / math.Pow(e.config.RopeTheta, float64(2*i)/qwenRotaryDim)
|
||||
}
|
||||
for token := 0; token < seq; token++ {
|
||||
freq := make([]float64, qwenRotaryHalfDim)
|
||||
for i := range freq {
|
||||
freq[i] = float64(position[token]) * inv[i]
|
||||
}
|
||||
for dim, offset := range []int{0, 1, 2} {
|
||||
if dim == 0 {
|
||||
continue
|
||||
}
|
||||
limit := e.config.MRopeSection[dim] * 3
|
||||
for i := offset; i < limit; i += 3 {
|
||||
freq[i] = float64(position[dim*seq+token]) * inv[i]
|
||||
}
|
||||
}
|
||||
for i, f := range freq {
|
||||
c, s := float32(math.Cos(f)), float32(math.Sin(f))
|
||||
cos[token*qwenRotaryDim+i] = c
|
||||
cos[token*qwenRotaryDim+qwenRotaryHalfDim+i] = c
|
||||
sin[token*qwenRotaryDim+i] = s
|
||||
sin[token*qwenRotaryDim+qwenRotaryHalfDim+i] = s
|
||||
}
|
||||
}
|
||||
return cos, sin
|
||||
}
|
||||
|
||||
func causalMask(seq int) []float32 {
|
||||
out := make([]float32, seq*seq)
|
||||
for row := 0; row < seq; row++ {
|
||||
for col := row + 1; col < seq; col++ {
|
||||
out[row*seq+col] = -math.MaxFloat32
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func normalize(raw []float32) []float64 {
|
||||
out := make([]float64, len(raw))
|
||||
var norm float64
|
||||
for i, v := range raw {
|
||||
out[i] = float64(v)
|
||||
norm += out[i] * out[i]
|
||||
}
|
||||
if norm > 0 {
|
||||
norm = math.Sqrt(norm)
|
||||
for i := range out {
|
||||
out[i] /= norm
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func (e *Embedder) Embed(_ context.Context, in embedding.Input) ([]float64, error) {
|
||||
switch in.Modality {
|
||||
case embedding.ModalityText:
|
||||
return e.VectorizeDense(in.Text)
|
||||
case embedding.ModalityImage:
|
||||
return e.EmbedImageDense(in.Data, in.MIME)
|
||||
case embedding.ModalityVideo:
|
||||
// 公共契约把视频交给 provider 自行解码,而本 provider 没有视频解码器
|
||||
// (Go 标准库不含 H.264/MP4)。这里必须明确说「本 provider 不提供视频
|
||||
// 文件编码」,而不是假装支持后再拿错数据算出一个语义错误的向量。
|
||||
//
|
||||
// 可用的视频路径是本 provider 自己的 EmbedVideoDense:调用方先抽帧,
|
||||
// 由本 provider 按自己的时序窗口分组。
|
||||
return nil, fmt.Errorf("%w: video(本 provider 不内嵌视频解码器;"+
|
||||
"请先抽帧并调用 qwen3vl 的 EmbedVideoDense)", embedding.ErrUnsupportedModality)
|
||||
default:
|
||||
return nil, fmt.Errorf("%w: %s", embedding.ErrUnsupportedModality, in.Modality)
|
||||
}
|
||||
}
|
||||
|
||||
// Info 只声明本 provider 能通过**公共契约**提供的模态。
|
||||
//
|
||||
// 视频不在其中:契约要求 provider 自行解码 Data,而本 provider 没有视频
|
||||
// 解码器;列进来会让核心据以创建 video 输入,然后在运行时全部失败。
|
||||
// 视频能力由本 provider 自己的 EmbedVideoDense(接收已解码帧)提供,
|
||||
// 待核心有了对 provider 不透明的多帧容器后再纳入契约。
|
||||
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.close.Do(func() {
|
||||
e.mu.Lock()
|
||||
defer e.mu.Unlock()
|
||||
if e.token != nil {
|
||||
e.token.Destroy()
|
||||
e.token = nil
|
||||
}
|
||||
if e.transform != nil {
|
||||
e.transform.Destroy()
|
||||
e.transform = nil
|
||||
}
|
||||
e.visionMu.Lock()
|
||||
for groups, sess := range e.vision {
|
||||
if sess != nil {
|
||||
sess.Destroy()
|
||||
}
|
||||
delete(e.vision, groups)
|
||||
}
|
||||
e.visionMu.Unlock()
|
||||
e.loaded = false
|
||||
})
|
||||
}
|
||||
|
||||
func computeFingerprint(modelDir string) string {
|
||||
h := sha256.New()
|
||||
graphNames := []string{"TokenEmbedding.onnx", "Transformer.onnx", "Vision.onnx", "embed_config.json"}
|
||||
// 视频是按时间组数各导一张图,因此每一张都必须进入指纹:
|
||||
// 漏掉它们会让「换了视频图但指纹没变」,历史向量不会重算。
|
||||
if entries, err := os.ReadDir(modelDir); err == nil {
|
||||
var extra []string
|
||||
for _, entry := range entries {
|
||||
if strings.HasPrefix(entry.Name(), "Vision_g") && strings.HasSuffix(entry.Name(), ".onnx") {
|
||||
extra = append(extra, entry.Name())
|
||||
}
|
||||
}
|
||||
sort.Strings(extra)
|
||||
graphNames = append(graphNames, extra...)
|
||||
}
|
||||
for _, name := range graphNames {
|
||||
if data, err := os.ReadFile(filepath.Join(modelDir, name)); err == nil {
|
||||
h.Write([]byte(name))
|
||||
h.Write([]byte{0})
|
||||
h.Write(data)
|
||||
h.Write([]byte{0})
|
||||
}
|
||||
}
|
||||
entries, _ := os.ReadDir(modelDir)
|
||||
var names []string
|
||||
for _, entry := range entries {
|
||||
n := entry.Name()
|
||||
if strings.HasPrefix(n, "embed_tokens.") || strings.HasPrefix(n, "layers.") || strings.HasPrefix(n, "onnx__") || strings.HasSuffix(n, ".onnx.data") {
|
||||
names = append(names, n)
|
||||
}
|
||||
}
|
||||
sort.Strings(names)
|
||||
for _, n := range names {
|
||||
if info, err := os.Stat(filepath.Join(modelDir, n)); err == nil {
|
||||
fmt.Fprintf(h, "%s:%d\n", n, info.Size())
|
||||
}
|
||||
}
|
||||
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 ""
|
||||
}
|
||||
556
providers/qwen3vl/embedder_onnx_test.go
Normal file
556
providers/qwen3vl/embedder_onnx_test.go
Normal file
@ -0,0 +1,556 @@
|
||||
//go:build onnxruntime
|
||||
|
||||
package qwen3vl
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"image"
|
||||
"image/png"
|
||||
"math"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"gitcode.com/JianFeeeee/HomeAgent/pkg/embedding"
|
||||
)
|
||||
|
||||
// onnxModelDir 返回三段式 Qwen 多模态 ONNX 产物目录。
|
||||
//
|
||||
// 产物约 8GB(含外部权重),不进仓库;由 scripts/export_qwen3vl_embedding_onnx.py
|
||||
// 自动拉取模型并导出。可通过 QWEN_ONNX_MODEL_DIR 指向别处;目录不存在时相关
|
||||
// 测试跳过,而不是失败——CI 与本机开发者都不一定有这份产物。
|
||||
func onnxModelDir() string {
|
||||
if v := os.Getenv("QWEN_ONNX_MODEL_DIR"); v != "" {
|
||||
return v
|
||||
}
|
||||
return "/home/newqqagent/models/qwen3-vl-embed-multimodal-onnx"
|
||||
}
|
||||
|
||||
// artifactDeclaresVideo 读产物自带的 embed_config.json,判断它是否声明支持原生视频。
|
||||
//
|
||||
// 用途:把「这个产物本来就不含视频」与「这个产物应该有视频,但参考里没有」分开。
|
||||
// 后者是产物/参考不匹配,必须报错而不是跳过——否则一个声明了视频支持的目录
|
||||
// 可以带着空视频参考一路「通过」。
|
||||
func artifactDeclaresVideo(t *testing.T, dir string) bool {
|
||||
t.Helper()
|
||||
data, err := os.ReadFile(filepath.Join(dir, "embed_config.json"))
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
var cfg struct {
|
||||
SupportsNativeVideo bool `json:"supports_native_video"`
|
||||
VideoGroups []int `json:"video_groups"`
|
||||
}
|
||||
if err := json.Unmarshal(data, &cfg); err != nil {
|
||||
return false
|
||||
}
|
||||
return cfg.SupportsNativeVideo || len(cfg.VideoGroups) > 0
|
||||
}
|
||||
|
||||
func requireONNXArtifacts(t *testing.T) string {
|
||||
t.Helper()
|
||||
dir := onnxModelDir()
|
||||
for _, name := range []string{"TokenEmbedding.onnx", "Transformer.onnx", "Vision.onnx", "embed_config.json", "tokenizer.json"} {
|
||||
if _, err := os.Stat(dir + "/" + name); err != nil {
|
||||
t.Skipf("ONNX 产物不完整(%s: %v),跳过;用 scripts/export_qwen3vl_embedding_onnx.py 导出", name, err)
|
||||
}
|
||||
}
|
||||
return dir
|
||||
}
|
||||
|
||||
// onnxReference 是导出脚本 `--emit-reference` 写出的冻结参考。
|
||||
//
|
||||
// 刻意不把浮点常量硬编码在测试里:参考值必须能追溯到「哪个模型、哪次导出、
|
||||
// 什么输入」,而不是一组无人知道出处的数字。参考文件里的 RGB/尺寸同时用来
|
||||
// 构造测试图片,保证输入与参考按构造一致,不会因测试改动而静默错位。
|
||||
type onnxReference struct {
|
||||
Text string `json:"text"`
|
||||
TextVectorPrefix []float64 `json:"text_vector_prefix"`
|
||||
ImageRGB []int `json:"image_rgb"`
|
||||
ImageSize int `json:"image_size"`
|
||||
ImageVectorPrefix []float64 `json:"image_vector_prefix"`
|
||||
Dim int `json:"dim"`
|
||||
|
||||
// 视频参考:相邻两帧构成一个时间组(tp0←帧2g、tp1←帧2g+1),
|
||||
// 帧颜色用来构造与导出脚本完全一致的测试输入。
|
||||
VideoGroups int `json:"video_groups"`
|
||||
VideoFrameRGB [][]int `json:"video_frame_rgb"`
|
||||
VideoVectorPrefix []float64 `json:"video_vector_prefix"`
|
||||
}
|
||||
|
||||
func loadReference(t *testing.T, dir string) *onnxReference {
|
||||
t.Helper()
|
||||
path := os.Getenv("QWEN_ONNX_REFERENCE")
|
||||
if path == "" {
|
||||
path = filepath.Join(dir, "qwen_reference.json")
|
||||
}
|
||||
raw, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
t.Skipf("缺少冻结参考 %s(由 scripts/export_qwen3vl_embedding_onnx.py --emit-reference 生成): %v", path, err)
|
||||
}
|
||||
var ref onnxReference
|
||||
if err := json.Unmarshal(raw, &ref); err != nil {
|
||||
t.Fatalf("解析参考 %s: %v", path, err)
|
||||
}
|
||||
if ref.Text == "" || len(ref.TextVectorPrefix) == 0 || len(ref.ImageRGB) != 3 || ref.ImageSize <= 0 {
|
||||
t.Fatalf("参考 %s 不完整: %+v", path, ref)
|
||||
}
|
||||
return &ref
|
||||
}
|
||||
|
||||
// solidPNG 生成一张 size×size 纯色 PNG,供跨语言冻结向量回归。
|
||||
//
|
||||
// 刻意用纯色且尺寸与视觉塔一致:Go 侧预处理对已是 768×768 的输入不做插值、
|
||||
// 不补边,于是 patch 张量只由布局决定。一旦 patch 排列写错(内层循环顺序、
|
||||
// merge 分组顺序、通道顺序),冻结向量立刻不匹配——而那类错误在人工看图时
|
||||
// 几乎发现不了。
|
||||
func solidPNG(t *testing.T, size int, r, g, b uint8) []byte {
|
||||
t.Helper()
|
||||
img := image.NewNRGBA(image.Rect(0, 0, size, size))
|
||||
for y := 0; y < size; y++ {
|
||||
for x := 0; x < size; x++ {
|
||||
i := img.PixOffset(x, y)
|
||||
img.Pix[i], img.Pix[i+1], img.Pix[i+2], img.Pix[i+3] = r, g, b, 255
|
||||
}
|
||||
}
|
||||
var buf bytes.Buffer
|
||||
if err := png.Encode(&buf, img); err != nil {
|
||||
t.Fatalf("encode png: %v", err)
|
||||
}
|
||||
return buf.Bytes()
|
||||
}
|
||||
|
||||
func newTestEmbedder(t *testing.T, dir string) *Embedder {
|
||||
t.Helper()
|
||||
e, err := New(dir)
|
||||
if err != nil {
|
||||
t.Fatalf("New: %v", err)
|
||||
}
|
||||
t.Cleanup(e.Close)
|
||||
info := e.Info()
|
||||
if info.Dimension != 2048 || info.Fingerprint == "" {
|
||||
t.Fatalf("元数据异常: dim=%d fingerprint=%q", info.Dimension, info.Fingerprint)
|
||||
}
|
||||
if err := embedding.ValidateInfo(info); err != nil {
|
||||
t.Fatalf("Info 不满足公共契约: %v", err)
|
||||
}
|
||||
return e
|
||||
}
|
||||
|
||||
func assertNormalized(t *testing.T, name string, got []float64, dim int) {
|
||||
t.Helper()
|
||||
if dim > 0 && len(got) != dim {
|
||||
t.Fatalf("%s 维度 = %d,期望 %d", name, len(got), dim)
|
||||
}
|
||||
var norm float64
|
||||
for _, v := range got {
|
||||
norm += v * v
|
||||
}
|
||||
if diff := math.Abs(math.Sqrt(norm) - 1); diff > 1e-6 {
|
||||
t.Errorf("%s L2 norm = %.9f,期望 1", name, math.Sqrt(norm))
|
||||
}
|
||||
}
|
||||
|
||||
func assertFrozenPrefix(t *testing.T, name string, got, want []float64) {
|
||||
t.Helper()
|
||||
if len(got) < len(want) {
|
||||
t.Fatalf("%s 向量过短: %d", name, len(got))
|
||||
}
|
||||
for i := range want {
|
||||
if diff := math.Abs(got[i] - want[i]); diff > 2e-5 {
|
||||
t.Errorf("%s 维度 %d = %.10g,参考 %.10g,差 %.3g", name, i, got[i], want[i], diff)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestEmbedderMatchesONNXReference 逐维对比导出脚本写出的冻结参考向量,
|
||||
// 验证完整 Go 路径:模板渲染 → BPE → TokenEmbedding → Transformer →
|
||||
// last-token 池化 → L2 normalize。
|
||||
//
|
||||
// 只覆盖前若干维不是因为放宽正确性(脚本侧的 PyTorch↔ONNX 校验是逐维的),
|
||||
// 而是避免把 2048 个浮点常量塞进仓库;这里负责捕获 Go 张量形状、输入名、
|
||||
// 输出名、池化/归一化或模板接线错误。
|
||||
func TestEmbedderMatchesONNXReference(t *testing.T) {
|
||||
dir := requireONNXArtifacts(t)
|
||||
e := newTestEmbedder(t, dir)
|
||||
ref := loadReference(t, dir)
|
||||
|
||||
ids, _, _, _, err := e.tok.textModelInput(e.config.Instruction, ref.Text, e.config.MaxLength)
|
||||
if err != nil {
|
||||
t.Fatalf("textModelInput: %v", err)
|
||||
}
|
||||
postID, ok := e.tok.SpecialID("<|endoftext|>")
|
||||
if !ok || ids[len(ids)-1] != postID {
|
||||
t.Fatalf("模型输入 post-processor 异常: len=%d tail=%v postID=%d ok=%v", len(ids), ids[len(ids)-1:], postID, ok)
|
||||
}
|
||||
|
||||
got, err := e.VectorizeDense(ref.Text)
|
||||
if err != nil {
|
||||
t.Fatalf("VectorizeDense: %v", err)
|
||||
}
|
||||
assertNormalized(t, "text", got, ref.Dim)
|
||||
assertFrozenPrefix(t, "text", got, ref.TextVectorPrefix)
|
||||
}
|
||||
|
||||
// TestEmbedderImageMatchesONNXReference 冻结一张纯色图的参考向量,
|
||||
// 验证 Go 侧的视觉预处理 + patch 排列 + 视觉注入 + 语言模型与 Python 参考一致。
|
||||
func TestEmbedderImageMatchesONNXReference(t *testing.T) {
|
||||
dir := requireONNXArtifacts(t)
|
||||
e := newTestEmbedder(t, dir)
|
||||
ref := loadReference(t, dir)
|
||||
|
||||
img := solidPNG(t, ref.ImageSize, uint8(ref.ImageRGB[0]), uint8(ref.ImageRGB[1]), uint8(ref.ImageRGB[2]))
|
||||
got, err := e.EmbedImageDense(img, "image/png")
|
||||
if err != nil {
|
||||
t.Fatalf("EmbedImageDense: %v", err)
|
||||
}
|
||||
assertNormalized(t, "image", got, ref.Dim)
|
||||
assertFrozenPrefix(t, "image", got, ref.ImageVectorPrefix)
|
||||
}
|
||||
|
||||
// TestEmbedderTextIsSensitiveToInput 阴性对照:冻结向量必须真的随输入变化。
|
||||
//
|
||||
// 没有这条对照,一个「永远返回同一向量」的错误实现也能通过上面的冻结回归
|
||||
// (只要那个常量恰好等于参考值)。这里验证不同文本给出不同向量,且相似文本
|
||||
// 的余弦高于无关文本——即嵌入确实携带语义,而不是常量。
|
||||
func TestEmbedderTextIsSensitiveToInput(t *testing.T) {
|
||||
dir := requireONNXArtifacts(t)
|
||||
e := newTestEmbedder(t, dir)
|
||||
|
||||
base, err := e.VectorizeDense("今天天气怎么样")
|
||||
if err != nil {
|
||||
t.Fatalf("VectorizeDense: %v", err)
|
||||
}
|
||||
same, err := e.VectorizeDense("今天天气怎么样")
|
||||
if err != nil {
|
||||
t.Fatalf("VectorizeDense: %v", err)
|
||||
}
|
||||
if cosine(base, same) < 0.999999 {
|
||||
t.Errorf("同一输入两次嵌入不一致: cos=%.9f(ONNX 会话被并发复用或存在非确定性)", cosine(base, same))
|
||||
}
|
||||
|
||||
other, err := e.VectorizeDense("数据库索引的选择性是怎么计算的")
|
||||
if err != nil {
|
||||
t.Fatalf("VectorizeDense: %v", err)
|
||||
}
|
||||
if cosine(base, other) > 0.999 {
|
||||
t.Errorf("无关文本的余弦高达 %.6f,嵌入可能是常量", cosine(base, other))
|
||||
}
|
||||
}
|
||||
|
||||
// TestEmbedderImageMatchesONNXReference 的替代:不依赖冻结参考的不变量检查。
|
||||
//
|
||||
// 即使参考文件缺失(没有导出产物)或未重新生成,这些不变量也应成立:
|
||||
// 图像路径必须真的走了视觉塔,且不同图片给出不同坐标。
|
||||
func TestEmbedderImageDiffersFromText(t *testing.T) {
|
||||
dir := requireONNXArtifacts(t)
|
||||
e := newTestEmbedder(t, dir)
|
||||
|
||||
imgVec, err := e.EmbedImageDense(solidPNG(t, qwenImageSize, 200, 30, 30), "image/png")
|
||||
if err != nil {
|
||||
t.Fatalf("EmbedImageDense: %v", err)
|
||||
}
|
||||
txtVec, err := e.VectorizeDense(DefaultInstruction)
|
||||
if err != nil {
|
||||
t.Fatalf("VectorizeDense: %v", err)
|
||||
}
|
||||
if cosine(imgVec, txtVec) > 0.999 {
|
||||
t.Error("图像向量与文本向量几乎相同,视觉塔可能没被真正执行")
|
||||
}
|
||||
|
||||
// 不同颜色的图必须给出不同向量
|
||||
blue, err := e.EmbedImageDense(solidPNG(t, qwenImageSize, 30, 150, 220), "image/png")
|
||||
if err != nil {
|
||||
t.Fatalf("EmbedImageDense: %v", err)
|
||||
}
|
||||
if cosine(imgVec, blue) > 0.999999 {
|
||||
t.Error("不同图片给出相同向量,视觉路径未生效")
|
||||
}
|
||||
}
|
||||
|
||||
// TestEmbedderVideoMatchesONNXReference 冻结一个视频用例的参考向量,
|
||||
// 验证 Go 侧完整的视频路径:多帧预处理(时间组布局)→ M-RoPE(<|video_pad|>)
|
||||
// → 视觉注入 → 语言模型。
|
||||
//
|
||||
// 帧颜色在参考里,用来构造与导出脚本一致的输入;帧顺序(组 g 的 tp0←帧2g、
|
||||
// tp1←帧2g+1)写错时这个测试会失败——而那类错误看图时发现不了。
|
||||
//
|
||||
// ⚠️ 当前**明确未通过**(因此跳过,而不是静默当通过):Go 侧的 video 模板
|
||||
// 与 HuggingFace processor 产出的不相等。已定位的差异:processor 会按时间组
|
||||
// 插入字面时间戳文本,逐 token 实测为
|
||||
//
|
||||
// <|vision_start|> <0.0 seconds> <|vision_start|> {576×video_pad} <|vision_end|>
|
||||
// <1.0 seconds> <|vision_start|> {576×video_pad} <|vision_end|>
|
||||
//
|
||||
// 而 Go 侧只生成 <|vision_start|>{G×576 pads}<|vision_end|>。实测同一输入
|
||||
// 下 Python seq=1190(1152 视觉 + 38 文本)、Go 侧只有 22 个文本 token。
|
||||
// 时间戳文本也会占用 M-RoPE 位置,因此 TestVideoModelInputMRope 的自洽断言
|
||||
// 虽然通过,也不能证明与官方实现一致。
|
||||
//
|
||||
// 修复位置在**本 provider 内部**(模型专属模板本就属于这里,不属于核心):
|
||||
// 按 processor 的规则生成同样的分组时间戳文本,然后取消本跳过。
|
||||
func TestEmbedderVideoMatchesONNXReference(t *testing.T) {
|
||||
dir := requireONNXArtifacts(t)
|
||||
e := newTestEmbedder(t, dir)
|
||||
ref := loadReference(t, dir)
|
||||
|
||||
if ref.VideoGroups < 2 || len(ref.VideoFrameRGB) != 2*ref.VideoGroups {
|
||||
// 产物声明了视频支持、参考里却没有视频用例 → 参考没跟上产物,这是缺陷。
|
||||
// 只有「产物本来就不含视频」才允许跳过。
|
||||
if artifactDeclaresVideo(t, dir) {
|
||||
t.Fatalf("产物声明支持原生视频,但参考缺少视频用例(video_groups=%d frames=%d):"+
|
||||
"参考与产物不匹配,请重跑导出脚本的 --verify-only",
|
||||
ref.VideoGroups, len(ref.VideoFrameRGB))
|
||||
}
|
||||
t.Skipf("产物不含原生视频(video_groups=%d),跳过视频回归", ref.VideoGroups)
|
||||
}
|
||||
|
||||
// 产物确实带视频用例:说明我们应当能验证。但 Go 侧模板尚未复现 processor
|
||||
// 的分组时间戳,现在跑必然失败。显式跳过并说明原因,避免出现
|
||||
// 「测试通过」与「视频实际未验证」混为一谈。
|
||||
if ref.VideoGroups > 0 {
|
||||
t.Skip("已知未修复:Go 侧 video 模板缺少 processor 插入的分组时间戳文本" +
|
||||
"(详见本测试注释);修复前视频冻结回归不得视为已验证")
|
||||
}
|
||||
|
||||
frames := make([][]byte, len(ref.VideoFrameRGB))
|
||||
for i, rgb := range ref.VideoFrameRGB {
|
||||
if len(rgb) != 3 {
|
||||
t.Fatalf("帧 %d 颜色字段异常: %v", i, rgb)
|
||||
}
|
||||
frames[i] = solidPNG(t, ref.ImageSize, uint8(rgb[0]), uint8(rgb[1]), uint8(rgb[2]))
|
||||
}
|
||||
|
||||
got, err := e.EmbedVideoDense(frames, "video/mp4")
|
||||
if err != nil {
|
||||
t.Fatalf("EmbedVideoDense: %v", err)
|
||||
}
|
||||
assertNormalized(t, "video", got, ref.Dim)
|
||||
assertFrozenPrefix(t, "video", got, ref.VideoVectorPrefix)
|
||||
}
|
||||
|
||||
// TestVideoModelInputMRope 逐 token 校验视频的 M-RoPE 位置。
|
||||
//
|
||||
// 对应 transformers 的 get_rope_index:它先把 video_grid_thw 按 grid_t 展开成
|
||||
// G 个 (1,h,w) 的 grid 项,每项单独算位置,项间 current_pos 前进
|
||||
// max(h,w)/spatial_merge。位置算错不会报错,只是嵌入慢慢变差,所以必须逐项验。
|
||||
func TestVideoModelInputMRope(t *testing.T) {
|
||||
dir := requireONNXArtifacts(t)
|
||||
e := newTestEmbedder(t, dir)
|
||||
|
||||
const groups = 3
|
||||
ids, _, position, visual, err := e.tok.videoModelInput("", groups, e.config.MaxLength)
|
||||
if err != nil {
|
||||
t.Fatalf("videoModelInput: %v", err)
|
||||
}
|
||||
seq := len(ids)
|
||||
|
||||
// 模板必须以 <|video_pad|> 填充(用成 <|image_pad|> 不会报错,只会错模态)。
|
||||
videoPad, ok := e.tok.SpecialID("<|video_pad|>")
|
||||
if !ok {
|
||||
t.Fatal("tokenizer 缺少 <|video_pad|>")
|
||||
}
|
||||
imagePad, _ := e.tok.SpecialID("<|image_pad|>")
|
||||
wantVisual := groups * qwenVisualTokens
|
||||
count := 0
|
||||
for i, id := range ids {
|
||||
if visual[i] {
|
||||
count++
|
||||
if id != videoPad {
|
||||
t.Fatalf("第 %d 个视觉 token id=%d,期望 video_pad=%d(image_pad=%d)", i, id, videoPad, imagePad)
|
||||
}
|
||||
}
|
||||
}
|
||||
if count != wantVisual {
|
||||
t.Fatalf("视觉 token 数 = %d,期望 %d", count, wantVisual)
|
||||
}
|
||||
|
||||
start := -1
|
||||
for i, v := range visual {
|
||||
if v {
|
||||
start = i
|
||||
break
|
||||
}
|
||||
}
|
||||
if start < 0 {
|
||||
t.Fatal("找不到视觉区间")
|
||||
}
|
||||
// 视觉区间必须连续(中间不能夹文本 token)。
|
||||
for i := start; i < start+wantVisual; i++ {
|
||||
if !visual[i] {
|
||||
t.Fatalf("视觉区间在 %d 处断裂", i)
|
||||
}
|
||||
}
|
||||
if start+wantVisual < seq && visual[start+wantVisual] {
|
||||
t.Fatal("视觉区间超出期望长度")
|
||||
}
|
||||
|
||||
// 视觉之前的文本 token 数就是 M-RoPE 的起始位置。
|
||||
base0 := int64(start)
|
||||
for g := 0; g < groups; g++ {
|
||||
base := base0 + int64(g*qwenVisionScale)
|
||||
for j := 0; j < qwenVisualTokens; j++ {
|
||||
i := start + g*qwenVisualTokens + j
|
||||
wantT := base
|
||||
wantH := base + int64(j/qwenVisionScale)
|
||||
wantW := base + int64(j%qwenVisionScale)
|
||||
if position[i] != wantT || position[seq+i] != wantH || position[2*seq+i] != wantW {
|
||||
t.Fatalf("组%d 第%d 个视觉 token 位置 = (%d,%d,%d),期望 (%d,%d,%d)",
|
||||
g, j, position[i], position[seq+i], position[2*seq+i], wantT, wantH, wantW)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestVideoInputRejectsUnsupportedShapes 帧数与档位不匹配时必须明确报错,
|
||||
// 而不是悄悄补齐/截断成另一个语义。
|
||||
func TestVideoInputRejectsUnsupportedShapes(t *testing.T) {
|
||||
if _, _, err := preprocessVideoFrames([][]byte{solidPNG(t, qwenImageSize, 1, 2, 3)}); err == nil {
|
||||
t.Error("单帧无法构成一个时间组,应报错")
|
||||
}
|
||||
many := make([][]byte, 2*(maxVideoGroupsSafety+1))
|
||||
if _, _, err := preprocessVideoFrames(many); err == nil {
|
||||
t.Errorf("超过分配安全上限 %d 应报错,而不是静默分配巨量内存", maxVideoGroupsSafety)
|
||||
}
|
||||
|
||||
dir := requireONNXArtifacts(t)
|
||||
e := newTestEmbedder(t, dir)
|
||||
if _, _, _, _, err := e.tok.visionModelInput("", "<|video_pad|>", 0, e.config.MaxLength); err == nil {
|
||||
t.Error("groups=0 应报错")
|
||||
}
|
||||
|
||||
// 未导出的档位必须明确报错并告知已加载哪些档,而不是默默找一个相近的。
|
||||
if _, err := e.EmbedVideoDense(framesOf(t, 2*(maxExportedGroupsInTest+1)), "video/mp4"); err == nil {
|
||||
t.Errorf("未导出的 G=%d 应报错", maxExportedGroupsInTest+1)
|
||||
}
|
||||
}
|
||||
|
||||
// maxExportedGroupsInTest 是测试环境预期导出的视频最大档(与导出脚本默认 2,3,4 一致)。
|
||||
const maxExportedGroupsInTest = 4
|
||||
|
||||
func framesOf(t *testing.T, n int) [][]byte {
|
||||
t.Helper()
|
||||
out := make([][]byte, n)
|
||||
for i := range out {
|
||||
out[i] = solidPNG(t, qwenImageSize, uint8(i), 100, 150)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// TestEmbedderRejectsUnsupportedModalities 音频必须显式报「不在本空间」。
|
||||
//
|
||||
// Qwen3-VL 模型卡与 config 双重确认无 audio_token_id;音频需要另一个真正的
|
||||
// 音频模型。若这里退化成普通错误,调用方会把它当「本次失败、下次重试」,
|
||||
// 于是每轮启动都重试一批永远不可能成功的条目。
|
||||
func TestEmbedderRejectsUnsupportedModalities(t *testing.T) {
|
||||
dir := requireONNXArtifacts(t)
|
||||
e := newTestEmbedder(t, dir)
|
||||
|
||||
for _, mime := range []string{"audio/wav", "audio/mpeg"} {
|
||||
_, err := e.EmbedImageDense([]byte("not-a-real-media"), mime)
|
||||
if err == nil {
|
||||
t.Fatalf("%s 应返回错误而不是造出向量", mime)
|
||||
}
|
||||
if !errors.Is(err, embedding.ErrUnsupportedModality) {
|
||||
t.Errorf("%s 错误应为 ErrUnsupportedModality,实际: %v", mime, err)
|
||||
}
|
||||
// 公共 SPI 路径也必须给出可识别的不支持信号。
|
||||
if _, err := e.Embed(context.Background(), embedding.Input{
|
||||
Modality: embedding.ModalityAudio, Data: []byte("x"), MIME: mime,
|
||||
}); !errors.Is(err, embedding.ErrUnsupportedModality) {
|
||||
t.Errorf("Embed(audio/%s) 应为 ErrUnsupportedModality,实际: %v", mime, err)
|
||||
}
|
||||
}
|
||||
|
||||
// 视频**文件**不能直接喂给单帧入口(Go 侧没有视频解码器),
|
||||
// 必须由调用方先抽帧再走 EmbedVideoDense。
|
||||
if _, err := e.EmbedImageDense([]byte("not-a-real-media"), "video/mp4"); !errors.Is(err, embedding.ErrUnsupportedModality) {
|
||||
t.Errorf("EmbedImageDense(video/mp4) 应为 ErrUnsupportedModality,实际: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func cosine(a, b []float64) float64 {
|
||||
if len(a) != len(b) || len(a) == 0 {
|
||||
return 0
|
||||
}
|
||||
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 0
|
||||
}
|
||||
return dot / (math.Sqrt(na) * math.Sqrt(nb))
|
||||
}
|
||||
|
||||
// TestProviderOpensThroughPublicSPI 走完核心真正的路径:
|
||||
// embedding.Open(名字) → 工厂 → Info 校验 → Embed。
|
||||
//
|
||||
// 它与 newTestEmbedder 的区别很关键:后者直接调 New(),只能证明「模型能加载」;
|
||||
// 本测试证明**注册表 + 公共契约**这条链路是通的——名字对得上、工厂能构造、
|
||||
// Info 满足契约、Embed 返回合法向量。核心升级后真正会走的就是这条路由。
|
||||
func TestProviderOpensThroughPublicSPI(t *testing.T) {
|
||||
dir := requireONNXArtifacts(t)
|
||||
|
||||
names := embedding.Names()
|
||||
found := false
|
||||
for _, n := range names {
|
||||
if n == "qwen3vl" {
|
||||
found = true
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Fatalf("qwen3vl 未注册到公共注册表;已注册: %v", names)
|
||||
}
|
||||
|
||||
provider, err := embedding.Open("qwen3vl", embedding.Config{
|
||||
Options: map[string]string{"model_dir": dir},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("embedding.Open(qwen3vl): %v", err)
|
||||
}
|
||||
defer provider.Close()
|
||||
|
||||
info := provider.Info()
|
||||
if info.Dimension != 2048 || info.Fingerprint == "" {
|
||||
t.Fatalf("Info 异常: dim=%d fp=%q", info.Dimension, info.Fingerprint)
|
||||
}
|
||||
// 公共契约路径只声明 text/image:本 provider 没有视频解码器,
|
||||
// 若这里出现 video 就意味着核心会创建一条注定失败的输入通道。
|
||||
for _, m := range info.Modalities {
|
||||
if m == embedding.ModalityVideo {
|
||||
t.Fatal("Info 不应声明 video(provider 无视频解码器,见文档)")
|
||||
}
|
||||
}
|
||||
|
||||
vec, err := provider.Embed(context.Background(), embedding.Input{
|
||||
Modality: embedding.ModalityText, Purpose: embedding.PurposeQuery, Text: "hello",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Embed(text): %v", err)
|
||||
}
|
||||
if err := embedding.ValidateVector(vec, info.Dimension); err != nil {
|
||||
t.Fatalf("返回向量不合法: %v", err)
|
||||
}
|
||||
|
||||
// 未知模态必须给出可识别的「本空间不支持」,而不是普通错误。
|
||||
_, err = provider.Embed(context.Background(), embedding.Input{
|
||||
Modality: embedding.ModalityAudio, Data: []byte("x"), MIME: "audio/wav",
|
||||
})
|
||||
if !errors.Is(err, embedding.ErrUnsupportedModality) {
|
||||
t.Fatalf("audio 应为 ErrUnsupportedModality,实际: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestOpenRejectsProviderWithoutModelDir 未配置 model_dir 时必须是明确的构造失败,
|
||||
// 而不是构造成功、每次 Embed 才报错(那会让启动日志看起来正常)。
|
||||
func TestOpenRejectsProviderWithoutModelDir(t *testing.T) {
|
||||
if _, err := embedding.Open("qwen3vl", embedding.Config{}); err == nil {
|
||||
t.Fatal("缺 model_dir 时应打开失败")
|
||||
}
|
||||
}
|
||||
28
providers/qwen3vl/embedder_stub.go
Normal file
28
providers/qwen3vl/embedder_stub.go
Normal file
@ -0,0 +1,28 @@
|
||||
//go:build !onnxruntime
|
||||
|
||||
package qwen3vl
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
|
||||
"gitcode.com/JianFeeeee/HomeAgent/pkg/embedding"
|
||||
)
|
||||
|
||||
func init() {
|
||||
embedding.Register("qwen3vl", func(embedding.Config) (embedding.Provider, error) {
|
||||
return nil, fmt.Errorf("qwen3vl provider requires build tag 'onnxruntime' (go build -tags onnxruntime)")
|
||||
})
|
||||
}
|
||||
|
||||
// Embedder 在未启用 onnxruntime 时不可用;保留类型是为了让引用它的代码在
|
||||
// 默认构建下也能编译。真正的 ONNX 实现见 embedder_onnx.go。
|
||||
type Embedder struct{}
|
||||
|
||||
func (e *Embedder) Embed(context.Context, embedding.Input) ([]float64, error) {
|
||||
return nil, errors.New("qwen3vl provider not available in this build")
|
||||
}
|
||||
|
||||
func (e *Embedder) Info() embedding.Info { return embedding.Info{} }
|
||||
func (e *Embedder) Close() {}
|
||||
218
providers/qwen3vl/image.go
Normal file
218
providers/qwen3vl/image.go
Normal file
@ -0,0 +1,218 @@
|
||||
//go:build onnxruntime
|
||||
|
||||
package qwen3vl
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"fmt"
|
||||
"image"
|
||||
"image/color"
|
||||
_ "image/gif"
|
||||
_ "image/jpeg"
|
||||
_ "image/png"
|
||||
"math"
|
||||
)
|
||||
|
||||
const (
|
||||
qwenImageSize = 768
|
||||
qwenPatchSize = 16
|
||||
qwenTemporalPatch = 2
|
||||
qwenSpatialMerge = 2
|
||||
qwenImagePatches = (qwenImageSize / qwenPatchSize) * (qwenImageSize / qwenPatchSize)
|
||||
qwenVisualTokens = qwenImagePatches / (qwenSpatialMerge * qwenSpatialMerge)
|
||||
qwenPatchVectorSize = 3 * qwenTemporalPatch * qwenPatchSize * qwenPatchSize
|
||||
|
||||
// qwenVisionScale 是视觉塔空间步长:grid_h / spatial_merge。
|
||||
// 每个时间组消耗这么多 M-RoPE 位置(见 model_input.go 的说明)。
|
||||
qwenVisionScale = (qwenImageSize / qwenPatchSize) / qwenSpatialMerge
|
||||
|
||||
// maxVideoGroupsSafety 是分配安全上限,**不是**能力上限。
|
||||
// 真正能导出哪些档由产物目录决定(Vision_g{N}.onnx);内核不硬编码
|
||||
// 导出清单,否则别人导出 G=8 就会被内核莫名拒绝。
|
||||
// 这个上限只用来防住「丢了上千帧进来」导致的巨量分配。
|
||||
maxVideoGroupsSafety = 64
|
||||
)
|
||||
|
||||
// fitCanvas 把任意图片解码并转成固定 768×768 画布。
|
||||
//
|
||||
// 保持宽高比缩放并在中心补中性灰(归一化后约为 0);直接强拉成正方形会
|
||||
// 破坏物体形状。已是 768×768 的输入不做插值,以便用跨语言冻结向量
|
||||
// 精确回归 patch 排列。
|
||||
func fitCanvas(raw []byte) (*image.NRGBA, error) {
|
||||
src, _, err := image.Decode(bytes.NewReader(raw))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("qwen: decode image: %w", err)
|
||||
}
|
||||
b := src.Bounds()
|
||||
if b.Dx() <= 0 || b.Dy() <= 0 {
|
||||
return nil, fmt.Errorf("qwen: empty image")
|
||||
}
|
||||
|
||||
scale := math.Min(float64(qwenImageSize)/float64(b.Dx()), float64(qwenImageSize)/float64(b.Dy()))
|
||||
w := max(1, int(math.Round(float64(b.Dx())*scale)))
|
||||
h := max(1, int(math.Round(float64(b.Dy())*scale)))
|
||||
if w > qwenImageSize {
|
||||
w = qwenImageSize
|
||||
}
|
||||
if h > qwenImageSize {
|
||||
h = qwenImageSize
|
||||
}
|
||||
|
||||
resized := resizeBicubic(src, w, h)
|
||||
canvas := image.NewNRGBA(image.Rect(0, 0, qwenImageSize, qwenImageSize))
|
||||
neutral := color.NRGBA{R: 128, G: 128, B: 128, A: 255}
|
||||
for i := 0; i < len(canvas.Pix); i += 4 {
|
||||
canvas.Pix[i], canvas.Pix[i+1], canvas.Pix[i+2], canvas.Pix[i+3] = neutral.R, neutral.G, neutral.B, neutral.A
|
||||
}
|
||||
ox, oy := (qwenImageSize-w)/2, (qwenImageSize-h)/2
|
||||
for y := 0; y < h; y++ {
|
||||
for x := 0; x < w; x++ {
|
||||
canvas.SetNRGBA(ox+x, oy+y, resized.NRGBAAt(x, y))
|
||||
}
|
||||
}
|
||||
return canvas, nil
|
||||
}
|
||||
|
||||
// appendPatches 按 Qwen2VLImageProcessor 的排列把一个时间组的两个画布写入 out。
|
||||
//
|
||||
// 排列:[grid_h/merge, grid_w/merge, merge_h, merge_w, channel,
|
||||
// temporal_patch, patch_h, patch_w],然后 flatten。
|
||||
//
|
||||
// 图像与视频共用本函数:图像的两个时间槽传同一张画布,视频传相邻两帧。
|
||||
// 共用是刻意的——两处各写一份排列,迟早会在某次修改后漂移,
|
||||
// 而排列错了只会得到一个语义偏移的向量,不会报错。
|
||||
func appendPatches(out []float32, slots *[qwenTemporalPatch]*image.NRGBA) []float32 {
|
||||
blocks := qwenImageSize / qwenPatchSize / qwenSpatialMerge
|
||||
for bh := 0; bh < blocks; bh++ {
|
||||
for bw := 0; bw < blocks; bw++ {
|
||||
for mh := 0; mh < qwenSpatialMerge; mh++ {
|
||||
for mw := 0; mw < qwenSpatialMerge; mw++ {
|
||||
baseY := (bh*qwenSpatialMerge + mh) * qwenPatchSize
|
||||
baseX := (bw*qwenSpatialMerge + mw) * qwenPatchSize
|
||||
for c := 0; c < 3; c++ {
|
||||
for temporal := 0; temporal < qwenTemporalPatch; temporal++ {
|
||||
canvas := slots[temporal]
|
||||
for py := 0; py < qwenPatchSize; py++ {
|
||||
for px := 0; px < qwenPatchSize; px++ {
|
||||
p := canvas.NRGBAAt(baseX+px, baseY+py)
|
||||
v := [3]uint8{p.R, p.G, p.B}[c]
|
||||
out = append(out, float32(v)/127.5-1)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// preprocessImage 把任意图片转成固定 768×768 视觉塔输入(单个时间组)。
|
||||
func preprocessImage(raw []byte) ([]float32, error) {
|
||||
canvas, err := fitCanvas(raw)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
slots := [qwenTemporalPatch]*image.NRGBA{canvas, canvas}
|
||||
out := make([]float32, 0, qwenImagePatches*qwenPatchVectorSize)
|
||||
return appendPatches(out, &slots), nil
|
||||
}
|
||||
|
||||
// preprocessVideoFrames 把已按时间排序的帧转成 G 个时间组的视觉塔输入,
|
||||
// 返回 patch 张量与时间组数 G。
|
||||
//
|
||||
// 时间组 g 的两个时间槽依次取帧 2g 与 2g+1,这与处理器实测逐字节一致
|
||||
// (纯色与异色两组对照均 torch.equal 通过);布局整体是
|
||||
// [G, blocks_h, blocks_w, merge_h, merge_w, c, temporal, patch_h, patch_w],
|
||||
// 即图像排列以 grid_t 为最外层堆叠。
|
||||
//
|
||||
// 帧数为奇数时不补帧:只用得上的帧参与编码,多余的一帧被丢弃,
|
||||
// 以免用重复帧伪造时序——那会改变跨帧注意力看到的运动。
|
||||
func preprocessVideoFrames(frames [][]byte) ([]float32, int, error) {
|
||||
if len(frames) < qwenTemporalPatch {
|
||||
return nil, 0, fmt.Errorf("qwen: video needs at least %d frames, got %d", qwenTemporalPatch, len(frames))
|
||||
}
|
||||
groups := len(frames) / qwenTemporalPatch
|
||||
if groups > maxVideoGroupsSafety {
|
||||
return nil, 0, fmt.Errorf("qwen: video groups %d exceeds safety limit %d(请先对帧采样)", groups, maxVideoGroupsSafety)
|
||||
}
|
||||
canvases := make([]*image.NRGBA, groups*qwenTemporalPatch)
|
||||
for i := 0; i < groups*qwenTemporalPatch; i++ {
|
||||
canvas, err := fitCanvas(frames[i])
|
||||
if err != nil {
|
||||
return nil, 0, fmt.Errorf("qwen: frame %d: %w", i, err)
|
||||
}
|
||||
canvases[i] = canvas
|
||||
}
|
||||
out := make([]float32, 0, groups*qwenImagePatches*qwenPatchVectorSize)
|
||||
for g := 0; g < groups; g++ {
|
||||
slots := [qwenTemporalPatch]*image.NRGBA{canvases[2*g], canvases[2*g+1]}
|
||||
out = appendPatches(out, &slots)
|
||||
}
|
||||
return out, groups, nil
|
||||
}
|
||||
|
||||
// resizeBicubic 使用半像素中心的 Catmull-Rom 三次卷积。
|
||||
func resizeBicubic(src image.Image, dstW, dstH int) *image.NRGBA {
|
||||
b := src.Bounds()
|
||||
if b.Dx() == dstW && b.Dy() == dstH {
|
||||
dst := image.NewNRGBA(image.Rect(0, 0, dstW, dstH))
|
||||
for y := 0; y < dstH; y++ {
|
||||
for x := 0; x < dstW; x++ {
|
||||
dst.SetNRGBA(x, y, color.NRGBAModel.Convert(src.At(b.Min.X+x, b.Min.Y+y)).(color.NRGBA))
|
||||
}
|
||||
}
|
||||
return dst
|
||||
}
|
||||
|
||||
dst := image.NewNRGBA(image.Rect(0, 0, dstW, dstH))
|
||||
sx, sy := float64(b.Dx())/float64(dstW), float64(b.Dy())/float64(dstH)
|
||||
for y := 0; y < dstH; y++ {
|
||||
fy := (float64(y)+0.5)*sy - 0.5
|
||||
y0 := int(math.Floor(fy))
|
||||
for x := 0; x < dstW; x++ {
|
||||
fx := (float64(x)+0.5)*sx - 0.5
|
||||
x0 := int(math.Floor(fx))
|
||||
var sum [4]float64
|
||||
var weight float64
|
||||
for j := -1; j <= 2; j++ {
|
||||
wy := cubicWeight(fy - float64(y0+j))
|
||||
yy := min(max(y0+j, 0), b.Dy()-1)
|
||||
for i := -1; i <= 2; i++ {
|
||||
w := wy * cubicWeight(fx-float64(x0+i))
|
||||
xx := min(max(x0+i, 0), b.Dx()-1)
|
||||
p := color.NRGBAModel.Convert(src.At(b.Min.X+xx, b.Min.Y+yy)).(color.NRGBA)
|
||||
sum[0] += float64(p.R) * w
|
||||
sum[1] += float64(p.G) * w
|
||||
sum[2] += float64(p.B) * w
|
||||
sum[3] += float64(p.A) * w
|
||||
weight += w
|
||||
}
|
||||
}
|
||||
if weight == 0 {
|
||||
weight = 1
|
||||
}
|
||||
dst.SetNRGBA(x, y, color.NRGBA{
|
||||
R: clampByte(sum[0] / weight), G: clampByte(sum[1] / weight),
|
||||
B: clampByte(sum[2] / weight), A: clampByte(sum[3] / weight),
|
||||
})
|
||||
}
|
||||
}
|
||||
return dst
|
||||
}
|
||||
|
||||
func cubicWeight(x float64) float64 {
|
||||
x = math.Abs(x)
|
||||
if x <= 1 {
|
||||
return 1.5*x*x*x - 2.5*x*x + 1
|
||||
}
|
||||
if x < 2 {
|
||||
return -0.5*x*x*x + 2.5*x*x - 4*x + 2
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func clampByte(v float64) uint8 {
|
||||
return uint8(min(255, max(0, int(math.Round(v)))))
|
||||
}
|
||||
130
providers/qwen3vl/model_input.go
Normal file
130
providers/qwen3vl/model_input.go
Normal file
@ -0,0 +1,130 @@
|
||||
//go:build onnxruntime
|
||||
|
||||
package qwen3vl
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// 视觉输入(图像与视频)的模型输入构造。
|
||||
//
|
||||
// 图像与视频的模板结构完全一致,只有两点不同:
|
||||
// 1. 占位符:<|image_pad|>(id 151655)vs <|video_pad|>(id 151656);
|
||||
// 2. 时间组数:图像恒为 1 组(576 个视觉 token),视频为 G 组(G×576)。
|
||||
//
|
||||
// 因此两者共用同一个构造器。分开写两份必然漂移,而漂移的表现是
|
||||
// 「嵌入略有不同」——不报错,只是检索慢慢变差。
|
||||
|
||||
// visionModelInput 构造视觉输入的 token 序列与 M-RoPE 位置。
|
||||
//
|
||||
// padToken 是 <|image_pad|> 或 <|video_pad|>;groups 是时间组数。
|
||||
func (t *Tokenizer) visionModelInput(instruction, padToken string, groups, maxLen int) (ids []int, attention, position []int64, visual []bool, err error) {
|
||||
if instruction == "" {
|
||||
instruction = DefaultInstruction
|
||||
}
|
||||
if groups < 1 {
|
||||
return nil, nil, nil, nil, fmt.Errorf("qwen: vision groups must be >= 1, got %d", groups)
|
||||
}
|
||||
if _, ok := t.SpecialID(padToken); !ok {
|
||||
return nil, nil, nil, nil, fmt.Errorf("tokenizer.json 缺少 %s", padToken)
|
||||
}
|
||||
text := "<|im_start|>system\n" + instruction +
|
||||
"<|im_end|>\n<|im_start|>user\n<|vision_start|>" +
|
||||
strings.Repeat(padToken, groups*qwenVisualTokens) +
|
||||
"<|vision_end|><|im_end|>\n<|im_start|>assistant\n"
|
||||
ids, err = t.encodeModelInput(text, maxLen)
|
||||
if err != nil {
|
||||
return nil, nil, nil, nil, err
|
||||
}
|
||||
padID, _ := t.SpecialID(padToken)
|
||||
|
||||
visual = make([]bool, len(ids))
|
||||
attention = make([]int64, len(ids))
|
||||
position = make([]int64, 3*len(ids))
|
||||
for i, id := range ids {
|
||||
attention[i] = 1
|
||||
visual[i] = id == padID
|
||||
}
|
||||
|
||||
current := int64(0)
|
||||
for start := 0; start < len(ids); {
|
||||
isVisual := visual[start]
|
||||
end := start + 1
|
||||
for end < len(ids) && visual[end] == isVisual {
|
||||
end++
|
||||
}
|
||||
if !isVisual {
|
||||
for i := start; i < end; i++ {
|
||||
p := current + int64(i-start)
|
||||
position[i] = p
|
||||
position[len(ids)+i] = p
|
||||
position[2*len(ids)+i] = p
|
||||
}
|
||||
current += int64(end - start)
|
||||
} else {
|
||||
run := end - start
|
||||
if run != groups*qwenVisualTokens {
|
||||
return nil, nil, nil, nil, fmt.Errorf("qwen: %s run=%d, want %d (groups=%d)",
|
||||
padToken, run, groups*qwenVisualTokens, groups)
|
||||
}
|
||||
// 每个时间组独立取位置:t 在组内固定为 base,h/w 在组内递增,
|
||||
// 组间 base 前进一个视觉步长。
|
||||
//
|
||||
// 与 transformers 的实现对应:get_rope_index 对视频先把
|
||||
// video_grid_thw 按 grid_t 展开成 G 个 (1,h,w) 的 grid 项,
|
||||
// 每项单独调用 get_vision_position_ids(current_pos, (1,h,w)),
|
||||
// 然后 current_pos += max(h,w)/spatial_merge。因为每项 t=1,
|
||||
// 其 temporal 分量就等于 current_pos,h/w 从 current_pos 起递增。
|
||||
for g := 0; g < groups; g++ {
|
||||
base := current
|
||||
for j := 0; j < qwenVisualTokens; j++ {
|
||||
i := start + g*qwenVisualTokens + j
|
||||
position[i] = base
|
||||
position[len(ids)+i] = base + int64(j/qwenVisionScale)
|
||||
position[2*len(ids)+i] = base + int64(j%qwenVisionScale)
|
||||
}
|
||||
current += int64(qwenVisionScale)
|
||||
}
|
||||
}
|
||||
start = end
|
||||
}
|
||||
return ids, attention, position, visual, nil
|
||||
}
|
||||
|
||||
// imageModelInput 构造 Qwen3-VL 单图对话模板及对应 M-RoPE 位置。
|
||||
// 固定 768×768 视觉塔产生 576 个合并后的视觉 token。
|
||||
func (t *Tokenizer) imageModelInput(instruction string, maxLen int) (ids []int, attention, position []int64, visual []bool, err error) {
|
||||
return t.visionModelInput(instruction, "<|image_pad|>", 1, maxLen)
|
||||
}
|
||||
|
||||
// videoModelInput 构造 Qwen3-VL 视频对话模板及对应 M-RoPE 位置。
|
||||
//
|
||||
// groups 是时间组数(每组合 2 帧),共 2×groups 帧、groups×576 个视觉 token。
|
||||
// 占位符是 <|video_pad|>(id 151656),与图像的 <|image_pad|> 不同——
|
||||
// 用错占位符不会报错,只会让模型把它当成另一种模态。
|
||||
//
|
||||
// 这里不限制 groups 上限:哪些档位真的可用由产物目录(Vision_g{N}.onnx)决定,
|
||||
// 硬编码一份清单在这里只会与导出脚本漂移。序列过长会因 tokenizer 截断
|
||||
// 而在下面的视觉区间长度校验处明确报错。
|
||||
func (t *Tokenizer) videoModelInput(instruction string, groups, maxLen int) (ids []int, attention, position []int64, visual []bool, err error) {
|
||||
return t.visionModelInput(instruction, "<|video_pad|>", groups, maxLen)
|
||||
}
|
||||
|
||||
// textModelInput 执行完整 tokenizer post_processor,并构造纯文本标准 RoPE 位置。
|
||||
func (t *Tokenizer) textModelInput(instruction, text string, maxLen int) (ids []int, attention, position []int64, visual []bool, err error) {
|
||||
ids, err = t.encodeModelInput(renderInstructionInput(instruction, text), maxLen)
|
||||
if err != nil {
|
||||
return nil, nil, nil, nil, err
|
||||
}
|
||||
attention = make([]int64, len(ids))
|
||||
position = make([]int64, 3*len(ids))
|
||||
visual = make([]bool, len(ids))
|
||||
for i := range ids {
|
||||
attention[i] = 1
|
||||
position[i] = int64(i)
|
||||
position[len(ids)+i] = int64(i)
|
||||
position[2*len(ids)+i] = int64(i)
|
||||
}
|
||||
return ids, attention, position, visual, nil
|
||||
}
|
||||
1376
providers/qwen3vl/testdata/qwen_tokenizer_reference.json
vendored
Normal file
1376
providers/qwen3vl/testdata/qwen_tokenizer_reference.json
vendored
Normal file
File diff suppressed because it is too large
Load Diff
503
providers/qwen3vl/tokenizer.go
Normal file
503
providers/qwen3vl/tokenizer.go
Normal file
@ -0,0 +1,503 @@
|
||||
// Package qwen 实现 Qwen3-VL-Embedding 的字节级 BPE 分词器。
|
||||
//
|
||||
// 为什么不复用 clip 的 tokenizer:CLIP 用的是「小写化 + 空白规整 + 词表 BPE」,
|
||||
// 而千问是 **GPT-2 式字节级 BPE**——先把输入按字节映射到一组可见 unicode,
|
||||
// 再对映射后的字符串做 BPE 合并。两者的预处理不可互换,硬套会在中文和
|
||||
// 空白较多的输入上产出完全不同的 token。
|
||||
//
|
||||
// 与上游(HuggingFace tokenizer.json 的 Rust 实现)对齐时的两处坑:
|
||||
//
|
||||
// 1. pre_tokenizer 正则里的 `\s+(?!\S)` 是**负向前瞻**,Go 的 RE2 不支持
|
||||
// lookaround。该分支只在「空白一直延伸到串尾」时命中,而此时贪婪的
|
||||
// `\s+` 会匹配完全相同的区间,所以直接删掉该分支即为等价改写。
|
||||
// 2. Go 的 `\s` 只覆盖 ASCII,而 Rust regex 的 `\s` 是 Unicode
|
||||
// `\p{White_Space}`。不换成 \p{White_Space} 的话,全角空格、NBSP、
|
||||
// 行分隔符等的切分点会与上游不一致。
|
||||
package qwen3vl
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"strings"
|
||||
"unicode"
|
||||
"unicode/utf8"
|
||||
)
|
||||
|
||||
// 空白判定统一用 unicode.IsSpace(Unicode White_Space 属性)。
|
||||
//
|
||||
// 不能用 Go 正则里的 \s——那只覆盖 ASCII;也不能写 \p{White_Space}——Go 的
|
||||
// regexp 只支持 script/category,不支持二进制属性(会报 invalid character
|
||||
// class range)。上游 Rust regex 的 \s 正是 White_Space,所以这里以
|
||||
// unicode.IsSpace 为准。
|
||||
|
||||
// specialToken 是一个 AddedToken:以整体形式优先匹配,不参与 BPE 拆分。
|
||||
type specialToken struct {
|
||||
content string
|
||||
id int
|
||||
}
|
||||
|
||||
// Tokenizer 是千问的字节级 BPE 分词器。
|
||||
type Tokenizer struct {
|
||||
vocab map[string]int
|
||||
ranks map[string]int
|
||||
|
||||
// byteEnc 是 GPT-2 的 byte→unicode 映射:把 0..255 每个字节映到一个
|
||||
// 「可见且不会与正常文本冲突」的 unicode 码点。因为 BPE 词表基于文本构建,
|
||||
// 直接放原始字节会与合法 UTF-8 冲突。
|
||||
byteEnc map[byte]rune
|
||||
|
||||
// specials 按 content 长度降序,保证「最长优先」——
|
||||
// 否则 `<|im_start|>` 可能被 `<|im_` 之类的短 token 先切走。
|
||||
specials []specialToken
|
||||
|
||||
// MaxLen 是嵌入用途的截断上限(与导出脚本的 MAX_LENGTH 一致)。
|
||||
MaxLen int
|
||||
}
|
||||
|
||||
// tokenizerJSON 只取我们需要的部分。
|
||||
type tokenizerJSON struct {
|
||||
Model struct {
|
||||
Vocab map[string]int `json:"vocab"`
|
||||
Merges []interface{} `json:"merges"`
|
||||
} `json:"model"`
|
||||
AddedTokens []struct {
|
||||
ID int `json:"id"`
|
||||
Content string `json:"content"`
|
||||
Special bool `json:"special"`
|
||||
} `json:"added_tokens"`
|
||||
}
|
||||
|
||||
// LoadTokenizer 从模型目录加载 tokenizer.json。
|
||||
func LoadTokenizer(modelDir string) (*Tokenizer, error) {
|
||||
raw, err := os.ReadFile(filepath.Join(modelDir, "tokenizer.json"))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("read tokenizer.json: %w", err)
|
||||
}
|
||||
var tj tokenizerJSON
|
||||
if err := json.Unmarshal(raw, &tj); err != nil {
|
||||
return nil, fmt.Errorf("parse tokenizer.json: %w", err)
|
||||
}
|
||||
if len(tj.Model.Vocab) == 0 {
|
||||
return nil, fmt.Errorf("tokenizer.json 的 model.vocab 为空")
|
||||
}
|
||||
|
||||
ranks := make(map[string]int, len(tj.Model.Merges))
|
||||
for i, m := range tj.Model.Merges {
|
||||
// merges 有两种形态:字符串 "a b",或数组 ["a","b"]。
|
||||
var pair string
|
||||
switch v := m.(type) {
|
||||
case string:
|
||||
pair = v
|
||||
case []interface{}:
|
||||
if len(v) == 2 {
|
||||
a, _ := v[0].(string)
|
||||
b, _ := v[1].(string)
|
||||
pair = a + " " + b
|
||||
}
|
||||
}
|
||||
if pair != "" {
|
||||
if _, seen := ranks[pair]; !seen {
|
||||
ranks[pair] = i
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
t := &Tokenizer{
|
||||
vocab: tj.Model.Vocab,
|
||||
ranks: ranks,
|
||||
byteEnc: bytesToUnicode(),
|
||||
MaxLen: 512,
|
||||
}
|
||||
for _, at := range tj.AddedTokens {
|
||||
if at.Special && at.Content != "" {
|
||||
t.specials = append(t.specials, specialToken{content: at.Content, id: at.ID})
|
||||
}
|
||||
}
|
||||
// 最长优先,避免短 token 抢走长 token 的前缀。
|
||||
sort.Slice(t.specials, func(i, j int) bool {
|
||||
return len(t.specials[i].content) > len(t.specials[j].content)
|
||||
})
|
||||
return t, nil
|
||||
}
|
||||
|
||||
// VocabSize 返回词表大小(诊断用)。
|
||||
func (t *Tokenizer) VocabSize() int { return len(t.vocab) }
|
||||
|
||||
// SpecialID 返回特殊 token 的 id;不存在时 ok=false。
|
||||
func (t *Tokenizer) SpecialID(content string) (int, bool) {
|
||||
for _, s := range t.specials {
|
||||
if s.content == content {
|
||||
return s.id, true
|
||||
}
|
||||
}
|
||||
return 0, false
|
||||
}
|
||||
|
||||
// DefaultInstruction 是导出脚本随 embed_config.json 写入的默认指令。
|
||||
const DefaultInstruction = "Represent the user's input."
|
||||
|
||||
// renderInstructionInput 按模型自带的对话模板拼输入(无构建标签,便于测试)。
|
||||
//
|
||||
// 必须与 HuggingFace processor 的 apply_chat_template(add_generation_prompt=True)
|
||||
// 产出完全一致:指令放 system、正文放 user、以 assistant 起始符结尾。差一个
|
||||
// 特殊 token,池化取到的「最后一个有效 token」位置就变了,嵌入也就不同——
|
||||
// 而且不会报错。参考数据集里有该模板串的用例,能逐 token 对齐验证。
|
||||
func renderInstructionInput(instruction, text string) string {
|
||||
if instruction == "" {
|
||||
instruction = DefaultInstruction
|
||||
}
|
||||
return "<|im_start|>system\n" + instruction +
|
||||
"<|im_end|>\n<|im_start|>user\n" + text +
|
||||
"<|im_end|>\n<|im_start|>assistant\n"
|
||||
}
|
||||
|
||||
// Encode 把文本编码为 token id 序列(识别输入中已有的特殊 token,
|
||||
// 但不执行 tokenizer.json 的 post_processor,也不做截断)。
|
||||
func (t *Tokenizer) Encode(text string) []int {
|
||||
var ids []int
|
||||
for _, seg := range t.splitSpecials(text) {
|
||||
if seg.specialID >= 0 {
|
||||
ids = append(ids, seg.specialID)
|
||||
continue
|
||||
}
|
||||
ids = append(ids, t.encodeOrdinary(seg.text)...)
|
||||
}
|
||||
return ids
|
||||
}
|
||||
|
||||
// encodeModelInput 执行 TextTower 输入所需的 tokenizer post_processor。
|
||||
//
|
||||
// tokenizer.json 的 TemplateProcessing 规则是 `$A <|endoftext|>`;HuggingFace
|
||||
// 在 truncation=true 时先把 A 截到 maxLen-1,再保留末尾 post token。漏掉它不会
|
||||
// 触发 ONNX 错误,却会改变池化位置和整条嵌入向量,因此不能直接用 Encode 的结果。
|
||||
func (t *Tokenizer) encodeModelInput(text string, maxLen int) ([]int, error) {
|
||||
postID, ok := t.SpecialID("<|endoftext|>")
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("tokenizer.json 缺少 post token <|endoftext|>")
|
||||
}
|
||||
if maxLen <= 0 {
|
||||
return nil, fmt.Errorf("maxLen 必须大于 0")
|
||||
}
|
||||
|
||||
ids := t.Encode(text)
|
||||
if len(ids) >= maxLen {
|
||||
ids = ids[:maxLen-1]
|
||||
}
|
||||
return append(ids, postID), nil
|
||||
}
|
||||
|
||||
// seg 是「普通文本」或「已识别的特殊 token」二选一。
|
||||
type seg struct {
|
||||
text string
|
||||
specialID int // -1 表示普通文本
|
||||
}
|
||||
|
||||
// splitSpecials 把输入切成普通片段与特殊 token 片段。
|
||||
//
|
||||
// 为什么必须先切:`<|im_start|>` 在词表里是一个整体 id(151644),若走 BPE
|
||||
// 会被拆成若干子 token,编码结果与上游不一致,模型看到的输入也就变了。
|
||||
func (t *Tokenizer) splitSpecials(text string) []seg {
|
||||
if len(t.specials) == 0 || text == "" {
|
||||
return []seg{{text: text, specialID: -1}}
|
||||
}
|
||||
var out []seg
|
||||
for len(text) > 0 {
|
||||
// 找最靠前的特殊 token 出现位置(同位置取最长)。
|
||||
bestIdx, bestLen, bestID := -1, 0, -1
|
||||
for _, s := range t.specials {
|
||||
i := strings.Index(text, s.content)
|
||||
if i < 0 {
|
||||
continue
|
||||
}
|
||||
if bestIdx == -1 || i < bestIdx || (i == bestIdx && len(s.content) > bestLen) {
|
||||
bestIdx, bestLen, bestID = i, len(s.content), s.id
|
||||
}
|
||||
}
|
||||
if bestIdx == -1 {
|
||||
out = append(out, seg{text: text, specialID: -1})
|
||||
break
|
||||
}
|
||||
if bestIdx > 0 {
|
||||
out = append(out, seg{text: text[:bestIdx], specialID: -1})
|
||||
}
|
||||
out = append(out, seg{specialID: bestID})
|
||||
text = text[bestIdx+bestLen:]
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// encodeOrdinary 对普通文本做「切分 → 字节映射 → BPE 合并」。
|
||||
func (t *Tokenizer) encodeOrdinary(text string) []int {
|
||||
if text == "" {
|
||||
return nil
|
||||
}
|
||||
var ids []int
|
||||
for _, piece := range t.preTokenize(text) {
|
||||
// 字节级映射:先把 piece 的 UTF-8 字节逐个映射成 unicode 字符。
|
||||
var sb strings.Builder
|
||||
for _, b := range []byte(piece) {
|
||||
sb.WriteRune(t.byteEnc[b])
|
||||
}
|
||||
for _, tok := range t.bpe(sb.String()) {
|
||||
if id, ok := t.vocab[tok]; ok {
|
||||
ids = append(ids, id)
|
||||
}
|
||||
// 词表里找不到的片段直接丢弃:正常情况不会发生
|
||||
//(词表覆盖全部 256 个字节级字符),发生即数据有问题。
|
||||
}
|
||||
}
|
||||
return ids
|
||||
}
|
||||
|
||||
// ---- pre_tokenizer ----
|
||||
//
|
||||
// 上游是一条正则(tokenizer.json 的 pre_tokenizer.pretokenizers[0].pattern):
|
||||
//
|
||||
// (?i:'s|'t|'re|'ve|'m|'ll|'d)|[^\r\n\p{L}\p{N}]?\p{L}+|\p{N}|
|
||||
// ?[^\s\p{L}\p{N}]+[\r\n]*|\s*[\r\n]+|\s+(?!\S)|\s+
|
||||
//
|
||||
// **为什么不用一个 Go 正则**:末两个分支里的 `\s+(?!\S)` 是负向前瞻,RE2 不
|
||||
// 支持 lookaround;而且它的真实语义依赖**回溯**——`\s+` 先贪婪吃完整段空白,
|
||||
// 发现后面是非空白导致 `(?!\S)` 失败,于是回退一个字符,正好留下末尾一个
|
||||
// 空白给前面那些以 ` ?` / `[^…]?` 开头的分支合并。这个“留一个”直接决定
|
||||
// 切分点(`" leading"` 会切成 `" "` + `" leading"` 而不是 `" "` + `"leading"`),
|
||||
// 近似改写必然对不上,所以按分支顺序显式实现。
|
||||
func (t *Tokenizer) preTokenize(text string) []string {
|
||||
var out []string
|
||||
for len(text) > 0 {
|
||||
switch {
|
||||
case matchApostrophe(text) > 0:
|
||||
n := matchApostrophe(text)
|
||||
out = append(out, text[:n])
|
||||
text = text[n:]
|
||||
case matchWord(text) > 0:
|
||||
n := matchWord(text)
|
||||
out = append(out, text[:n])
|
||||
text = text[n:]
|
||||
case matchDigit(text) > 0:
|
||||
n := matchDigit(text)
|
||||
out = append(out, text[:n])
|
||||
text = text[n:]
|
||||
case matchPunct(text) > 0:
|
||||
n := matchPunct(text)
|
||||
out = append(out, text[:n])
|
||||
text = text[n:]
|
||||
case matchNewline(text) > 0:
|
||||
n := matchNewline(text)
|
||||
out = append(out, text[:n])
|
||||
text = text[n:]
|
||||
default:
|
||||
// `\s+(?!\S)|\s+` 合一:空白段。
|
||||
total, lastStart := wsRun(text)
|
||||
if total == 0 {
|
||||
// 兜底:不应到达(分支覆盖全部字符),防御性前进一个 rune。
|
||||
_, size := utf8.DecodeRuneInString(text)
|
||||
out = append(out, text[:size])
|
||||
text = text[size:]
|
||||
continue
|
||||
}
|
||||
n := total
|
||||
if total < len(text) && lastStart > 0 {
|
||||
n = lastStart // 后面还有非空白 → 回退掉末尾那一个空白
|
||||
}
|
||||
out = append(out, text[:n])
|
||||
text = text[n:]
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func runeAt(s string) (rune, int) { return utf8.DecodeRuneInString(s) }
|
||||
|
||||
func isLetter(r rune) bool { return unicode.IsLetter(r) }
|
||||
func isNumber(r rune) bool { return unicode.IsNumber(r) }
|
||||
func isWS(r rune) bool { return unicode.IsSpace(r) }
|
||||
|
||||
// wsRun 返回开头连续空白段的字节长度,以及最后一个空白 rune 的起始字节位置。
|
||||
func wsRun(s string) (total, lastStart int) {
|
||||
lastStart = -1
|
||||
i := 0
|
||||
for i < len(s) {
|
||||
r, size := runeAt(s[i:])
|
||||
if !isWS(r) {
|
||||
break
|
||||
}
|
||||
lastStart = i
|
||||
i += size
|
||||
}
|
||||
return i, lastStart
|
||||
}
|
||||
|
||||
// matchApostrophe:`(?i:'s|'t|'re|'ve|'m|'ll|'d)`
|
||||
func matchApostrophe(s string) int {
|
||||
if len(s) == 0 || s[0] != '\'' {
|
||||
return 0
|
||||
}
|
||||
rest := s[1:]
|
||||
// 各后缀互为前缀关系(re/ve/ll/s/t/m/d),所以先试长的。
|
||||
for _, suf := range []string{"re", "ve", "ll", "s", "t", "m", "d"} {
|
||||
if len(rest) >= len(suf) && strings.EqualFold(rest[:len(suf)], suf) {
|
||||
return 1 + len(suf)
|
||||
}
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
// matchWord:`[^\r\n\p{L}\p{N}]?\p{L}+`
|
||||
//
|
||||
// 注意可选字符**排除** \r \n;若吃了可选字符却没有字母跟上,整个分支失败
|
||||
// (与正则的“该分支不匹配”一致,不能把可选字符当已消耗)。
|
||||
func matchWord(s string) int {
|
||||
i := 0
|
||||
if r, size := runeAt(s); r != '\r' && r != '\n' && !isLetter(r) && !isNumber(r) {
|
||||
i = size
|
||||
}
|
||||
r, size := runeAt(s[i:])
|
||||
if !isLetter(r) {
|
||||
return 0
|
||||
}
|
||||
i += size
|
||||
for i < len(s) {
|
||||
r, size := runeAt(s[i:])
|
||||
if !isLetter(r) {
|
||||
break
|
||||
}
|
||||
i += size
|
||||
}
|
||||
return i
|
||||
}
|
||||
|
||||
// matchDigit:`\p{N}` —— 只吃**一个**数字。
|
||||
func matchDigit(s string) int {
|
||||
if r, size := runeAt(s); isNumber(r) {
|
||||
return size
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
// matchPunct:` ?[^\s\p{L}\p{N}]+[\r\n]*`
|
||||
//
|
||||
// 开头是**字面空格**(不是 \s),所以只可能吃掉一个 U+0020。
|
||||
func matchPunct(s string) int {
|
||||
i := 0
|
||||
if strings.HasPrefix(s, " ") {
|
||||
i = 1
|
||||
}
|
||||
n := 0
|
||||
for i+n < len(s) {
|
||||
r, size := runeAt(s[i+n:])
|
||||
if isWS(r) || isLetter(r) || isNumber(r) {
|
||||
break
|
||||
}
|
||||
n += size
|
||||
}
|
||||
if n == 0 {
|
||||
return 0
|
||||
}
|
||||
i += n
|
||||
for i < len(s) && (s[i] == '\r' || s[i] == '\n') {
|
||||
i++
|
||||
}
|
||||
return i
|
||||
}
|
||||
|
||||
// matchNewline:`\s*[\r\n]+`
|
||||
//
|
||||
// 贪婪+回溯的真实语义:`\s*` 先吃完整段空白,`[\r\n]+` 无可匹配而回退,
|
||||
// 最终停在段内**最后一个** \r 或 \n 之前,再把它之后的连续 \r\n 吃掉。
|
||||
func matchNewline(s string) int {
|
||||
total, _ := wsRun(s)
|
||||
if total == 0 {
|
||||
return 0
|
||||
}
|
||||
last := -1
|
||||
for j := total - 1; j >= 0; j-- {
|
||||
if s[j] == '\r' || s[j] == '\n' {
|
||||
last = j
|
||||
break
|
||||
}
|
||||
}
|
||||
if last < 0 {
|
||||
return 0
|
||||
}
|
||||
end := last
|
||||
for end < len(s) && (s[end] == '\r' || s[end] == '\n') {
|
||||
end++
|
||||
}
|
||||
return end
|
||||
}
|
||||
|
||||
// bpe 是标准字节级 BPE:反复合并 rank 最小的相邻对,直到无可合并。
|
||||
func (t *Tokenizer) bpe(word string) []string {
|
||||
symbols := make([]string, 0, len(word))
|
||||
for _, r := range word {
|
||||
symbols = append(symbols, string(r))
|
||||
}
|
||||
if len(symbols) < 2 {
|
||||
return symbols
|
||||
}
|
||||
|
||||
for {
|
||||
bestRank, bestIdx := -1, -1
|
||||
for i := 0; i+1 < len(symbols); i++ {
|
||||
r, ok := t.ranks[symbols[i]+" "+symbols[i+1]]
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if bestRank == -1 || r < bestRank {
|
||||
bestRank, bestIdx = r, i
|
||||
}
|
||||
}
|
||||
if bestIdx == -1 {
|
||||
return symbols
|
||||
}
|
||||
merged := symbols[bestIdx] + symbols[bestIdx+1]
|
||||
symbols = append(symbols[:bestIdx], append([]string{merged}, symbols[bestIdx+2:]...)...)
|
||||
if len(symbols) < 2 {
|
||||
return symbols
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// bytesToUnicode 是 GPT-2 的字节↔unicode 映射表。
|
||||
//
|
||||
// 让每个字节都有一个「安全」的可见码点表示,避免原始控制字节混进 BPE 词表。
|
||||
// 可打印 ASCII 与拉丁补充区保持原样,其余字节映射到 256 之后的码点。
|
||||
func bytesToUnicode() map[byte]rune {
|
||||
bs := make([]int, 0, 256)
|
||||
for b := int('!'); b <= int('~'); b++ {
|
||||
bs = append(bs, b)
|
||||
}
|
||||
for b := 0xA1; b <= 0xAC; b++ {
|
||||
bs = append(bs, b)
|
||||
}
|
||||
for b := 0xAE; b <= 0xFF; b++ {
|
||||
bs = append(bs, b)
|
||||
}
|
||||
|
||||
inBS := make(map[int]bool, len(bs))
|
||||
for _, b := range bs {
|
||||
inBS[b] = true
|
||||
}
|
||||
|
||||
cs := make([]int, len(bs))
|
||||
copy(cs, bs)
|
||||
n := 0
|
||||
for b := 0; b < 256; b++ {
|
||||
if inBS[b] {
|
||||
continue
|
||||
}
|
||||
bs = append(bs, b)
|
||||
cs = append(cs, 256+n)
|
||||
n++
|
||||
}
|
||||
|
||||
out := make(map[byte]rune, 256)
|
||||
for i, b := range bs {
|
||||
out[byte(b)] = rune(cs[i])
|
||||
}
|
||||
return out
|
||||
}
|
||||
221
providers/qwen3vl/tokenizer_test.go
Normal file
221
providers/qwen3vl/tokenizer_test.go
Normal file
@ -0,0 +1,221 @@
|
||||
package qwen3vl
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// modelDir 是本地千问模型目录。不存在则跳过——参考数据已固化在 testdata,
|
||||
// 但分词器本身要从 tokenizer.json 加载词表与 merges(11MB,不入库)。
|
||||
const modelDir = "/home/newqqagent/models/models/qwen--Qwen3-VL-Embedding-2B/snapshots/master"
|
||||
|
||||
type tokenizerRef struct {
|
||||
VocabSize int `json:"vocab_size"`
|
||||
Cases []struct {
|
||||
Text string `json:"text"`
|
||||
IDs []int `json:"ids"`
|
||||
Tokens []string `json:"tokens"`
|
||||
} `json:"cases"`
|
||||
AddedTokens []struct {
|
||||
Content string `json:"content"`
|
||||
ID int `json:"id"`
|
||||
Special bool `json:"special"`
|
||||
} `json:"added_tokens"`
|
||||
}
|
||||
|
||||
func loadRef(t *testing.T) *tokenizerRef {
|
||||
t.Helper()
|
||||
raw, err := os.ReadFile(filepath.Join("testdata", "qwen_tokenizer_reference.json"))
|
||||
if err != nil {
|
||||
t.Fatalf("读取参考数据: %v", err)
|
||||
}
|
||||
var ref tokenizerRef
|
||||
if err := json.Unmarshal(raw, &ref); err != nil {
|
||||
t.Fatalf("解析参考数据: %v", err)
|
||||
}
|
||||
return &ref
|
||||
}
|
||||
|
||||
func loadTokenizer(t *testing.T) *Tokenizer {
|
||||
t.Helper()
|
||||
if _, err := os.Stat(filepath.Join(modelDir, "tokenizer.json")); err != nil {
|
||||
t.Skipf("模型目录不可用,跳过: %v", err)
|
||||
}
|
||||
tok, err := LoadTokenizer(modelDir)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadTokenizer: %v", err)
|
||||
}
|
||||
return tok
|
||||
}
|
||||
|
||||
// 与 HuggingFace 的真实 tokenizer 逐条对齐。
|
||||
//
|
||||
// 这是本包唯一的正确性判据:字节级 BPE 的失败模式是「看起来能跑但 token 不同」,
|
||||
// 而 token 不同会让模型收到完全不同的输入,嵌入自然也就错了——不会报任何错。
|
||||
// 所以必须拿真实输出对照,不能靠读代码断言。
|
||||
func TestTokenizerMatchesReference(t *testing.T) {
|
||||
ref := loadRef(t)
|
||||
tok := loadTokenizer(t)
|
||||
|
||||
if got := tok.VocabSize(); got != ref.VocabSize {
|
||||
t.Errorf("词表大小 = %d,参考 %d", got, ref.VocabSize)
|
||||
}
|
||||
|
||||
failed := 0
|
||||
for _, c := range ref.Cases {
|
||||
got := tok.Encode(c.Text)
|
||||
if !sameIDs(got, c.IDs) {
|
||||
failed++
|
||||
t.Errorf("不一致 text=%q\n got %v\n want %v", c.Text, got, c.IDs)
|
||||
}
|
||||
}
|
||||
if failed > 0 {
|
||||
t.Fatalf("%d/%d 条用例不一致", failed, len(ref.Cases))
|
||||
}
|
||||
}
|
||||
|
||||
func sameIDs(a, b []int) bool {
|
||||
if len(a) != len(b) {
|
||||
return false
|
||||
}
|
||||
for i := range a {
|
||||
if a[i] != b[i] {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// 特殊 token 必须整体匹配:走 BPE 会被拆成子 token,模型看到的输入就变了。
|
||||
func TestSpecialTokensMatchWhole(t *testing.T) {
|
||||
ref := loadRef(t)
|
||||
tok := loadTokenizer(t)
|
||||
|
||||
for _, at := range ref.AddedTokens {
|
||||
if !at.Special {
|
||||
continue
|
||||
}
|
||||
got, ok := tok.SpecialID(at.Content)
|
||||
if !ok {
|
||||
t.Errorf("特殊 token %q 未从 tokenizer.json 载入", at.Content)
|
||||
continue
|
||||
}
|
||||
if got != at.ID {
|
||||
t.Errorf("特殊 token %q id=%d,参考 %d", at.Content, got, at.ID)
|
||||
}
|
||||
|
||||
// 单独出现时必须编码成恰好一个 id。
|
||||
ids := tok.Encode(at.Content)
|
||||
if len(ids) != 1 || ids[0] != at.ID {
|
||||
t.Errorf("特殊 token %q 应整体编码为 [%d],实际 %v", at.Content, at.ID, ids)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 最长优先:`<|im_start|>` 不能被更短的 `<|im_end|>` 之类前缀抢走。
|
||||
func TestSpecialTokenLongestFirst(t *testing.T) {
|
||||
tok := loadTokenizer(t)
|
||||
text := "<|im_start|>user\n你好<|im_end|>"
|
||||
|
||||
ids := tok.Encode(text)
|
||||
startID, _ := tok.SpecialID("<|im_start|>")
|
||||
endID, _ := tok.SpecialID("<|im_end|>")
|
||||
|
||||
if len(ids) == 0 || ids[0] != startID {
|
||||
t.Fatalf("应以 <|im_start|>(%d) 开头,实际 %v", startID, ids)
|
||||
}
|
||||
if last := ids[len(ids)-1]; last != endID {
|
||||
t.Fatalf("应以 <|im_end|>(%d) 结尾,实际 %v", endID, ids)
|
||||
}
|
||||
}
|
||||
|
||||
// 空串与单字符边界。
|
||||
func TestTokenizerEdgeCases(t *testing.T) {
|
||||
tok := loadTokenizer(t)
|
||||
if got := tok.Encode(""); len(got) != 0 {
|
||||
t.Errorf("空串应产出 0 个 token,实际 %v", got)
|
||||
}
|
||||
for _, s := range []string{"a", "中", "1", " "} {
|
||||
if got := tok.Encode(s); len(got) == 0 {
|
||||
t.Errorf("%q 应至少产出 1 个 token", s)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 模板渲染必须与参考数据里的整串完全一致,且逐 token 对齐。
|
||||
//
|
||||
// 这是嵌入正确性的前提:模板差一个字符,池化取到的「最后一个有效 token」
|
||||
// 位置就变了,向量也就不同——而且不会报错。
|
||||
func TestRenderInstructionInputMatchesTemplate(t *testing.T) {
|
||||
ref := loadRef(t)
|
||||
tok := loadTokenizer(t)
|
||||
|
||||
const want = "<|im_start|>system\nRepresent the user's input.<|im_end|>\n" +
|
||||
"<|im_start|>user\n你好<|im_end|>\n<|im_start|>assistant\n"
|
||||
|
||||
got := renderInstructionInput("", "你好")
|
||||
if got != want {
|
||||
t.Fatalf("模板渲染不一致:\n got %q\n want %q", got, want)
|
||||
}
|
||||
|
||||
for _, c := range ref.Cases {
|
||||
if c.Text != want {
|
||||
continue
|
||||
}
|
||||
if ids := tok.Encode(got); !sameIDs(ids, c.IDs) {
|
||||
t.Fatalf("模板串 token 不一致:\n got %v\n want %v", ids, c.IDs)
|
||||
}
|
||||
return
|
||||
}
|
||||
t.Fatal("参考数据里缺少该模板串用例")
|
||||
}
|
||||
|
||||
// 模型输入还要执行 tokenizer.json 的 TemplateProcessing:末尾追加
|
||||
// <|endoftext|>;超长输入先给正文留 maxLen-1 个位置,再保留 post token。
|
||||
func TestEncodeModelInputPostProcessor(t *testing.T) {
|
||||
tok := loadTokenizer(t)
|
||||
postID, ok := tok.SpecialID("<|endoftext|>")
|
||||
if !ok {
|
||||
t.Fatal("tokenizer 缺少 <|endoftext|>")
|
||||
}
|
||||
|
||||
shortRaw := tok.Encode("你好")
|
||||
short, err := tok.encodeModelInput("你好", 512)
|
||||
if err != nil {
|
||||
t.Fatalf("短文本 encodeModelInput: %v", err)
|
||||
}
|
||||
if len(short) != len(shortRaw)+1 || short[len(short)-1] != postID {
|
||||
t.Fatalf("短文本 post-processor 异常: raw=%v model=%v", shortRaw, short)
|
||||
}
|
||||
|
||||
longRaw := tok.Encode(strings.Repeat("记忆", 600))
|
||||
long, err := tok.encodeModelInput(strings.Repeat("记忆", 600), 512)
|
||||
if err != nil {
|
||||
t.Fatalf("长文本 encodeModelInput: %v", err)
|
||||
}
|
||||
if len(long) != 512 || long[511] != postID {
|
||||
t.Fatalf("长文本截断异常: len=%d tail=%v", len(long), long[len(long)-1:])
|
||||
}
|
||||
if !sameIDs(long[:511], longRaw[:511]) {
|
||||
t.Fatal("长文本正文未按 maxLen-1 截断")
|
||||
}
|
||||
}
|
||||
|
||||
// byteEnc 必须是双射:256 个字节映射到 256 个互不相同的码点。
|
||||
// 有碰撞就会让不同字节编成同一个 token,静默产生错误输入。
|
||||
func TestBytesToUnicodeBijective(t *testing.T) {
|
||||
m := bytesToUnicode()
|
||||
if len(m) != 256 {
|
||||
t.Fatalf("映射应覆盖 256 个字节,实际 %d", len(m))
|
||||
}
|
||||
seen := map[rune]byte{}
|
||||
for b, r := range m {
|
||||
if prev, dup := seen[r]; dup {
|
||||
t.Fatalf("码点冲突:字节 %d 与 %d 都映射到 %q", prev, b, r)
|
||||
}
|
||||
seen[r] = b
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user