Tensorflow TypeError:仅整数标量数组可转换为标量索引
问题分析与解决
错误根源
你的text_encoder输出的query_embedding带有额外的128维度,导致最终results形状为(1,128,9),对应函数返回的结构是包含1个元素的列表,该元素又是包含128个元素的列表,每个元素是9个图片路径的列表。
两个直接触发错误的点:
- 函数内部的列表推导式中,
idx是numpy数组类型,而image_paths是Python列表——列表仅支持整数标量索引,用数组索引会抛出TypeError。 - 你执行
find_matches(...)[0]后得到的是128组9路径的集合,后续循环直接取matches[i](i从0到8),实际取到的是一个包含9个路径的列表,传给mpimg.imread()必然出错。
解决方法
方法一:压缩多余维度(推荐,符合单查询取9个匹配的需求)
修改find_matches函数,对text_encoder输出的128维度做聚合(比如取均值,也可根据模型设计选择取第一个元素等方式),将query_embedding从(1,128,256)压缩为(1,256):
def find_matches(image_embeddings, queries, k=9, normalize=True): # 获取查询向量并压缩多余维度 query_embedding = text_encoder(tf.convert_to_tensor(queries)) query_embedding = tf.reduce_mean(query_embedding, axis=1) if normalize: image_embeddings = tf.math.l2_normalize(image_embeddings, axis=1) query_embedding = tf.math.l2_normalize(query_embedding, axis=1) dot_similarity = tf.matmul(query_embedding, image_embeddings, transpose_b=True) results = tf.math.top_k(dot_similarity, k).indices.numpy() # 此时results形状为(1,9),返回结构为[[路径1, 路径2, ..., 路径9]] return [[image_paths[idx] for idx in indices] for indices in results]
修改后,find_matches(...)[0]会直接返回9个图片路径的列表,后续绘图代码无需改动即可正常运行。
方法二:适配现有输出结构并修复索引问题
如果确实需要保留128维度的多组结果,需先将numpy数组索引转为整数适配Python列表,再调整调用代码获取对应组的结果:
第一步:修改find_matches函数
def find_matches(image_embeddings, queries, k=9, normalize=True): query_embedding = text_encoder(tf.convert_to_tensor(queries)) if normalize: image_embeddings = tf.math.l2_normalize(image_embeddings, axis=1) query_embedding = tf.math.l2_normalize(query_embedding, axis=1) dot_similarity = tf.matmul(query_embedding, image_embeddings, transpose_b=True) results = tf.math.top_k(dot_similarity, k).indices.numpy() # 将numpy数组索引转为整数,适配Python列表索引 return [[[image_paths[int(idx)] for idx in group] for group in indices] for indices in results]
第二步:调整调用代码
指定取128组结果中的某一组(示例取第一组):
query = "a family standing next to the ocean on a sandy beach with a surf board" # 取第1组的9个匹配路径 matches = find_matches(image_embeddings, [query], normalize=True)[0][0] plt.figure(figsize=(20, 20)) for i in range(9): ax = plt.subplot(3, 3, i + 1) plt.imshow(mpimg.imread(matches[i])) plt.axis("off")
内容的提问来源于stack exchange,提问作者albert
相关产品推荐
相关产品推荐

