PyTorch Dataset
Dags att repetera PyTorch Datasets!
Innan modellträning kan börja behöver du läsa in data och skicka det till modellen i rätt format. I PyTorch hanteras detta av Datasets och DataLoaders. Vi börjar med att bygga ett PyTorch Dataset för våra data om dricksvattenkvalitet.
I den här övningen definierar du en klass som heter WaterDataset för att läsa in data från en CSV-fil. För att göra det behöver du implementera de tre metoder som PyTorch förväntar sig att ett Dataset ska ha:
.__init__()för att läsa in data,.__len__()för att returnera datamängdens storlek,.__getitem()__för att extrahera särdrag och etikett för ett enskilt sampel.
Följande importer som du behöver har redan gjorts åt dig:
import pandas as pd
from torch.utils.data import Dataset
Den här övningen är en del av kursen
Fördjupad djupinlärning med PyTorch
Interaktiv övning med praktiskt arbete
Testa den här övningen genom att slutföra den här exempelkoden.
class WaterDataset(Dataset):
def __init__(self, csv_path):
super().__init__()
# Load data to pandas DataFrame
df = ____
# Convert data to a NumPy array and assign to self.data
____ = ____.____