温馨提示×

centos pytorch怎么导入数据

小樊
37
2026-09-18 12:58:43
栏目: 智能运维

CentOS 上使用 PyTorch 导入数据,方式跟在其他 Linux 系统或 Windows 上基本一致,核心是用 Dataset + DataLoader。下面从环境准备 → 常见数据导入方式 → 示例给你一个完整说明。


一、CentOS 上 PyTorch 环境准备(简要)

1. 安装 Python(推荐 3.8+)

python3 --version

如果没有:

sudo yum install python3 python3-pip

2. 安装 PyTorch(CPU 示例)

pip3 install torch torchvision torchaudio

GPU 版本(如 CUDA 11.8):

pip3 install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118

验证:

import torch
print(torch.__version__)

二、PyTorch 导入数据的核心方式

✅ 方式 1:使用 torch.utils.data.Dataset(最常用)

1)自定义 Dataset(适合 CSV / 图片 / 文本)

from torch.utils.data import Dataset
import pandas as pd

class MyDataset(Dataset):
    def __init__(self, csv_path):
        self.data = pd.read_csv(csv_path)

    def __len__(self):
        return len(self.data)

    def __getitem__(self, idx):
        x = self.data.iloc[idx, :-1].values.astype('float32')
        y = self.data.iloc[idx, -1]
        return x, y

2)使用 DataLoader

from torch.utils.data import DataLoader

dataset = MyDataset("data.csv")
loader = DataLoader(dataset, batch_size=32, shuffle=True)

for x, y in loader:
    print(x.shape, y.shape)

✅ 方式 2:使用 torchvision.datasets(图像数据)

from torchvision import datasets, transforms

transform = transforms.Compose([
    transforms.ToTensor()
])

train_set = datasets.ImageFolder(
    root="/data/train",
    transform=transform
)

loader = DataLoader(train_set, batch_size=64, shuffle=True)

目录结构要求:

/data/train/
    cat/
        xxx.jpg
    dog/
        yyy.jpg

✅ 方式 3:直接加载 NumPy / TXT 数据

import numpy as np
import torch

data = np.loadtxt("data.txt", delimiter=",")
x = torch.tensor(data[:, :-1], dtype=torch.float32)
y = torch.tensor(data[:, -1], dtype=torch.long)

✅ 方式 4:读取大规模数据(CentOS 服务器推荐)

  • CSV 太大:用 pandas 分块
for chunk in pd.read_csv("big.csv", chunksize=10000):
    ...
  • 图片太多ImageFolder + DataLoader + num_workers>0
DataLoader(..., num_workers=4)

⚠️ CentOS 上如果报错:

BrokenPipeError

解决:

if __name__ == "__main__":
    main()

三、CentOS 常见坑

1. 缺少共享库

sudo yum install libgomp

2. 中文路径/编码问题

pd.read_csv("data.csv", encoding="utf-8")

3. 权限问题

chmod -R 755 /data

四、最小完整示例(CSV → PyTorch)

import torch
from torch.utils.data import Dataset, DataLoader
import pandas as pd

class CSVData(Dataset):
    def __init__(self, path):
        self.df = pd.read_csv(path)

    def __len__(self):
        return len(self.df)

    def __getitem__(self, idx):
        x = self.df.iloc[idx, :-1].values.astype('float32')
        y = int(self.df.iloc[idx, -1])
        return torch.tensor(x), torch.tensor(y)

loader = DataLoader(CSVData("train.csv"), batch_size=32, shuffle=True)

for x, y in loader:
    print(x.shape, y.shape)
    break

如果你能告诉我:

  • ✅ 数据类型(图片 / 文本 / CSV / 音频)
  • ✅ 是否在 GPU 服务器
  • ✅ 数据规模(几 MB 还是几十 GB)

我可以直接给你最适合 CentOS 的导入方案

0