如何为Faster R-CNN Inception V2生成TSV与Sprite实现TensorBoard可视化?
嘿,我之前折腾过类似的需求,刚好能给你一套落地的方案,亲测有效!要把TensorBoard Projector里的特征点和你的图像、标签对应上,核心就是搞定三个东西:特征TSV、标签TSV和Sprite图像,再给TensorBoard做个配置关联就行。下面一步步来:
一、先搞清楚要提取哪个特征张量
首先你得确认TensorBoard里可视化的36维特征来自模型的哪个节点。毕竟Faster R-CNN是检测模型,特征层不少——你可以打开TensorBoard的Graph标签页,找到对应36维输出的张量名称(比如可能是SecondStageBoxPredictor/Reshape_1:0,或者是ROI池化后的全连接层输出)。记下来这个张量名,后面脚本要用。
二、写脚本提取特征,生成特征TSV和标签TSV
你需要加载训练好的模型,遍历那1024张图像,跑前向传播提取对应特征,同时把标签也存下来。注意图像顺序必须和TensorBoard里的1024个点完全一致,不然对应关系会乱!
给你个参考脚本(基于TensorFlow 1.x,毕竟Faster R-CNN的经典实现多在TF1):
import tensorflow as tf import numpy as np import os # 替换成你的路径和参数 IMAGE_DIR = "你的1024张图像所在文件夹路径" CHECKPOINT_PATH = "训练好的模型ckpt路径(不带后缀)" FEATURE_TENSOR_NAME = "刚才找到的特征张量名,比如SecondStageBoxPredictor/Reshape_1:0" LABEL_FILE = "图像标签文件,每行格式:图像名 标签" OUTPUT_FEATURE_TSV = "features.tsv" OUTPUT_LABEL_TSV = "labels.tsv" # 加载模型图和权重 graph = tf.Graph() with graph.as_default(): saver = tf.train.import_meta_graph(f"{CHECKPOINT_PATH}.meta") with tf.Session() as sess: saver.restore(sess, CHECKPOINT_PATH) # 获取模型的输入张量(Faster R-CNN默认是image_tensor:0) image_tensor = graph.get_tensor_by_name("image_tensor:0") feature_tensor = graph.get_tensor_by_name(FEATURE_TENSOR_NAME) # 按顺序加载图像(务必和TensorBoard里的点顺序一致!) image_paths = [os.path.join(IMAGE_DIR, f) for f in os.listdir(IMAGE_DIR) if f.endswith(('.jpg', '.png'))] # 可以用sort或者手动指定顺序,确保和训练/可视化时的顺序匹配 image_paths.sort() # 读取标签字典 label_dict = {} with open(LABEL_FILE, 'r') as f: for line in f: img_name, label = line.strip().split() label_dict[img_name] = label features_list = [] labels = [] # 图像预处理(和训练时的预处理保持一致!) def preprocess(img_path): img = tf.io.read_file(img_path) img = tf.image.decode_image(img, channels=3) img = tf.image.resize(img, (600, 600)) # 对应你config里的图像尺寸 img = tf.expand_dims(img, 0) img = (img - 127.5) / 127.5 # 归一化,按模型要求调整 return img.eval(session=sess) # 遍历提取特征 for idx, img_path in enumerate(image_paths): print(f"处理第 {idx+1}/{len(image_paths)} 张图") img_input = preprocess(img_path) # 跑前向传播拿特征 raw_features = sess.run(feature_tensor, feed_dict={image_tensor: img_input}) # 因为Faster R-CNN会输出多个目标的特征,这里取置信度最高的那个或者做平均 # 比如取第一个目标的特征,或者全局平均成36维向量 target_feature = np.mean(raw_features, axis=1).flatten() features_list.append(target_feature) # 拿对应标签 img_name = os.path.basename(img_path) labels.append(label_dict.get(img_name, "unknown")) # 保存特征TSV with open(OUTPUT_FEATURE_TSV, 'w') as f: for feat in features_list: f.write('\t'.join(map(str, feat)) + '\n') # 保存标签TSV(第一行是表头,方便TensorBoard识别) with open(OUTPUT_LABEL_TSV, 'w') as f: f.write("Label\n") for label in labels: f.write(f"{label}\n")
三、生成Sprite图像
Sprite图就是把所有小图像拼成一张大网格图,TensorBoard会自动按顺序对应每个特征点。步骤很简单:
from PIL import Image import numpy as np import os IMAGE_DIR = "你的1024张图像所在文件夹路径" SPRITE_IMG_SIZE = 64 # 每个小图的尺寸,比如64x64 OUTPUT_SPRITE = "sprite.png" # 同样,图像顺序要和之前完全一致! image_paths = [os.path.join(IMAGE_DIR, f) for f in os.listdir(IMAGE_DIR) if f.endswith(('.jpg', '.png'))] image_paths.sort() # 计算网格大小(1024是32x32) grid_size = int(np.sqrt(len(image_paths))) # 创建空白Sprite图 sprite = Image.new('RGB', (grid_size * SPRITE_IMG_SIZE, grid_size * SPRITE_IMG_SIZE)) # 逐个粘贴图像 for idx, img_path in enumerate(image_paths): img = Image.open(img_path).convert('RGB') img = img.resize((SPRITE_IMG_SIZE, SPRITE_IMG_SIZE), Image.LANCZOS) # 计算当前小图的位置 x = (idx % grid_size) * SPRITE_IMG_SIZE y = (idx // grid_size) * SPRITE_IMG_SIZE sprite.paste(img, (x, y)) # 保存Sprite图 sprite.save(OUTPUT_SPRITE)
四、配置TensorBoard Projector
现在你有了features.tsv、labels.tsv、sprite.png,接下来要告诉TensorBoard怎么关联它们:
- 在你的TensorBoard日志目录下,创建一个
projector_config.pbtxt文件,内容如下:
embeddings { tensor_name: "你之前用的特征张量名,比如SecondStageBoxPredictor/Reshape_1:0" metadata_path: "labels.tsv" sprite { image_path: "sprite.png" single_image_dim: 64 single_image_dim: 64 } }
- 把
features.tsv、labels.tsv、sprite.png和projector_config.pbtxt全部放到TensorBoard的日志目录里。 - 重启TensorBoard,刷新Projector页面,就能看到每个特征点对应的图像和标签了!
关键提醒
- 顺序!顺序!顺序! 这是最容易踩坑的地方,特征、标签、Sprite图的图像顺序必须完全匹配,不然可视化出来的图像和标签会对不上。
- 特征张量要对应:确保你提取的特征就是TensorBoard里可视化的那36维,不然生成的TSV和Projector里的点不匹配。
- 预处理要一致:图像预处理的方式(尺寸、归一化)要和训练时的config保持一致,不然提取的特征会有偏差。
内容的提问来源于stack exchange,提问作者José Luis
相关产品推荐
相关产品推荐

