在 CentOS 上使用 PyTorch 导入数据,方式跟在其他 Linux 系统或 Windows 上基本一致,核心是用 Dataset + DataLoader。下面从环境准备 → 常见数据导入方式 → 示例给你一个完整说明。
python3 --version
如果没有:
sudo yum install python3 python3-pip
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__)
torch.utils.data.Dataset(最常用)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
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)
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
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)
pandas 分块for chunk in pd.read_csv("big.csv", chunksize=10000):
...
ImageFolder + DataLoader + num_workers>0DataLoader(..., num_workers=4)
⚠️ CentOS 上如果报错:
BrokenPipeError
解决:
if __name__ == "__main__":
main()
sudo yum install libgomp
pd.read_csv("data.csv", encoding="utf-8")
chmod -R 755 /data
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
如果你能告诉我:
我可以直接给你最适合 CentOS 的导入方案。