You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Tensorflow TypeError:仅整数标量数组可转换为标量索引

问题分析与解决

错误根源

你的text_encoder输出的query_embedding带有额外的128维度,导致最终results形状为(1,128,9),对应函数返回的结构是包含1个元素的列表,该元素又是包含128个元素的列表,每个元素是9个图片路径的列表。

两个直接触发错误的点:

  1. 函数内部的列表推导式中,idx是numpy数组类型,而image_paths是Python列表——列表仅支持整数标量索引,用数组索引会抛出TypeError。
  2. 你执行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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.24 17:48:20