နားလည်ထားရမယ့် အချက်
Dataset class ၏ တာဝန်မှာ 'sample တစ်ခု ဘယ်လိုရနိုင်သလဲ' ဆိုတဲ့ contract ကို ဖော်ပြပေးရုံသာဖြစ်သည် — __len__() က sample စုစုပေါင်း ဘယ်နှစ်ခုရှိသလဲ ဖော်ပြပေးပြီး __getitem__(idx) က index တစ်ခုအတွက် (image, label) pair ကို ဘယ်လို ပြန်ပေးရမလဲ ဖော်ပြပေးသည်။ Image data ကို disk ပေါ်ရှိ file ကနေ load လုပ်ရင်တောင် (ဥပမာ PIL.Image.open ဖြင့်) ဒီ interface တစ်ခုတည်းကိုပဲ implement လုပ်ရုံနဲ့ လုံလောက်ပြီး၊ data ကို ဘယ်ကနေ ရရှိသည်ဆိုတာနဲ့ DataLoader က ဘယ်လို consume လုပ်သည်ဆိုတာကို completely decouple လုပ်ပေးထားသည်။ ဒီနေရာမှာ synthetic in-memory tensor တွေကို wrap လုပ်ခြင်းက file I/O ကင်းလွတ်စေသော်လည်း interface tချင်း အတူတူပင်ဖြစ်သည်။
DataLoader ကတော့ ဒီ Dataset ကို consume လုပ်ပြီး batching, shuffling, (နောက်ထပ် worker process များဖြင့် parallel loading) စတဲ့ practical concern တွေကို ကိုင်တွယ်ပေးသည်။ shuffle=True ထားလိုက်ခြင်းက epoch တစ်ခုစီတိုင်း sample order ကို ပြောင်းလဲပေးသောကြောင့် model သည် data ၏ sequential ordering ကို memorize လုပ်မိသွားခြင်းကို ကာကွယ်ပေးသည် — ဥပမာ dataset ကို class အလိုက် sort ထားလိုက်လျှင် shuffle မလုပ်ပါက batch တစ်ခုစီသည် class တစ်ခုတည်းသာ ပါဝင်နိုင်ပြီး training instability ဖြစ်စေနိုင်သည်။ DataLoader က iterator interface (for batch in loader) ကို ပေးထားသောကြောင့် training loop ကို dataset size ဘယ်လောက်ကြီးမကြီးနှင့် သီးခြားစီ ရေးနိုင်စေသည်။
လက်တွေ့ scenario နဲ့ ချိတ်ကြည့်မယ်
Tutorial Platform တွင် lesson screenshot preview image များကို 'duplicate/near-duplicate' ဟုတ်မဟုတ် detect လုပ်မည့် internal model တစ်ခု train လုပ်နေသည်ဆိုပါစို့။ Screenshot image များကို disk ပေါ်တွင် တစ်ခုချင်းစီ save မလုပ်မီ preprocessing pipeline အတွင်း in-memory tensor အဖြစ်ပြောင်းထားပြီးဖြစ်ပါက ဒီ lesson ကဲ့သို့ custom Dataset class ကို ဒီ in-memory tensor များပေါ် တိုက်ရိုက် wrap လုပ်၍ file I/O overhead လုံးဝမလိုအပ်ဘဲ DataLoader ကို တည်ဆောက်နိုင်သည်။
အတူတူ စမ်းရေးကြည့်မယ်
import torch
from torch.utils.data import Dataset, DataLoader
class FakeImageDataset(Dataset):
def __init__(self, num_samples=100, num_classes=10):
# Synthetic in-memory data: no files touched
self.images = torch.randn(num_samples, 3, 32, 32)
self.labels = torch.randint(0, num_classes, (num_samples,))
def __len__(self):
return len(self.images)
def __getitem__(self, idx):
return self.images[idx], self.labels[idx]
dataset = FakeImageDataset(num_samples=100, num_classes=10)
loader = DataLoader(dataset, batch_size=16, shuffle=True)
images, labels = next(iter(loader))
print(f"batch images shape: {tuple(images.shape)}")
print(f"batch labels shape: {tuple(labels.shape)}")
print(f"dataset length: {len(dataset)}")
'batch images shape: (16, 3, 32, 32)', 'batch labels shape: (16,)', 'dataset length: 100' ဟု line သုံးကြောင်း print ထုတ်မည် — batch_size=16 ကို သတ်မှတ်ထားသောကြောင့် ပထမ batch တွင် sample ၁၆ ခု ပါဝင်ပြီး dataset စုစုပေါင်း length ၁၀၀ ဖြစ်သည်။၅ မိနစ် စမ်းကြည့်
FakeImageDataset ထဲသို့ transform parameter (optional callable) တစ်ခု ထပ်ထည့်ပြီး __getitem__ အတွင်း image ကို ပြန်ပေးမီ transform ရှိလျှင် apply လုပ်ရန် ပြင်ဆင်ကြည့်ပါ၊ ပြီးလျှင် lambda x: x * 2 ကဲ့သို့ simple transform တစ်ခုနှင့် dataset ကို instantiate လုပ်ပြီး output values ပြောင်းလဲသွားခြင်း ရှိမရှိ စစ်ဆေးပါ။
သတိလေးတစ်ချက်
__getitem__ အတွင်း self.images[idx] အစား self.images ကို တစ်ခါတည်း slice ခွဲပေးမိလျှင် (ဥပမာ self.images[idx:idx+1]) batch dimension ပါလာသောကြောင့် DataLoader ၏ default collate function က unexpected extra dimension တစ်ခု ထပ်ဖြည့်ပေးနိုင်သည်။
num_samples ထက် ကြီးသော batch_size သတ်မှတ်ပြီး drop_last=True ထားခဲ့လျှင် batch တစ်ခုမှ မရရှိတော့ဘဲ next(iter(loader)) သည် StopIteration error တက်နိုင်သည်။
PyTorch Tutorials — Writing Custom Datasets, DataLoaders and Transforms — Computer Vision