0. 真实照片切块¶
一张 224×224 图片会被切成 14×14 个图块。每个图块展平后是一条向量,后续模型把这些向量当作序列读取。
In [2]:
# 使用真实建筑照片,按 ViT 常见设置切成 16x16 图块。
raw_photo = load_sample_image("china.jpg")
vit_image = np.asarray(Image.fromarray(raw_photo).resize((224, 224))) / 255.0
patch_size = 16
patch_grid = vit_image.shape[0] // patch_size
patches = vit_image.reshape(patch_grid, patch_size, patch_grid, patch_size, 3).swapaxes(1, 2)
patch_tokens = patches.reshape(-1, patch_size * patch_size * 3)
patch_summary = []
for patch_id, patch in enumerate(patches.reshape(-1, patch_size, patch_size, 3)):
row, col = divmod(patch_id, patch_grid)
patch_summary.append({
"图块编号": patch_id,
"行": row,
"列": col,
"向量维度": patch_tokens.shape[1],
"R均值": patch[:, :, 0].mean(),
"G均值": patch[:, :, 1].mean(),
"B均值": patch[:, :, 2].mean(),
"亮度标准差": patch.mean(axis=2).std(),
})
patch_df = pd.DataFrame(patch_summary)
display(patch_df.head(12).round(3))
| 图块编号 | 行 | 列 | 向量维度 | R均值 | G均值 | B均值 | 亮度标准差 | |
|---|---|---|---|---|---|---|---|---|
| 0 | 0 | 0 | 0 | 768 | 0.703 | 0.805 | 0.917 | 0.007 |
| 1 | 1 | 0 | 1 | 768 | 0.718 | 0.819 | 0.937 | 0.005 |
| 2 | 2 | 0 | 2 | 768 | 0.732 | 0.833 | 0.951 | 0.004 |
| 3 | 3 | 0 | 3 | 768 | 0.746 | 0.845 | 0.962 | 0.006 |
| 4 | 4 | 0 | 4 | 768 | 0.763 | 0.861 | 0.966 | 0.014 |
| 5 | 5 | 0 | 5 | 768 | 0.784 | 0.879 | 0.977 | 0.005 |
| 6 | 6 | 0 | 6 | 768 | 0.808 | 0.896 | 0.990 | 0.006 |
| 7 | 7 | 0 | 7 | 768 | 0.838 | 0.915 | 0.999 | 0.006 |
| 8 | 8 | 0 | 8 | 768 | 0.864 | 0.932 | 0.996 | 0.004 |
| 9 | 9 | 0 | 9 | 768 | 0.892 | 0.945 | 0.996 | 0.005 |
| 10 | 10 | 0 | 10 | 768 | 0.909 | 0.954 | 0.997 | 0.005 |
| 11 | 11 | 0 | 11 | 768 | 0.925 | 0.963 | 0.999 | 0.003 |
In [3]:
# 绘制真实图片图块网格和若干具体图块。
selected_patches = [18, 43, 87, 112, 145, 181]
fig = plt.figure(figsize=(10.8, 6.6))
gs = fig.add_gridspec(2, 6, height_ratios=[3.4, 1.5], hspace=0.22, wspace=0.08)
ax_img = fig.add_subplot(gs[0, :])
ax_img.imshow(vit_image)
ax_img.set_title("真实图片切成 14x14 个图块", loc="left", fontweight="bold")
ax_img.set_xticks(np.arange(0, 225, patch_size))
ax_img.set_yticks(np.arange(0, 225, patch_size))
ax_img.grid(color="#ffffff", linewidth=0.8)
ax_img.tick_params(labelbottom=False, labelleft=False, length=0)
for slot, patch_id in enumerate(selected_patches):
row, col = divmod(patch_id, patch_grid)
ax = fig.add_subplot(gs[1, slot])
ax.imshow(patches[row, col])
ax.set_title(f"图块 {patch_id}\n({row},{col})", fontsize=9, fontweight="bold")
ax.set_xticks([])
ax.set_yticks([])
fig.suptitle("ViT 图块切分:真实照片转成图块向量序列", x=0.08, ha="left", fontsize=14, fontweight="bold", color="#0f172a")
plt.tight_layout()
plt.show()