Kom igångKom igång gratis

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

Visa kurs

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
        ____ = ____.____
Redigera och kör kod