MLX-VLM 中的 DINOv2 移植:Channel-Last 视觉骨干网络详解与实战
【免费下载链接】mlx-vlmMLX-VLM is a package for inference and fine-tuning of Vision Language Models (VLMs) on your Mac using MLX.项目地址: https://gitcode.com/GitHub_Trending/ml/mlx-vlm
DINOv2 是 Meta AI(facebookresearch)发布的基于自监督学习的视觉 Transformer(ViT)模型系列,本仓库在 MLX 框架下将其移植为 channel-last((B, H, W, C))实现,作为 MLX-VLM 的一个独立图像编码器模块,同时也作为 Video Depth Anything、MoGe-3 等稠密预测模型的共享骨干网络。阅读本文后,你将掌握如何在 Apple Silicon 上直接加载官方 DINOv2 权重进行图像特征提取、理解其配置字段与位置编码插值策略,并能把它作为视觉骨干复用到下游稠密预测任务中。
一、模块定位与包内组成
本模块位于 mlx_vlm/models/dinov2/,是一个基于 MLX 的 channel-last DINOv2 vision transformer 移植。与常见 PyTorch 实现最大的区别在于:张量布局采用(B, H, W, C)(channel-last),这与 MLX 的整体设计风格一致,也直接影响了预处理与模型输入的组织方式。
包内包含两类组件(见init.py 的导出):
| 组件 | 角色 | 说明 |
|---|---|---|
Model | 独立图像编码器 | 直接加载 Hugging Face 的facebook/dinov2-{small,base,large,giant}与facebook/dinov2-with-registers-{small,base,large,giant}检查点,支持 register tokens、patch mask(bool_masked_pos)、forward_features、get_intermediate_layers |
DINOv2 | 共享骨干主干 | 被稠密预测模型复用的纯骨干网络(含完整 Transformer 编码器) |
DINOv2Encoder | 共享骨干封装 | 负责把[0, 1]区间的 RGB 图像缩放到 token 网格、做 ImageNet 归一化,并按层返回特征网格 |
一个重要的边界需要明确:训练阶段的 DINO 头(classification head 等)和 stochastic depth(drop path)没有移植。Model.sanitize在加载检查点时也会显式丢弃classifier.*权重(详见 dinov2.py 与 test_models.py 的测试验证)。因此本模块定位为“推理 / 特征提取 / 骨干复用”,而非完整复现 DINOv2 自监督训练管线。
二、快速上手:加载官方检查点并提取图像特征
README 给出了一个最小可运行的示例,这里将其完整展开并逐行解释:
import mlx.core as mx from mlx_vlm import load model, processor = load("facebook/dinov2-with-registers-base") # processor 返回 channel-first 的 pixel values;模型需要 (B, H, W, C) pixel_values = processor(images=[image], return_tensors="np")["pixel_values"] out = model(mx.array(pixel_values).transpose(0, 2, 3, 1)) out["last_hidden_state"] # (B, 1 + registers + patches, D),已归一化 out["pooler_output"] # (B, D) —— 归一化后的 cls token out["hidden_patch_tokens"] # (B, patches, D)关键点解析:
load直接接受 Hugging Face Hub 的模型标识。facebook/dinov2-with-registers-base的 Hub 配置model_type为dinov2_with_registers,而 utils.py 中的MODEL_REMAPPING会将其重映射到本包("dinov2_with_registers": "dinov2"),因此无需手动转换权重。- 通道顺序转换是显式发生的:processor 输出
(B, C, H, W),而模型接受(B, H, W, C),因此必须调用transpose(0, 2, 3, 1)。这是 channel-last 移植最直接的体现。 - 输出是一个字典(见
Model.__call__,dinov2.py):last_hidden_state:完整归一化 token 序列,形状(B, 1 + num_registers + num_patches, D),当不启用 registers 时即(B, 1 + patches, D);pooler_output:last_hidden_state[:, 0],即归一化后的 cls token,形状(B, D);hidden_patch_tokens:剔除 cls 与 register token 后的纯 patch token,形状(B, num_patches, D),适合直接作为稠密特征;x_prenorm:归一化之前的原始 token 序列(供需要 pre-norm 表示的下游任务使用)。
如果以自监督 / 特征提取场景使用,还可以调用与官方 DINOv2 对齐的forward_features,它返回参考实现的特征字典:x_norm_clstoken、x_norm_regtokens、x_norm_patchtokens、x_prenorm和masks(见 dinov2.py)。测试 test_models.py 验证了带 registers 时的形状约定:cls 在最前、registers 紧随其后、patch tokens 最后。
三、模型配置:HF 字段对齐与架构预设
配置类ModelConfig定义于 config.py,其字段名刻意与 Hugging Face 的dinov2/dinov2_with_registers配置保持一致,从而保证 Hub 检查点无需改动即可加载。同时通过@property别名桥接到共享骨干(DINOv2/DINOv2Encoder)所使用的字段名。
核心字段及默认值:
| 字段 | 默认值 | 说明 |
|---|---|---|
model_type | "dinov2" | 模型类型标识,加载dinov2_with_registers配置时自动重映射 |
hidden_size | 768 | 嵌入维度,别名embed_dim |
num_hidden_layers | 12 | Transformer 层数,别名depth |
num_attention_heads | 12 | 注意力头数,别名num_heads |
mlp_ratio | 4.0 | MLP 隐藏层放大比例 |
layer_norm_eps | 1e-6 | LayerNorm 的 epsilon |
image_size | 224 | 输入图像尺寸(可传[518, 518]列表,__post_init__取第一个值),别名img_size |
patch_size | 14 | patch 大小(14 对应 ViT-*/14 系列) |
num_channels | 3 | 输入通道数 |
qkv_bias | True | QKV 投影是否带 bias |
layerscale_value | 1.0 | LayerScale 初始值 |
use_swiglu_ffn | False | 是否使用 SwiGLU FFN(ViT-g 需要),别名ffn返回"swiglu"或"mlp" |
num_register_tokens | 0 | register token 数量 |
interpolate_offset | 0.0 | 位置编码插值偏移,见下一节 |
interpolate_antialias | False | 位置编码插值是否启用抗锯齿 |
同时内置了DINOV2_PRESETS架构预设,覆盖官方四个规模的 checkpoint:
| 预设名 | 对应模型 | embed_dim | depth | num_heads | FFN 类型 |
|---|---|---|---|---|---|
vits14 | ViT-S/14 | 384 | 12 | 6 | mlp |
vitb14 | ViT-B/14 | 768 | 12 | 12 | mlp |
vitl14 | ViT-L/14 | 1024 | 24 | 16 | mlp |
vitg14 | ViT-g/14 | 1536 | 40 | 24 | swiglu |
注意vitg14使用SwiGlu FFN(SwiGLUFFN类,dinov2.py),其隐藏维度会被向上取整为 8 的倍数:hidden_dim = (int(hidden_dim * 2 / 3) + 7) // 8 * 8,权重结构是融合的w12+w3。这与其余三个规模使用的双层 GELU MLP 不同,也是sanitize需要区分mlp.fc1/fc2与mlp.weights_in/weights_out两套键名的原因(键映射见 dinov2.py)。
四、位置编码插值:两种策略与interpolate_offset
当输入分辨率与训练分辨率(默认 224,patch 14)不一致时,位置编码必须插值到新的 token 网格。本实现支持两种与参考实现一一对应的策略(实现见interpolate_pos_encoding,dinov2.py):
- Hugging Face 参考风格(默认,
interpolate_offset=0.0):基于显式输出尺寸(h // patch_size, w // patch_size)做双三次插值,并且当配置设置interpolate_antialias: true时启用抗锯齿(对应 ATen 的_upsample_bicubic2d_aa,权重实现见 interpolate.py)。 - 原仓库 scale-factor 风格(
interpolate_offset: 0.1):按缩放因子插值,并加入 0.1 的偏移以避免浮点误差(参考实现中该偏移用于推导采样位置,注意代码注释里(sy, sx)的顺序——原仓库从像素高度推导w0并作用到 W 轴)。此时patch_pos_embed通过resize_bicubic_nhwc(..., scale_factor=(sy, sx))完成。
两种方式都会在插值完成后把 cls token 与 patch 位置编码重新拼接,并恢复输入张量的 dtype。该行为与 Hugging Face 参考实现对齐,是保证任意分辨率下特征语义一致的关键细节。
五、源码级解析:组件结构与权重加载
5.1 Transformer 编码器组件
dinov2.py 以标准 ViT 结构组织编码器:
PatchEmbed:nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size)将(B, H, W, C)映射为(B, N, D)的 token 序列,其中N = (H/patch_size) * (W/patch_size);Attention:融合 QKV 的nn.Linear(dim, dim * 3),通过mx.fast.scaled_dot_product_attention实现缩放点积注意力,scale = head_dim ** -0.5;Block:x + LayerScale(Attn(LayerNorm(x)))与x + LayerScale(FFN(LayerNorm(x)))的标准 pre-norm 残差结构,FFN 根据config.ffn在Mlp与SwiGLUFFN间选择;- 可学习 token:
cls_token、pos_embed、mask_token以及可选的register_tokens(形状(1, num_register_tokens, D))均为初始化为零的可学习参数。
5.2 Patch mask 与prepare_tokens
prepare_tokens(dinov2.py)实现了 DINOv2 的 token 准备流程:patch embed →(可选)用mask_token替换被 mask 的 patch → 拼接 cls → 叠加插值后的位置编码 → 若启用 registers 则插入到 cls 之后。bool_masked_pos是形状(B, N)的布尔数组,测试 test_models.py 验证了 mask 位置的 patch embedding 会被 mask token 精确替换。
5.3 权重映射:sanitize与 QKV 融合
Model.sanitize负责把 HF 检查点键名改写到本模块的参数布局,主要包括:
- 剥离
dinov2.前缀、丢弃classifier.*(分类头不属于编码器); - 顶层嵌入键映射:
embeddings.cls_token → cls_token、embeddings.position_embeddings → pos_embed、embeddings.register_tokens → register_tokens等(完整映射见 dinov2.py); - QKV 融合:把 HF 分离的
query/key/value权重沿axis=0拼接为单一的attn.qkv权重(mx.concatenate),与Attention类的融合投影结构一一对应; - 卷积权重转置:
patch_embeddings.projection.weight从(O, I, H, W)转置为(O, H, W, I),适配 channel-last 的nn.Conv2d; - 已经是原始 DINOv2 布局的检查点则原样通过(测试 test_models.py 验证了这一点)。
测试 test_models.py 还对 sanitize 结果做了严格校验:映射后的键集合必须与tree_flatten(model.parameters())完全一致,并能以strict=True成功加载。
六、作为共享骨干:Video Depth Anything 与 MoGe-3
README 明确指出DINOv2/DINOv2Encoder被两个稠密预测模型复用,这是 DINOv2 模块在仓库中最具实战价值的应用场景:
| 模型 | 使用的骨干规模 | 仓库位置 |
|---|---|---|
video_depth_anything | ViT-S/14、ViT-B/14、ViT-L/14 | mlx_vlm/models/video_depth_anything/ |
moge3 | ViT-L/14、ViT-g/14 | mlx_vlm/models/moge3/ |
Video Depth Anything在其 config.py 中直接引用DINOV2_PRESETS派生编码器维度:vits/vitb/vitl分别对应vits14/vitb14/vitl14,再叠加 DPT 解码头预设(features、out_channels、intermediate_layer_idx)。其默认img_size=518、patch_size=14,并且interpolate_offset=0.1——这正是 README 所说“使用原仓库 scale-factor 行为”的落点。骨干通过DINOv2.get_intermediate_layers输出指定层的(patch_tokens, cls_token)供 DPT 头做多尺度融合。
MoGe-3则通过继承DINOv2Encoder构建了MoGe3Encoder(moge3/vision.py):在骨干输出的每一层特征网格上追加一个 1×1 卷积投影,再把各层投影结果求和得到dim_out维的稠密特征。DINOv2Encoder.__call__(dinov2.py)负责把[0, 1]图像双线性缩放(抗锯齿)到(token_rows * patch_size, token_cols * patch_size),随后按 ImageNet 统计归一化(mean[0.485, 0.456, 0.406]、std[0.229, 0.224, 0.225]),最终按指定层返回((B, rows, cols, D), (B, D) cls)的网格列表。测试 test_models.py 验证了该封装能按请求的 token 网格尺寸输出正确的逐层特征形状。
七、测试覆盖与使用边界
TestDinov2测试类(test_models.py)系统验证了模块的契约,可作为二次开发的参考清单:
test_encoder_feature_grids:DINOv2Encoder输出逐层网格与 cls token 的形状;test_config_aliases:HF 风格字段 ↔ 骨干字段别名映射,以及dinov2_with_registersHub 配置的加载(未知键被丢弃);test_forward_features_registers:register token 的位置约定与get_intermediate_layers输出;test_prepare_tokens_masks:mask token 替换逻辑;test_model_call_output:Model.__call__的字典输出约定;test_sanitize_hf_checkpoint/test_sanitize_strips_prefix_and_classifier:权重映射的完备性与前缀/分类头剥离。
使用边界需要特别说明:本模块不包含DINOv2 训练阶段的 DINO 头与 stochastic depth(drop path),因此它适用于“加载官方权重做推理与特征提取”“作为稠密预测骨干复用”两类场景;如果你需要完整的自监督训练实现,应参考原始 DINOv2 仓库(README 顶部给出的 facebookresearch/dinov2 为移植来源)。运行时请确保 MLX 环境已就绪,并通过from mlx_vlm import load使用统一的模型加载入口。
【免费下载链接】mlx-vlmMLX-VLM is a package for inference and fine-tuning of Vision Language Models (VLMs) on your Mac using MLX.项目地址: https://gitcode.com/GitHub_Trending/ml/mlx-vlm
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考