定義資料集、載入器
## 1. Dataset 與 DataLoader 的用途
### `Dataset`(資料集)
`Dataset` 用來描述「資料本身」以及「如何取得每一筆資料」。
在 PyTorch 中,通常自訂資料集需要繼承:
```python
torch.utils.data.Dataset
```
並實作兩個方法:
```python
__len__() # 回傳資料總筆數
__getitem__() # 根據索引取得一筆資料
```
一筆資料通常包含:
- 輸入資料:`input`、`x`
- 對應的答案或標籤:`target`、`y`
---
### `DataLoader`(資料載入器)
`DataLoader` 負責將 `Dataset` 中的資料載入模型訓練流程,常見功能包括:
- 將資料分成批次(batch)
- 隨機打亂資料
- 平行載入資料
- 逐批產生輸入與標籤
基本語法:
```python
from torch.utils.data import DataLoader
dataloader = DataLoader(
dataset,
batch_size=32,
shuffle=True
)
```
常用參數:
| 參數 | 說明 |
|---|---|
| `dataset` | 要載入的資料集 |
| `batch_size` | 每一批包含幾筆資料 |
| `shuffle` | 是否在每個 epoch 開始時打亂資料 |
| `num_workers` | 使用幾個程序平行載入資料,通常可設為 `0` 或更大值 |
---
## 2. 自訂 Dataset 的基本語法
```python
import torch
from torch.utils.data import Dataset
class MyDataset(Dataset):
def __init__(self):
# 建立或讀取資料
pass
def __len__(self):
# 回傳資料筆數
pass
def __getitem__(self, index):
# 回傳第 index 筆資料
pass
```
---
## 3. 簡單整數輸入與期待值資料集範例
以下建立一個簡單資料集:
y = 2x + 1
| 輸入 `x` | 期待值 `y` |
|---:|---:|
| 0 | 1 |
| 1 | 3 |
| 2 | 5 |
| 3 | 7 |
### 完整程式碼
```python
import torch
from torch.utils.data import Dataset, DataLoader
class IntegerDataset(Dataset):
def __init__(self, start=0, end=10):
"""
建立整數輸入資料:
x = start, start+1, ..., end-1
y = 2x + 1
"""
self.x = list(range(start, end))
self.y = [2 * value + 1 for value in self.x]
def __len__(self):
"""回傳資料總筆數"""
return len(self.x)
def __getitem__(self, index):
"""回傳第 index 筆輸入與期待值"""
x_value = torch.tensor(
[self.x[index]],
dtype=torch.float32
)
y_value = torch.tensor(
[self.y[index]],
dtype=torch.float32
)
return x_value, y_value
# 建立 Dataset
dataset = IntegerDataset(start=0, end=10)
# 建立 DataLoader
dataloader = DataLoader(
dataset,
batch_size=3,
shuffle=False
)
# 查看資料集大小
print("資料筆數:", len(dataset))
# 查看單筆資料
x, y = dataset[2]
print("第 3 筆資料:")
print("輸入 x =", x)
print("期待值 y =", y)
# 逐批讀取資料
for batch_x, batch_y in dataloader:
print("batch_x =", batch_x)
print("batch_y =", batch_y)
```
---
## 4. 可能的輸出
```text
資料筆數: 10
第 3 筆資料:
輸入 x = tensor([2.])
期待值 y = tensor([5.])
batch_x = tensor([[0.],
[1.],
[2.]])
batch_y = tensor([[1.],
[3.],
[5.]])
batch_x = tensor([[3.],
[4.],
[5.]])
batch_y = tensor([[7.],
[9.],
[11.]])
...
```
因為設定:
```python
batch_size=3
```
所以每次從 `DataLoader` 取出最多 3 筆資料。總共有 10 筆資料,因此會分成 4 個 batch,最後一批只有 1 筆。
---
簡單來說:
- `Dataset`:定義資料是什麼,以及如何取得單筆資料
- `DataLoader`:決定資料如何以批次方式提供給模型使用
相關學習地圖、教學課程
Python 人工智慧