按索引分组ndarray元素:图像嵌入数组的位置分组方法咨询
解决方法
你可以通过NumPy的索引或轴操作轻松实现这个需求,以下是两种高效的实现方式:
方法一:直接索引提取(最直观)
假设你的嵌入数组名为embeddings,形状为(1000, 512, 256),直接遍历每个嵌入位置,提取所有图像对应位置的向量:
import numpy as np # 假设embeddings是你的原始嵌入数组 observation_groups = [] for idx in range(512): # 提取所有图像的第idx个嵌入,得到形状为(1000, 256)的数组 group = embeddings[:, idx, :] observation_groups.append(group)
最终observation_groups是一个包含512个元素的列表,每个元素对应一组观测数据,形状为(1000, 256)。
方法二:转置+拆分(更高效)
通过转置调整数组轴的顺序,再拆分得到所有观测组:
import numpy as np # 转置数组,将嵌入位置轴移到最前面,得到形状(512, 1000, 256)的数组 transposed_embeddings = embeddings.transpose(1, 0, 2) # 沿第0轴拆分,得到512个(1, 1000, 256)的子数组 observation_groups = np.split(transposed_embeddings, 512, axis=0) # 去除每个子数组多余的维度,得到(1000, 256)的观测组 observation_groups = [group.squeeze(axis=0) for group in observation_groups]
原理说明
原始数组的轴定义为(图像数量, 嵌入位置, 向量维度),通过[:, idx, :]索引可以直接筛选出所有图像中第idx个位置的嵌入向量,正好符合每个观测组的需求。
内容的提问来源于stack exchange,提问作者Djanger
相关产品推荐
相关产品推荐

