Zacznij terazZacznij za darmo

Zbiór danych PyTorch

Czas odświeżyć wiedzę o zbiorach danych w PyTorch!

Zanim rozpocznie się trenowanie modelu, trzeba wczytać dane i przekazać je do modelu w odpowiednim formacie. W PyTorch zajmują się tym klasy Dataset i DataLoader. Zacznijmy od zbudowania zbioru danych PyTorch dla naszych danych o jakości wody.

W tym ćwiczeniu zdefiniujesz klasę WaterDataset, która wczyta dane z pliku CSV. W tym celu zaimplementujesz trzy metody, których PyTorch oczekuje od każdego zbioru danych:

  • .__init__() – do wczytania danych,
  • .__len__() – do zwracania rozmiaru zbioru danych,
  • .__getitem()__ – do wyodrębniania cech i etykiety dla pojedynczej próbki.

Następujące importy zostały już wykonane za ciebie:

import pandas as pd
from torch.utils.data import Dataset

To ćwiczenie jest częścią kursu

Głębokie uczenie z PyTorch – poziom średnio zaawansowany

Zobacz kurs

Interaktywne ćwiczenie praktyczne

Spróbuj tego ćwiczenia, uzupełniając ten przykładowy kod.

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
        ____ = ____.____
Edytuj i uruchom kod