跳到主要内容

VGGT 中间层特征匹配实验:frame/global token 逐层评价

这篇文章记录一个针对 VGGT 的实验脚本:不修改 VGGT 源码,通过 PyTorch forward hooks 把 aggregator.frame_blocksaggregator.global_blocks 的中间 token 特征取出来,然后用双视图 GT tracks 评价这些 token 的跨视图 patch matching 能力。

实验想回答的问题是:

VGGT 的 frame-wise attention 和 global attention,
到底在哪些层开始形成可用于跨视图匹配的 patch features?

VGGT 中间层 patch matching 实验流程

图源说明:根据本实验脚本的数据流重绘。

⭐ 我自己写的 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 可以改成 kinddepth 可以改成 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 中心点更合理:

u=(px+0.5)patch_sizeu=(p_x+0.5)\cdot patch\_sizev=(py+0.5)patch_sizev=(p_y+0.5)\cdot patch\_size

这意味着实验精度上限本来就是 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_view1feature_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_similarityhard 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 里如果有 modelstate_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)

这段实验最终能说明什么

如果实验跑通,可以观察几类趋势。

第一,frameglobal 哪个更适合跨视图匹配。

直觉上:

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_accuracypck_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 中间层分析的参照。