VGGT 中间层特征匹配实验:frame/global token 逐层评价
这篇文章记录一个针对 VGGT 的实验脚本:不修改 VGGT 源码,通过 PyTorch forward hooks 把 aggregator.frame_blocks 和 aggregator.global_blocks 的中间 token 特征取出来,然后用双视图 GT tracks 评价这些 token 的跨视图 patch matching 能力。
实验想回答的问题是:
VGGT 的 frame-wise attention 和 global attention,
到底在哪些层开始形成可用于跨视图匹配的 patch features?
图源说明:根据本实验脚本的数据流重绘。
⭐ 我自己写的 collector 草稿
下面这段是自己写的初版思路,核心目标是:
通过 hook 抓取 VGGT aggregator 中每层 block 的输出;
把 frame/global 两种 token 形状统一成 [B, S, P, C];
剔除 camera token 和 register tokens;
保存到 CPU,供后续分析。
import numpy as np
import torch
from collections import OrderedDict
from typing import Dict, Iterable, Tuple
class FeatureExtraction:
def __init__(
self,
aggregator,
batch_size,
sequence_length,
):
self.aggregator = aggregator
self.batch_size = batch_size
self.sequence_length = sequence_length
self.stored: OrderedDict[Tuple[str,int],torch.Tensor] = OrderedDict()
self.patch_start_idx = aggregator.patch_start_idx
self.hook_list = []
def save(self, sorts: str, depth: int, out: torch.Tensor):
if sorts != "global" and sorts != "frame":
raise ValueError(f"{sorts} is not a legal class.")
if out.shape[0] % self.batch_size != 0 and out.shape[1] % self.sequence_length != 0:
raise ValueError(f"invalid tensor shape.")
if sorts == "global":
x = out.reshape(
self.batch_size,
self.sequence_length,
out.shape[1] // self.sequence_length,
out.shape[2],
)
else:
x = out.reshape(
self.batch_size,
self.sequence_length,
out.shape[1],
out.shape[2],
)
x = x[:, :, self.patch_start_idx:, :]
self.stored[(sorts, depth)] = x.detach().cpu()
def register(self):
for i, block in enumerate(self.aggregator.frame_blocks):
self.hook_list.append(
block.register_forward_hook(
lambda _a, _b, out, depth=i: self.save("frame", depth, out)
)
)
for i, block in enumerate(self.aggregator.global_blocks):
self.hook_list.append(
block.register_forward_hook(
lambda _a, _b, out, depth=i: self.save("global", depth, out)
)
)
def __enter__(self):
self.register()
return self
def remove(self):
for _a in self.hook_list:
_a.remove()
self.hook_list.clear()
这段初版已经抓住了最关键的点:
- 类定义里只写类名或父类,不写构造函数参数。
- 初始化参数放在
__init__()。 - 对象成员后续要用
self.访问。 - Python 字符串类型是
str。 reshape()返回新 tensor,必须接住返回值。out.shape[1] // self.sequence_length必须用整数除法。- hook 要写成
block.register_forward_hook(...)。 - lambda 里用
depth=i固定当前循环层号,否则闭包会拿到循环结束后的最后一个i。
⭐ 这个草稿还可以继续修正的地方
第一,形状检查条件应该更严格。
草稿里:
if out.shape[0] % self.batch_size != 0 and out.shape[1] % self.sequence_length != 0:
这里用 and 容易漏掉错误。只要任意一个维度不满足预期,就应该报错。Codex 版本改成了针对 frame/global 分别检查:
if output.shape[0] != b * s:
raise ValueError(...)
if output.shape[0] != b or output.shape[1] % s != 0:
raise ValueError(...)
第二,缺少 __exit__()。
如果只写 __enter__() 和 remove(),虽然可以手动调用 remove(),但不能完整支持:
with collector:
...
上下文管理器应该在退出时自动 remove hooks,避免重复注册 hook 或内存泄漏。
第三,变量命名可以更清楚。
sorts 可以改成 kind,depth 可以改成 layer。这样后面 CSV 里也更自然:
kind = frame/global
layer = 第几层
Codex 版本脚本总体目标
Codex 版本把上面的 collector 扩展成完整实验入口。它做了几件事:
1. 读取包含 images 和 tracks 的 NPZ。
2. 把 GT 像素坐标转为 VGGT patch token index。
3. 加载 VGGT 模型。
4. 注册 hook,采集每层 frame/global tokens。
5. 对每层特征做 view1 -> view2 最近邻匹配。
6. 计算误差、PCK、hard-negative margin 等指标。
7. 保存 CSV、曲线图和匹配可视化。
输入 NPZ 约定:
| 字段 | 形状 | 作用 |
|---|---|---|
images | [2,3,H,W] 或 [2,H,W,3] | 两张预处理后图像 |
tracks | [2,Q,2] | 两视图对应点,坐标为 (u,v) |
visibility | [2,Q],可选 | 两张图中是否可见 |
positive_mask | [Q],可选 | 是否为真实正对应 |
patch 坐标转换
VGGT 的 DINO patch size 默认为 14,因此每个 patch token 对应图像上的一个:
14 x 14 像素网格
像素坐标到 patch index:
def uv_to_patch_index(
uv: torch.Tensor,
height: int,
width: int,
patch_size: int = 14,
) -> Tuple[torch.Tensor, torch.Tensor]:
hp = height // patch_size
wp = width // patch_size
px = torch.floor(uv[:, 0] / patch_size).long()
py = torch.floor(uv[:, 1] / patch_size).long()
valid = (px >= 0) & (px < wp) & (py >= 0) & (py < hp)
return py * wp + px, valid
这里:
u是横坐标,对应 patch 列号px。v是纵坐标,对应 patch 行号py。py * wp + px是按行展开的 patch index。valid用来过滤落在图像外的 GT track。
patch index 回到像素坐标:
def patch_index_to_uv(
index: torch.Tensor,
width_patches: int,
patch_size: int,
) -> torch.Tensor:
py = index // width_patches
px = index % width_patches
return torch.stack(
((px.float() + 0.5) * patch_size, (py.float() + 0.5) * patch_size),
dim=-1,
)
它返回的是 patch 中心坐标,而不是 patch 左上角。
⭐ 为什么评价误差时用 patch 中心坐标?
因为中间层特征是 patch token,而不是原始像素 token。
如果一个预测 patch index 是:
第 py 行,第 px 列
它并没有表达这个 patch 内具体哪个像素,而是表达整个 patch 的特征。因此可视化或计算像素误差时,用 patch 中心点更合理:
这意味着实验精度上限本来就是 patch-level 的,不能把它理解成亚像素级 matching。
LayerFeatureCollector:不改源码采集 VGGT 中间层
Codex 版本的 collector:
class LayerFeatureCollector:
"""不修改 Aggregator 源码,通过 hooks 采集各层输出。
frame block 原始输出:[B*S, P_total, C]
global block 原始输出:[B, S*P_total, C]
统一保存为:[B, S, H_p*W_p, C],并剔除特殊 token。
"""
def __init__(self, aggregator, batch_size: int, num_views: int):
self.aggregator = aggregator
self.batch_size = batch_size
self.num_views = num_views
self.patch_start_idx = int(aggregator.patch_start_idx)
self.features: "OrderedDict[Tuple[str, int], torch.Tensor]" = OrderedDict()
self.handles = []
这里最关键的是 patch_start_idx。
VGGT 每一帧前面不全是 patch tokens,通常还包含:
1 个 camera token
4 个 register tokens
后面才是 image patch tokens
所以后面必须执行:
x = x[:, :, self.patch_start_idx:, :]
否则 patch index 会整体错位。
frame block 输出形状
frame-wise attention 是每张图内部独立处理,因此常见输出形状:
[B*S, P_total, C]
Codex 版本恢复成:
x = output.reshape(b, s, output.shape[-2], output.shape[-1])
也就是:
[B, S, P_total, C]
global block 输出形状
global attention 把多张图 token 拼在一起,因此常见输出形状:
[B, S*P_total, C]
Codex 版本先算每张图的 token 数:
tokens_per_view = output.shape[1] // s
再恢复:
x = output.reshape(b, s, tokens_per_view, output.shape[-1])
最终也统一成:
[B, S, P_total, C]
剔除特殊 token 后变成:
[B, S, H_p*W_p, C]
⭐ 为什么要用 forward hook,而不是改 VGGT 源码?
hook 的好处是实验侵入性低。
如果直接改 VGGT 源码:
需要改 aggregator.forward()
需要维护额外返回值
可能影响官方 checkpoint 或其他调用路径
而 forward hook 可以临时挂在模块上:
handle = block.register_forward_hook(hook)
forward 时自动拿到该 block 的输出:
hook(module, inputs, output)
实验结束后调用:
handle.remove()
就能恢复模型原样。
这很适合做“观察中间层表示”的实验。
最近邻匹配指标
核心评价函数是:
@torch.no_grad()
def evaluate_feature_matching(
feature_view1: torch.Tensor,
feature_view2: torch.Tensor,
source_patch_indices: torch.Tensor,
target_uv_gt: torch.Tensor,
target_patch_indices: torch.Tensor,
height: int,
width: int,
patch_size: int = 14,
exclusion_radius: int = 1,
) -> Dict[str, object]:
输入的 feature_view1 和 feature_view2 形状应为:
[H_p*W_p, C]
先做 L2 normalize:
f1 = F.normalize(feature_view1.float(), dim=-1)
f2 = F.normalize(feature_view2.float(), dim=-1)
然后取 source patches 对应的 query:
query = f1[source_patch_indices]
计算和 view2 所有 patch 的相似度:
similarity = query @ f2.T
最近邻预测:
predicted_indices = similarity.argmax(dim=-1)
于是每个 source GT patch 都会在 target view 的所有 patches 中找到一个最相似 patch。
输出指标
函数返回:
| 指标 | 含义 |
|---|---|
mean_error | 预测 patch 中心与 GT 像素坐标的平均像素距离 |
median_error | 中位误差 |
pck_1_patch | 误差小于等于 1 个 patch size 的比例 |
pck_2_patch | 误差小于等于 2 个 patch size 的比例 |
exact_patch_accuracy | 预测 patch index 是否严格等于 GT patch index |
patch_within_1_accuracy | 预测 patch 是否在 GT 周围 1 邻域内 |
mean_positive_similarity | 正样本匹配相似度 |
mean_hard_negative_similarity | hard negative 相似度 |
mean_margin | 正样本相似度减 hard negative 相似度 |
⭐ hard-negative margin 指标在看什么?
对每个 query patch,有一个 GT target patch:
positive
也有很多错误 target patches:
negative
普通 negative 可能很容易区分,所以没有太大诊断价值。hard negative 指的是:
排除 GT 附近 patch 后,剩下所有错误 patch 里最像 query 的那个。
代码里用 exclusion_radius 排除 GT patch 周围区域:
exclude = (
(px[None] - gt_px[:, None]).abs() <= exclusion_radius
) & (
(py[None] - gt_py[:, None]).abs() <= exclusion_radius
)
然后找最大错误相似度:
hard_negative = similarity.masked_fill(exclude, float("-inf")).max(dim=-1).values
margin:
margin = positive - hard_negative
如果 margin 越大,说明这一层特征越能把正确对应和最难错误对应拉开。
所以这个实验不只是看“有没有找对”,还看“对得是否有信心”。
GT correspondence 筛选
load_experiment_npz() 负责读取并检查数据:
images可以是[2,3,H,W],也可以是[2,H,W,3]。- 如果是 channel-last,会转成 channel-first。
- 如果是
uint8 [0,255],会归一化到[0,1]。 tracks必须是[2,Q,2]。- 没有
visibility时默认全部可见。 - 没有
positive_mask时默认全是真正对应。
select_valid_correspondences() 负责筛选:
两张图都可见;
是正样本;
source 和 target 都在图像范围内;
可选过滤小视差;
可选按 source patch 去重。
小视差过滤:
gt_patch_displacement = torch.maximum(
(target_px - source_px).abs(),
(target_py - source_py).abs(),
)
valid &= gt_patch_displacement >= min_gt_patch_displacement
这用于排除太简单的对应。例如两张图几乎没动时,同一个 patch 对同一个 patch 的匹配很容易,不足以评价 global attention 是否真的学到跨视图几何。
同一 source patch 多数投票
由于 GT tracks 是像素级的,但评价是 patch-level 的,同一个 source patch 里可能有多个 GT 点,对应到不同 target patches。
Codex 版本采用多数投票:
values, counts = torch.unique(candidate_targets, return_counts=True)
target_mode = values[counts.argmax()]
然后只保留投到这个 target mode 的点,并取平均 source/target UV。
这让每个 source patch 只有一个主要 GT target patch,避免同一个 patch 被多个真值目标拉扯。
模型加载
def load_model(model_source: str, device: torch.device):
from vggt.models.vggt import VGGT
source = Path(model_source)
if source.is_file():
model = VGGT()
checkpoint = torch.load(source, map_location="cpu")
...
model.load_state_dict(checkpoint, strict=True)
else:
model = VGGT.from_pretrained(model_source)
return model.eval().to(device)
这个函数支持两种来源:
本地 checkpoint 文件
或
Hugging Face 模型名,例如 facebook/VGGT-1B
本地 checkpoint 里如果有 model 或 state_dict 包装,会自动拆出来;如果 key 前面有 module.,也会去掉。
主流程 main()
主流程可以概括成:
parse args
-> load NPZ
-> select valid correspondences
-> load VGGT
-> attach LayerFeatureCollector hooks
-> run model.aggregator(images)
-> evaluate every collected layer
-> save CSV / plots / match visualization / run summary
关键推理部分:
collector = LayerFeatureCollector(model.aggregator, batch_size=1, num_views=2)
with collector, torch.inference_mode(), amp_context:
model.aggregator(model_images)
这里没有调用完整 VGGT 输出 heads,而是只跑:
model.aggregator(model_images)
因为实验只关心 aggregator 的中间 tokens,不关心最终 camera/depth/pointmap heads。
⭐ 为什么只跑 aggregator 就够?
VGGT 的主干可以粗略分成:
images
-> aggregator / transformer backbone
-> prediction heads
你这个实验要评价的是:
frame/global attention 层里的 patch token 是否已经有跨视图匹配能力。
这些 token 在 aggregator 内部就已经存在。
如果跑完整模型,当然也可以,但会额外计算:
camera head
depth head
pointmap head
tracking head
这会占用更多显存和时间,而且和当前实验目标无关。
所以直接跑 model.aggregator(model_images) 是合理的。
保存结果
脚本设计上会输出:
metrics.csv
mean_error.png
pck_1_patch.png
exact_patch_accuracy.png
positive_vs_hard_negative.png
matches_global_*.png 或 matches_frame_*.png
run_summary.json
其中 run_summary.json 记录:
{
"image_shape": [2, 3, "H", "W"],
"patch_grid": ["H/14", "W/14"],
"num_valid_matches": "...",
"patch_start_idx": "...",
"layers_collected": "...",
"note": "frame/global 共应为 depth*2 组;默认 VGGT depth=24,因此应为 48。"
}
如果默认 VGGT 有 24 层 frame/global blocks,那么理想情况下:
frame 24 层
global 24 层
总计 48 组中间特征
你粘贴的 Codex 脚本里有一处截断
你给出的代码中这一段明显被截断了:
def save_csv(rows: Iterable[dict], output_path: Path) -> None:
columns = [
"kind", "layer", "num_matches", "kind), key=lambda r: r["layer"])
ax.plot(
这里从 save_csv() 的 columns 列表突然跳到了绘图逻辑,说明中间缺了:
columns 定义剩余部分;
CSV 写入逻辑;
save_plots() 函数定义开头;
各曲线图的调用。
这不是算法问题,而是粘贴文本缺失。补齐思路如下:
def save_csv(rows: Iterable[dict], output_path: Path) -> None:
rows = list(rows)
columns = [
"kind",
"layer",
"num_matches",
"mean_error",
"median_error",
"pck_1_patch",
"pck_2_patch",
"exact_patch_accuracy",
"patch_within_1_accuracy",
"mean_positive_similarity",
"mean_hard_negative_similarity",
"mean_margin",
]
with output_path.open("w", newline="", encoding="utf-8") as f:
writer = csv.DictWriter(f, fieldnames=columns)
writer.writeheader()
for row in rows:
writer.writerow({key: row.get(key, "") for key in columns})
def save_plots(rows: Iterable[dict], output_dir: Path) -> None:
import matplotlib.pyplot as plt
rows = list(rows)
def plot_metric(key: str, ylabel: str, filename: str) -> None:
fig, ax = plt.subplots(figsize=(8, 5))
for kind in ("frame", "global"):
selected = sorted(
(r for r in rows if r["kind"] == kind),
key=lambda r: r["layer"],
)
ax.plot(
[r["layer"] for r in selected],
[r[key] for r in selected],
marker="o",
markersize=3,
label=kind,
)
ax.set_xlabel("Layer")
ax.set_ylabel(ylabel)
ax.grid(alpha=0.3)
ax.legend()
fig.tight_layout()
fig.savefig(output_dir / filename, dpi=180)
plt.close(fig)
plot_metric("mean_error", "Mean pixel error", "mean_error.png")
plot_metric("median_error", "Median pixel error", "median_error.png")
plot_metric("pck_1_patch", "PCK @ 1 patch", "pck_1_patch.png")
plot_metric("pck_2_patch", "PCK @ 2 patches", "pck_2_patch.png")
plot_metric("exact_patch_accuracy", "Exact patch accuracy", "exact_patch_accuracy.png")
plot_metric("patch_within_1_accuracy", "Patch within 1 accuracy", "patch_within_1_accuracy.png")
plot_metric("mean_margin", "Positive - hard negative", "mean_margin.png")
fig, ax = plt.subplots(figsize=(8, 5))
for kind, key, style in (
("frame", "mean_positive_similarity", "-"),
("frame", "mean_hard_negative_similarity", "--"),
("global", "mean_positive_similarity", "-"),
("global", "mean_hard_negative_similarity", "--"),
):
selected = sorted(
(r for r in rows if r["kind"] == kind),
key=lambda r: r["layer"],
)
ax.plot(
[r["layer"] for r in selected],
[r[key] for r in selected],
linestyle=style,
label=f"{kind}-{'positive' if 'positive' in key else 'hard-neg'}",
)
ax.set_xlabel("Layer")
ax.set_ylabel("Cosine similarity")
ax.grid(alpha=0.3)
ax.legend()
fig.tight_layout()
fig.savefig(output_dir / "positive_vs_hard_negative.png", dpi=180)
plt.close(fig)
这段实验最终能说明什么
如果实验跑通,可以观察几类趋势。
第一,frame 和 global 哪个更适合跨视图匹配。
直觉上:
frame block 主要做单图内部建模;
global block 允许不同 view tokens 交互。
所以如果 VGGT 的 global attention 真在建立跨视图几何关系,那么 global 层的:
mean_error 应下降;
PCK 应上升;
positive-hard negative margin 应变大。
第二,跨视图匹配能力在哪些层出现。
早期层可能更像局部纹理:
边缘、颜色、纹理 patch
中后期层可能更像几何语义:
同一表面、同一物体部件、跨视图对应
逐层画曲线就能看到这个变化。
第三,VGGT 的中间 token 是否能当作 dense matching descriptor。
如果最后几层 global token 的 exact_patch_accuracy 和 pck_1_patch 很高,说明这些 token 已经具备一定 matching 能力。反之,如果误差仍然很大,说明 VGGT 的最终几何输出并不一定依赖显式可分离的 patch descriptor。
运行示例
假设已经在 VGGT 工程环境里,并且有一个实验数据:
pair_tracks.npz
可以运行:
python vggt_layer_matching_eval.py \
--input-npz pair_tracks.npz \
--output-dir outputs/vggt_layer_matching \
--model facebook/VGGT-1B \
--device cuda \
--dtype bfloat16 \
--patch-size 14 \
--min-gt-patch-displacement 2 \
--visualize-kind global \
--visualize-layer -1
如果显存紧张,可以先用:
--device cpu
验证数据和代码逻辑,但速度会慢很多。
我对这段代码的整体评价
这个实验设计是合理的,尤其适合验证你前面读 VGGT 时的一个关键疑问:
VGGT 的中间层是否真的学到了跨视图对应关系?
它的优点是:
- 不改 VGGT 源码。
- 同时比较 frame/global 两类 block。
- 使用 GT tracks 做定量评价。
- 不只看最近邻正确率,还看 hard negative margin。
- 输出 CSV、曲线和可视化,方便观察层间变化。
需要注意的限制:
- 评价是 patch-level,不是 pixel-level。
- GT tracks 必须和输入图像使用同一预处理坐标系。
- 如果图像 resize/crop 后没有同步变换 tracks,结果会完全失真。
save_csv/save_plots那段需要补齐后才能运行。model.aggregator(model_images)的返回和内部形状依赖具体 VGGT 版本,升级仓库后要重新检查patch_start_idx和 block 输出形状。
这篇实验文章可以作为后续实际跑 VGGT 中间层分析的参照。