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

如何将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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.20 05:42:40