如何从列表存储的PyTorch张量中批量提取指定位置元素
解决方案
方法1:堆叠为批量张量操作(推荐)
你当前的samples1是由多个形状为(2,)的张量组成的列表,先通过torch.stack()将整个列表拼接为形状为(num_samples, 2)的批量张量,之后直接做维度切片即可,PyTorch原生张量操作的运算效率远高于循环处理列表:
# 列表堆叠为批量张量,输出形状为 [250, 2] batch_samples = torch.stack(samples1) # 提取所有张量的第一个元素,输出形状为 [250] first_elements = batch_samples[:, 0] # 提取所有张量的第二个元素,输出形状为 [250] second_elements = batch_samples[:, 1]
方法2:列表推导(适合简单场景)
如果仅需要提取结果、不需要后续做批量张量运算,可以直接遍历列表取对应位置的元素:
# 提取结果为列表格式 first_elements = [tensor[0] for tensor in samples1] second_elements = [tensor[1] for tensor in samples1] # 如需结果为张量格式,额外做一次拼接即可 first_elements_tensor = torch.tensor(first_elements) second_elements_tensor = torch.tensor(second_elements)
内容的提问来源于stack exchange,提问作者user785099
相关产品推荐
相关产品推荐

