
在实际的计算机视觉项目中使用 Vision TransformerViT时大家很快会遇到一个问题模型确实能分类但没人能直观说清楚它到底“看”了图像的哪些区域。VIT-2动画这个名字指的就是一类以动画方式可视化 ViT 内部运行过程的工程实践。它不是官方发布的模型版本也不是一个固定代码仓库而是把 ViT 的 patch 切分、位置编码、注意力权重变化完整展示出来让开发者在调参与解释模型时有一套可以观察和比对的工具。本文会带读者完成一个最小可运行的 Python 可视化项目输入一张图片输出一个 GIF 动画动画按顺序展示 ViT 每一层中 CLS token 对不同图像块的注意力分布。这个主题适合正在学习和使用 Transformer 系列模型的算法工程师也适合需要给非技术人员解释模型行为的开发同学。完成本文的实践后你得到的不只是一个 GIF而是一套可以放在任何图像分类任务中验证模型行为的可视化工具。你可以拿它观察浅层与深层注意力差异也可以对比不同模型、不同输入对注意力分布的影响。1. 为什么要用动画观察 ViT1.1 ViT 的工作流程Vision Transformer 的核心思想是把图片当作一串 token 来处理。一张 224x224 的图片先切成 16x16 的小块得到 14x14 的网格也就是 196 个 patch。每个 patch 展平后经过一个线性映射变成向量再与一个可学习的位置编码相加最后送入标准的 Transformer 编码器。在这个过程中CLS token 会被拼在序列最前面。它的角色很像一个“汇总容器”经过多层注意力计算后分类头只读取最后一个 CLS token 的输出来预测类别。而每一层注意力中CLS token 会与图像中的所有 patch token 计算相关性得到 196 个注意力分数。这些分数可以理解为模型当前认为哪些图像区域对最终分类更重要。VIT-2动画要展示的就是这 196 个分数的分布随层数变化的过程。只截取某一层的结果做热力图只能看到静态信息用动画连续展示每一层才能看清注意力的迁移过程。1.2 静态热力图和动画的区别静态热力图通常只提供一张图。常见做法是取最后一层所有注意力头的平均权重叠加到原图上看。这种方式对最终结果有一定的解释力但会掩盖层间差异。实际上ViT 的注意力在不同层有明显变化。底层往往关注局部纹理、边缘、颜色块高层会逐步聚焦到与类别强相关的区域。如果只看最后一层就无法知道模型在哪个阶段开始从“看纹理”转向“看结构”。动画的优势在于时间维度的叠加每一帧代表一层读者可以明显感受到注意力从分散到集中的过程。此外在排查模型“为什么错分”时静态图只能告诉你最后错在哪动画能告诉你哪一层开始出现误导性注意力。这个信息对调试非常有用。1.3 本文目标本文会选择 HuggingFace Transformers 中的google/vit-base-patch16-224作为示例模型。先说明依赖和项目结构再给出完整的可视化脚本然后逐段解释关键代码的维度变化和参数含义最后讲解运行验证、常见问题和生产环境落地建议。整个实践不需要 GPUCPU 也可以完成。只要网络能下载预训练模型并且能安装 Python 依赖就能复现。如果你没有合适的测试图片脚本也会自动生成一张测试图保证流程可以跑通。2. 环境准备和项目结构2.1 依赖清单与版本选择可视化 ViT 需要的第三方库并不复杂核心是模型加载、图像处理、绘图和动画导出四部分。下表列出了主要依赖及其用途依赖库用途推荐版本Python运行环境3.9 及以上transformers加载预训练 ViT 模型4.30 及以上torch模型推理和 Tensor 操作2.0 及以上Pillow图像读取与缩放9.0 及以上matplotlib绘图和 GIF 动画导出3.7 及以上numpy数组处理1.23 及以上这里把transformers作为模型来源是因为它提供了非常方便的output_attentions开关不需要修改模型源码就能拿到每层注意力权重。如果使用torchvision的 ViT 实现通常需要注册 hook 或者魔改源码实现成本更高不推荐作为入门方案。安装命令如下pip install torch transformers pillow matplotlib numpy如果本机有 NVIDIA GPU 且安装了 CUDA 版本 PyTorch计算会自动走 GPU。但本文脚本的计算量非常小CPU 完全够用因此不强制要求 GPU 环境。2.2 项目目录建议新建一个专门目录避免模型缓存和输出文件散落各处vit2-animation/ ├── requirements.txt ├── vit2_animation.py ├── images/ │ └── demo.jpg └── output/ └── vit2_attn.gif其中images/存放输入的测试图片output/用来保存生成的动画。没有图片时脚本会自动用 NumPy 生成一张红蓝分界的测试图所以即使images为空程序也能运行。这样组织的好处是后续如果要批量处理多张图片只需要在images/下放多个文件再略微修改脚本即可。不要把所有文件都堆在根目录否则排查问题时很难分清输入、输出和代码。3. 最小可视化动画实现3.1 加载预训练 ViT 模型脚本的第一部分负责模型和图像处理器加载。以google/vit-base-patch16-224为例它接受 224x224 的 RGB 图像patch 大小为 16序列长度是 197其中 1 个 CLS token 加上 196 个 patch token。import sys import numpy as np import torch import matplotlib.pyplot as plt import matplotlib.animation as animation from PIL import Image from transformers import AutoImageProcessor, AutoModelForImageClassification MODEL_NAME google/vit-base-patch16-224 DEVICE torch.device(cuda if torch.cuda.is_available() else cpu) processor AutoImageProcessor.from_pretrained(MODEL_NAME) model AutoModelForImageClassification.from_pretrained( MODEL_NAME, output_attentionsTrue ) model.to(DEVICE) model.eval()这里使用AutoModelForImageClassification而不是AutoModel是因为分类模型会在最后输出logits可以直接查看预测类别和概率。output_attentionsTrue会让前向传播额外返回每层注意力张量。注意AutoImageProcessor在旧版本 Transformers 中叫做AutoFeatureExtractor。如果你使用的是 4.30 以下版本可以把第一行模型加载改成from transformers import AutoFeatureExtractor processor AutoFeatureExtractor.from_pretrained(MODEL_NAME)否则会提示模块不存在。3.2 读取并预处理输入图像模型要求的输入不能直接使用原始图像文件。需要先经过AutoImageProcessor完成缩放、归一化、通道顺序调整并转换成 PyTorch Tensor。def load_image(path: str): if path: img Image.open(path).convert(RGB) else: # 生成一张红蓝上下分界的测试图 arr np.zeros((224, 224, 3), dtypenp.uint8) arr[:112, :, 0] 255 arr[112:, :, 1] 255 img Image.fromarray(arr) return img def process_image(processor, img): inputs processor(img, return_tensorspt) pixel_values inputs[pixel_values].to(DEVICE) return pixel_values, img这里有两个容易忽略的细节第一Image.open后必须执行.convert(RGB)。有些图片是 RGBA 四通道或灰度单通道直接送入模型会报维度错误。统一转成 RGB 可以规避大部分输入问题。第二processor内部已经完成了Resize、Normalize不要再手动调用torchvision.transforms重复归一化否则会出现注意力热力图严重偏移或分类概率异常。3.3 获取多层注意力权重加载模型并输入图像后前向传播会返回一个BaseModelOutputWithPoolingAndCrossAttentions或分类输出对象。重点观察attentions字段。with torch.no_grad(): outputs model(pixel_values, output_attentionsTrue) attention_tuple outputs.attentions # attention_tuple 是一个元组长度为 Transformer 编码器层数 # attention_tuple[0].shape: [batch_size, num_heads, num_tokens, num_tokens] # 对于 vit-base-patch16-224层数为 12注意力头数为 12在 ViT-base 中attention_tuple共有 12 个元素对应 12 个编码器层。每个元素的形状是[1, 12, 197, 197]。197x197中的第 0 行是 CLS token 对所有 token 的注意力权重第 0 列是以 CLS token 作为 key 时各 token 对它的注意力权重。可视化通常看第 0 行也就是 CLS token 如何看待图像上的各个 patch。3.4 将注意力权重变成可视化帧这一步需要把197个 token 的权重转成14x14的二维网格因为原始图像被切成了 14 行 14 列。下面的函数用于提取某一层中 CLS token 对所有 patch 的注意力权重并缩放回原始图像尺寸。def attention_to_image(attention, img_size, layer_idx, head_idxNone): # attention: [heads, seq_len, seq_len] if head_idx is None: attn_avg attention.mean(dim0) else: attn_avg attention[head_idx] cls_attn attn_avg[0, 1:] # 去掉 CLS token patch_attn cls_attn.reshape(14, 14).float().cpu().numpy() # 使用双三次插值放大到原图尺寸 attn_img Image.fromarray(patch_attn) attn_img attn_img.resize( (img_size[0], img_size[1]), resampleImage.Resampling.BICUBIC ) return np.asarray(attn_img)cls_attn的长度是 196正好是 14x14 个 patch。如果模型输入尺寸不是 224x224这里的 14 也要相应改变。更通用的写法是grid_size int(patch_num ** 0.5) patch_attn cls_attn.reshape(grid_size, grid_size)但为了演示清晰这里先固定成 14。3.5 合成 GIF 动画Matplotlib 的matplotlib.animation.FuncAnimation可以逐帧更新图像。我们为每一层生成一帧帧内容包含原图、注意力热力图叠加层和标题信息。def create_animation(attention_tuple, original_img, output_path, interval600): fig, ax plt.subplots(figsize(5, 5)) frames [] for layer_idx in range(len(attention_tuple)): frame_data attention_to_image( attention_tuple[layer_idx][0], original_img.size, layer_idx ) frames.append(frame_data) def update(frame_idx): ax.clear() ax.imshow(original_img) attn_map frames[frame_idx] ax.imshow(attn_map, cmapjet, alpha0.6) ax.set_title(fLayer {frame_idx 1}) ax.axis(off) anim animation.FuncAnimation( fig, update, frameslen(frames), intervalinterval, repeatTrue ) anim.save(output_path, writerpillow) plt.close(fig)interval600表示每帧间隔 600 毫秒也就是 0.6 秒。12 层一共 12 帧完整播放一遍约 7.2 秒。如果觉得播放太快可以增大 interval如果希望动画更平滑可以把frames扩展成每一层内的多个注意力头但 GIF 文件体积会明显变大。writerpillow是保存 GIF 的推荐方式因为它不依赖系统中是否安装了 ImageMagick。只要确认 Pillow 已安装即可。3.6 主程序入口主程序把前面几个函数串起来支持命令行传入图片路径和输出路径。如果没有传入图片路径就使用自动生成的测试图。if __name__ __main__: image_path sys.argv[1] if len(sys.argv) 1 else None output_path sys.argv[2] if len(sys.argv) 2 else output/vit2_attn.gif original_img load_image(image_path) pixel_values, original_img process_image(processor, original_img) with torch.no_grad(): outputs model(pixel_values, output_attentionsTrue) create_animation(outputs.attentions, original_img, output_path) print(f动画已保存到: {output_path})这里的输出路径如果没有创建output目录需要先手动创建。也可以在主程序里加上os.makedirs(output, exist_okTrue)避免首次运行报错。4. 关键代码的工作原理和参数解读4.1 输入张量的维度变化图像进入 ViT 之前需要把H x W x C的数组转换成模型期待的B x C x H x W格式并且归一化到模型训练时的分布。AutoImageProcessor会自动完成这些操作但理解维度变化仍然重要。阶段形状说明原始图像H x W x 3PIL Imageprocessor 输出1 x 3 x 224 x 224BatchSize1RGB 三通道patch embedding1 x 197 x 768196 个 patch 加上 1 个 CLS token注意力张量1 x 12 x 197 x 197每层 12 个头每个头输出完整注意力矩阵第 3 行的768是 ViT-base 的 hidden size。如果使用 ViT-large这个数字会变成1024注意力头数也会变化但可视化代码的结构不需要改动。4.2 attention 输出维度解读outputs.attentions是嵌套元组最外层对应层数第二层对应 batch第三层是注意力头最后是注意力矩阵的两个维度。# 假设取出第 4 层第 0 个头 attn_tensor outputs.attentions[3][0, 0] # shape [197, 197]这 197 个 token 的排序是固定的0 号是 CLS token1 到 196 号是按行扫描排列的 patch token。也就是说原始图像左上角 patch 是 1 号右上角是 14 号下一行第一个 patch 是 15 号。理解这个顺序后才不至于在热力图叠加时出现区域错位。4.3 可视化参数推荐参数默认值建议范围影响interval600400-1000帧间隔越大播放越慢便于观察中间层变化alpha0.60.3-0.8越大热力图越遮挡原图越小越偏向原图纹理cmapjetjet、coolwarm、viridis影响色带对比不同色带对注意力高低敏感度不同head_idxNoneNone 表示取均值或指定 0-11观察单一注意力头时更能发现特征但噪声也更大如果只是快速验证取alpha0.6和cmapjet最直观。如果要深入分析某个具体注意力头建议把cmap换成viridis它对高注意力区域的识别更友好。5. 运行验证与结果分析5.1 执行方式确保当前目录有output目录然后运行python vit2_animation.py images/demo.jpg output/vit2_attn.gif如果不传入图片路径程序会自动生成测试图并输出到默认路径python vit2_animation.py正常运行时控制台会显示模型加载进度最后输出动画已保存到: output/vit2_attn.gif如果是在服务器上执行且没有图形界面需要在使用 Matplotlib 前设置后端为Agg。脚本中可以在导入matplotlib.pyplot之前加入import matplotlib matplotlib.use(Agg)这会告诉 Matplotlib 只输出文件不弹出窗口。5.2 预期输出生成的 GIF 会包含 12 帧。每帧展示一层中 CLS token 对所有 patch 的注意力分布。叠加层越亮、越偏红的区域说明模型对那里的关注度越高。对于测试图上半红、下半蓝如果模型工作正常注意力一开始可能会分散到大量边缘位置因为上下色块边界是最容易区分两个颜色的特征。随着层数加深注意力会逐渐集中到有利于分类的区域。5.3 不同 layer 的注意力差异观察不同层的注意力图时可以按这个顺序检查前两层是否呈现大范围分散状态。中间层是否开始出现明显的局部集中。最后几层是否聚焦在语义上更明确的物体区域。这种现象是 ViT 预训练模型的常见特征但不是所有输入都会严格遵循。比如纯纹理图片、复杂街景、遮挡物体注意力迁移规律会有变化。如果你发现最后一层还在大面积分散可以先检查注意力头均值计算是否正确再检查是否选择了过低层数的输出。5.4 在无显示环境服务器中生成 GIF生产环境通常没有桌面环境直接执行脚本可能会遇到类似TclError的报错。解决办法是在脚本顶部设置matplotlib.use(Agg)。如果使用 Jupyter Notebook则不需要额外设置但保存 GIF 时建议滴加dpi80控制文件大小。保存较大 GIF 时可以降低帧数或图像缩放比例。比如只展示 6 个关键层或者在create_animation中把原图缩小到 128x128 后再叠加热力图这样文件体积可以从数 MB 降到几百 KB。6. 常见问题排查6.1 注意力图错位或重叠位置不对现象热力图叠加在原图上看起来很“碎”明明关注的是猫脸高亮区域却跑到背景上。可能原因图像被processor缩放但可视化叠加时没有对齐原始图像尺寸。patch 排列顺序理解错误。直接使用了最后一帧的outputs.attentions但没有确认num_patches (224 // 16)^2 196。检查方式打印cls_attn.shape确认是 196。打印original_img.size确认视觉效果与输入尺寸一致。手动构造一张左白右黑的纯色图运行脚本后观察热力图是否也呈现左右不对称。解决方案在attention_to_image中强制将注意力权重按grid_sizereshape并用双三次插值缩放到原始图片大小。不要使用resize((14, 14))那样会把热力图缩到 patch 网格大小。6.2 模型下载慢或网络异常现象第一次运行时长时间停在下载进度条上或者报Connection error。原因HuggingFace 模型权重需要从远端下载如果网络条件不稳定很容易失败。检查方式查看报错信息中是否包含google/vit-base-patch16-224。检查磁盘缓存目录确认已下载文件是否完整。解决方案可以预先下载模型到本地目录再通过from_pretrained(./vit-local)加载。或者设置环境变量HF_ENDPOINT指向可访问的镜像。这个问题和代码逻辑无关重点是要提前准备好模型权重避免在演示环节卡住。建议在正式使用前先执行一次模型加载确认权重完整落到本机缓存。后续运行就不需要重复下载了。6.3 显存或内存不足现象运行时报CUDA out of memory或RuntimeError: DataLoader worker ...。原因如果同时加载大模型并处理高分辨率图片显存占用会上升。ViT-base 只有 86M 参数本不应该爆显存但如果 batch 过大或processor处理了过大的原始图可能出现问题。解决方案强制使用 CPUDEVICE torch.device(cpu)在预处理前把图片缩放到较小尺寸再交给processor。减少动画帧数不一次生成 144 帧多头动画先取 12 层均值。img img.resize((224, 224), Image.Resampling.BICUBIC)这样能显著降低 Tensor 计算压力且对可视化演示影响不大。6.4 保存 GIF 报错现象ValueError: Cannot save animation: no suitable writer found或ModuleNotFoundError: No module named PIL。原因缺少 GIF writer或者 Pillow 未安装。解决方案首先确认 Pillow 版本python -c import PIL; print(PIL.__version__)然后安装或升级 Pillowpip install --upgrade Pillow保存动画时指定writerpillow这是最稳妥的方式。避免依赖 ImageMagick因为它在不同系统上配置差异较大。6.5 分类概率没有变化现象输出的logits始终是同一个值更换图片后没有明显变化。原因可能模型仍处于训练模式或者没有对输入做归一化。训练模式下dropout 会随机失活部分节点导致输出不稳定没有归一化时图片分布与训练分布差别过大。解决方案检查model.eval()是否已经调用。检查processor是否正确应用到输入。查看outputs.logits的 shape确认不是[1, 1]这种异常输出。这类问题通常不是可视化代码本身的问题而是模型加载或预处理链路的问题需要从输入开始逐层打印排查。7. 最佳实践与扩展方向7.1 可视化代码尽量独立不要把可视化逻辑混在训练或被监控的推理代码里。建议单独维护一个visualize_attention.py文件只接收图像路径和模型路径输出 GIF 或单张热力图。这样做的好处是模型升级后不影响可视化脚本排查问题时也不会误改到核心推理流程。如果确实要在服务中集成可视化建议把注意力提取放在with torch.no_grad()块中并且只提取用户指定层避免每次请求都保存完整 attention 张量否则会带来不必要的内存和性能开销。7.2 多头注意力对比均值热力图适合整体观察也会掩盖单个注意力头的语义。实际调试时可以按下面的方式把一个层内多个注意力头横向排列fig, axes plt.subplots(3, 4, figsize(12, 9)) for head_idx in range(12): row, col divmod(head_idx, 4) axes[row][col].imshow(original_img) axes[row][col].imshow(attention_to_image(attention, original_img.size, head_idx, layer_idx5), alpha0.6) axes[row][col].set_title(fHead {head_idx})这样做可以看到不同头关注边、颜色、高频区域等不同特征。也能帮助判断某一层的注意力是否偏向背景噪声。7.3 将动画嵌入 Web 或 Notebook动画保存为 GIF 后可以直接插入到 Markdown 文档、Web 页面或者 Jupyter Notebook 中。如果需要更灵活的交互可以改用 Gradio、Streamlit 或 Flask 搭建一个简单的可视化服务让用户上传图片后即时生成动画。在 Web 场景下建议把生成结果缓存起来避免同一张图片反复生成。也可以用视频格式代替 GIF因为 GIF 在色彩复杂场景下体积较大播放帧率受限。7.4 与 Grad-CAM 等解释方法结合注意力热力图解释的是 Transformer 内部的注意力分布Grad-CAM 解释的是模型最后输出的梯度对特征图的响应。两者并不矛盾甚至可以互相验证。当注意力热力图和 Grad-CAM 结果一致时说明模型的决策路径相对清晰。当两者出现明显分歧时往往意味着模型依赖了背景信息或纹理信息需要进一步检查数据分布是否与训练集一致。7.5 生产环境落地注意事项如果把 VIT-2动画做成团队内部工具至少要考虑以下几点事项建议模型缓存把权重提前下载到内部文件服务避免每次启动从外网拉取日志记录记录输入图片 hash、模型版本、生成耗时方便回溯资源限制限制输入图片大小增加超时时间避免大图导致内存暴涨异常处理捕获图片损坏、模型推理失败、文件写入失败等异常批量处理使用队列将多张图片逐个生成避免并发推理打满 GPU可视化工具的价值不只是让模型看起来“可解释”更重要的是让开发者从注意力变化中发现问题。如果你发现某一层开始出现大面积背景注意力且最终分类错误这就是一个重要的探测信号。顺着这个信号继续排查数据和标注问题往往比直接换模型更有效。建议下一步试着把多头对比和 Grad-CAM 都加入你的工具箱用多角度证据还原模型真正的决策逻辑。