如何将np.int16类型NumPy数组转换为torch.ShortTensor以节省内存?
解决方法
首先明确:torch.from_numpy默认会保留原NumPy数组的dtype,如果你得到了float64类型的张量,大概率是你的NumPy数组在转换前已经被隐式转为float64(比如执行过除法、浮点运算等操作)。先确认你的数组确实是np.int16类型:
print(array.dtype) # 正常输出应为int16
正确转换方式
1. 直接使用torch.from_numpy(数组本身是np.int16时)
如果数组已经是np.int16,torch.from_numpy会直接生成torch.int16类型的张量(即ShortTensor):
import torch import numpy as np # 示例np.int16数组 np_array = np.array([1, 2, 3], dtype=np.int16) torch_tensor = torch.from_numpy(np_array) print(torch_tensor.dtype) # 输出torch.int16
2. 强制转换(数组已变为浮点类型时)
如果数组因操作变为float64,可通过以下方式强制转为torch.int16:
- 方法一:用
torch.tensor直接指定dtype
torch_tensor = torch.tensor(np_array, dtype=torch.int16)
- 方法二:先修正NumPy数组的dtype再转换
np_int16_array = np_array.astype(np.int16) torch_tensor = torch.from_numpy(np_int16_array)
- 方法三:转换后再指定dtype
torch_tensor = torch.from_numpy(np_array).to(torch.int16)
Dataset中的实际应用示例
from torch.utils.data import Dataset class CustomDataset(Dataset): def __init__(self, data_paths): self.data_paths = data_paths def __getitem__(self, idx): # 加载NumPy数组 np_array = np.load(self.data_paths[idx]) # 确保数组为np.int16类型后转换 if np_array.dtype != np.int16: np_array = np_array.astype(np.int16) tensor = torch.from_numpy(np_array) return tensor def __len__(self): return len(self.data_paths)
注意:
torch.from_numpy会与原NumPy数组共享内存,若不需要共享,可改用torch.tensor()或在转换后调用.clone()。
内容的提问来源于stack exchange,提问作者Intrastellar Explorer
相关产品推荐
相关产品推荐

