如何使用TensorFlow提取图像特征向量?附相关实现代码片段
嘿,你选的这个TF Hub模型确实很适合提取图像特征向量!我把你的代码补全并整理成更完整的实现,顺便给你提几个关键点:
使用TF Hub ResNet提取图像特征向量
完整代码实现(TensorFlow 1.x兼容版)
def get_2k_descriptor(self, im): # 将输入图像转换为模型要求的float32类型 image = tf.image.convert_image_dtype(im, tf.float32) # 加载TF Hub上的ResNet V2 101特征向量模型 module = hub.Module("https://tfhub.dev/google/imagenet/resnet_v2_101/feature_vector/1") # 自动获取模型期望的输入尺寸(无需硬编码224x224) height, width = hub.get_expected_image_size(module) # 调整图像尺寸到模型要求的大小 resized_image = tf.image.resize_images(image, [height, width]) # 为图像添加batch维度(模型默认接受批量输入) expanded_image = tf.expand_dims(resized_image, 0) # 计算得到特征向量 features = module(expanded_image) # 去除batch维度,得到单张图像的特征向量 feature_vector = tf.squeeze(features) return feature_vector
关键细节提示
- 注意你提到的链接是ResNet V1 101,但代码里用的是ResNet V2 101,两者特征提取逻辑略有差异,建议保持模型链接和代码一致
hub.get_expected_image_size(module)是更通用的写法,不同模型输入尺寸可能不同,用这个方法能避免硬编码出错- 模型要求输入是批量数据,所以必须用
tf.expand_dims给单张图像添加batch维度 - 这个模型输出的特征向量维度是2048,正好对应你想要的“2k descriptor”
TensorFlow 2.x适配版本
如果你的环境是TF2.x,推荐使用更现代的API:
import tensorflow as tf import tensorflow_hub as hub def get_2k_descriptor(im): image = tf.image.convert_image_dtype(im, tf.float32) # TF2.x版本的模型加载方式 model = hub.load("https://tfhub.dev/google/imagenet/resnet_v2_101/feature_vector/1") height, width = hub.get_expected_image_size(model) resized_image = tf.image.resize(image, [height, width]) expanded_image = tf.expand_dims(resized_image, 0) features = model(expanded_image) # 转换为numpy数组方便后续处理 feature_vector = tf.squeeze(features).numpy() return feature_vector
内容的提问来源于stack exchange,提问作者naman kothari
相关产品推荐
相关产品推荐

