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

如何基于索引列表高效获取多维训练数据集中的指定样本?

高效从多维数据中按随机索引提取样本

在有限训练数据下评估模型性能时,你需要随机选取原训练数据的20%样本,现有代码已生成随机索引列表filter_indices,以下是几种高效获取对应全维度样本的实现方式:

针对Numpy数组(最常见场景)

如果你的data是Numpy数组(从data.shape=(100,50,50,3)的格式来看大概率是),直接利用Numpy的整数数组索引即可,这是底层优化过的高效操作,会完整保留所有维度:

import numpy as np

# 将索引列表转为Numpy数组(可选,但能进一步提升索引效率)
filter_indices_np = np.array(filter_indices)
selected_data = data[filter_indices_np]

最终selected_data的形状为(len(filter_indices), 50, 50, 3),完全保留原数据的特征维度。

优化:避免重复采样

你的原代码使用random.randrange生成索引,可能会出现重复采样的情况(同一个样本被多次选中)。如果需要无重复的20%样本,可以修改索引生成逻辑:

sample_count = int(data.shape[0] * 0.2)
filter_indices = random.sample(range(data.shape[0]), sample_count)

再配合上面的Numpy索引方法,就能得到无重复的随机样本集。

针对PyTorch/TensorFlow张量

如果data是PyTorch张量:

import torch
# 若data已是张量,直接索引即可
selected_data = data[torch.tensor(filter_indices)]

如果是TensorFlow张量:

import tensorflow as tf
selected_data = tf.gather(data, filter_indices)

内容的提问来源于stack exchange,提问作者Ahmad

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.23 19:45:05