View a markdown version of this page

Terapkan SageMaker penyaringan cerdas ke skrip Anda PyTorch - Amazon SageMaker AI

Terjemahan disediakan oleh mesin penerjemah. Jika konten terjemahan yang diberikan bertentangan dengan versi bahasa Inggris aslinya, utamakan versi bahasa Inggris.

Terapkan SageMaker penyaringan cerdas ke skrip Anda PyTorch

Instruksi ini menunjukkan cara mengaktifkan penyar SageMaker ingan cerdas dengan skrip pelatihan Anda.

  1. Konfigurasikan antarmuka SageMaker penyaringan cerdas.

    Pustaka penyaringan SageMaker cerdas menerapkan teknik pengambilan sampel berbasis kerugian ambang relatif yang membantu menyaring sampel dengan dampak yang lebih rendah pada pengurangan nilai kerugian. Algoritma penyaringan SageMaker cerdas menghitung nilai kerugian dari setiap sampel data input menggunakan forward pass, dan menghitung persentil relatifnya terhadap nilai kerugian data sebelumnya.

    Dua parameter berikut adalah apa yang perlu Anda tentukan ke RelativeProbabilisticSiftConfig kelas untuk membuat objek konfigurasi penyaringan.

    • Tentukan proporsi data yang harus digunakan untuk pelatihan ke beta_value parameter.

    • Tentukan jumlah sampel yang digunakan dalam perbandingan dengan loss_history_length parameter.

    Contoh kode berikut menunjukkan pengaturan objek RelativeProbabilisticSiftConfig kelas.

    from smart_sifting.sift_config.sift_configs import ( RelativeProbabilisticSiftConfig LossConfig SiftingBaseConfig ) sift_config=RelativeProbabilisticSiftConfig( beta_value=0.5, loss_history_length=500, loss_based_sift_config=LossConfig( sift_config=SiftingBaseConfig(sift_delay=0) ) )

    Untuk informasi selengkapnya tentang loss_based_sift_config parameter dan kelas terkait, lihat SageMaker modul konfigurasi penyaringan cerdas di bagian referensi SageMaker Smart Sifting Python SDK.

    Ob sift_config jek dalam contoh kode sebelumnya digunakan pada langkah 4 untuk mengatur SiftingDataloader kelas.

  2. (Opsional) Konfigurasikan kelas transformasi batch pengilingan SageMaker cerdas.

    Kasus penggunaan pelatihan yang berbeda memerlukan format data pelatihan yang berbeda. Mengingat berbagai format data, algoritma SageMaker penyaringan cerdas perlu mengidentifikasi cara melakukan penyaringan pada batch tertentu. Untuk mengatasi hal ini, SageMaker smart sifting menyediakan modul transformasi batch yang membantu mengonversi batch menjadi format standar yang dapat disaring secara efisien.

    1. SageMaker smart sifting menangani transformasi batch data pelatihan dalam format berikut: daftar Python, kamus, tupel, dan tensor. Untuk format data ini, SageMaker penyaringan cerdas secara otomatis menangani konversi format data batch, dan Anda dapat melewati sisa langkah ini. Jika Anda melewati langkah ini, pada langkah 4 untuk mengonfigurasiSiftingDataloader, biarkan batch_transforms parameter SiftingDataloader ke nilai defaultnya, yaituNone.

    2. Jika dataset Anda tidak dalam format ini, Anda harus melanjutkan ke sisa langkah ini untuk membuat transformasi batch khusus menggunakanSiftingBatchTransform.

      Dalam kasus di mana dataset Anda tidak dalam salah satu format yang didukung oleh pen SageMaker yaringan cerdas, Anda mungkin mengalami kesalahan. Kesalahan format data tersebut dapat diselesaikan dengan menambahkan batch_transforms parameter batch_format_index atau ke SiftingDataloader kelas, yang Anda atur pada langkah 4. Berikut ini menunjukkan contoh kesalahan karena format data yang tidak kompatibel dan resolusi untuk mereka.

      Pesan Kesalahan Resolusi

      Batch jenis {type(batch)} tidak didukung secara default.

      Kesalahan ini menunjukkan format batch tidak didukung secara default. Anda harus mengimplementasikan kelas transformasi batch kustom, dan gunakan ini dengan menent batch_transforms ukannya ke parameter SiftingDataloader kelas.

      Tidak dapat mengindeks kumpulan jenis {type(batch)}

      Kesalahan ini menunjukkan objek batch tidak dapat diindeks secara normal. Pengguna harus menerapkan transformasi batch khusus dan meneruskannya menggunakan batch_transforms parameter.

      Ukuran batch {batch_size} tidak cocok dengan ukuran dimensi 0 atau dimensi 1

      Kesalahan ini terjadi ketika ukuran batch yang disediakan tidak cocok dengan dimensi ke-0 atau ke-1 dari batch. Pengguna harus menerapkan transformasi batch khusus dan meneruskannya menggunakan batch_transforms parameter.

      Kedua dimensi 0 dan dimensi 1 cocok dengan ukuran batch

      Kesalahan ini menunjukkan bahwa karena beberapa dimensi cocok dengan ukuran batch yang disediakan, informasi lebih lanjut diperlukan untuk menyaring batch. Pengguna dapat memberikan batch_format_index parameter untuk menunjukkan apakah batch dapat diindeks oleh sampel atau fitur. Pengguna juga dapat menerapkan transformasi batch khusus, tetapi ini lebih banyak pekerjaan daripada yang diperlukan.

      Untuk menyelesaikan masalah yang disebutkan di atas, Anda perlu membuat kelas transformasi batch khusus menggunakan SiftingBatchTransform modul. Kelas transformasi batch harus terdiri dari sepasang fungsi transform dan reverse transform. Pasangan fungsi mengonversi format data Anda ke format yang dapat diproses oleh algoritma penyaringan SageMaker cerdas. Setelah Anda membuat kelas transformasi batch, kelas mengembalikan SiftingBatch objek yang akan Anda berikan ke SiftingDataloader kelas pada langkah 4.

      Berikut ini adalah contoh kelas transformasi batch kustom dari SiftingBatchTransform modul.

      • Contoh implementasi transformasi batch daftar kustom dengan penyar SageMaker ingan cerdas untuk kasus di mana potongan dataloader memiliki input, masker, dan label.

        from typing import Any import torch from smart_sifting.data_model.data_model_interface import SiftingBatchTransform from smart_sifting.data_model.list_batch import ListBatch class ListBatchTransform(SiftingBatchTransform): def transform(self, batch: Any): inputs = batch[0].tolist() labels = batch[-1].tolist() # assume the last one is the list of labels return ListBatch(inputs, labels) def reverse_transform(self, list_batch: ListBatch): a_batch = [torch.tensor(list_batch.inputs), torch.tensor(list_batch.labels)] return a_batch
      • Contoh implementasi transformasi batch daftar kustom dengan penyaringan SageMaker cerdas untuk kasus di mana tidak ada label yang diperlukan untuk transformasi terbalik.

        class ListBatchTransformNoLabels(SiftingBatchTransform): def transform(self, batch: Any): return ListBatch(batch[0].tolist()) def reverse_transform(self, list_batch: ListBatch): a_batch = [torch.tensor(list_batch.inputs)] return a_batch
      • Contoh implementasi batch tensor khusus dengan penyaringan SageMaker cerdas untuk kasus di mana potongan pemuat data memiliki input, masker, dan label.

        from typing import Any from smart_sifting.data_model.data_model_interface import SiftingBatchTransform from smart_sifting.data_model.tensor_batch import TensorBatch class TensorBatchTransform(SiftingBatchTransform): def transform(self, batch: Any): a_tensor_batch = TensorBatch( batch[0], batch[-1] ) # assume the last one is the list of labels return a_tensor_batch def reverse_transform(self, tensor_batch: TensorBatch): a_batch = [tensor_batch.inputs, tensor_batch.labels] return a_batch

      Setelah Anda membuat kelas transform SiftingBatchTransform asi batch -implemented, Anda menggunakan kelas ini di langkah 4 untuk menyiapkan kelas. SiftingDataloader Sisa panduan ini mengasumsikan bahwa ListBatchTransform kelas dibuat. Pada langkah 4, kelas ini diteruskan kebatch_transforms.

  3. Buat kelas untuk mengimplementasikan Loss antarmuka pen SageMaker yaringan cerdas. Tutorial ini mengasumsikan bahwa kelas diberi namaSiftingImplementedLoss. Saat menyiapkan kelas ini, kami menyarankan Anda menggunakan fungsi kerugian yang sama dalam loop pelatihan model. Ikuti sublangkah berikut untuk membuat kelas yang di Loss implementasikan SageMaker penyaringan cerdas.

    1. SageMaker smart sifting menghitung nilai kerugian untuk setiap sampel data pelatihan, sebagai lawan dari menghitung nilai kerugian tunggal untuk batch. Untuk memastikan bahwa pen SageMaker yaringan cerdas menggunakan logika perhitungan kerugian yang sama, buat fungsi kerugian yang diimplementasikan secara cerdas menggunakan Loss modul pengayakan SageMaker pintar yang menggunakan fungsi kerugian Anda dan menghitung kerugian per sampel pelatihan.

      Tip

      SageMaker algoritma penyaringan pintar berjalan pada setiap sampel data, bukan pada seluruh batch, jadi Anda harus menambahkan fungsi inisialisasi untuk mengatur fungsi kerugian tanpa PyTorch strategi pengurangan apa pun.

      class SiftingImplementedLoss(Loss): def __init__(self): self.loss = torch.nn.CrossEntropyLoss(reduction='none')

      Ini juga ditunjukkan dalam contoh kode berikut.

    2. Tentukan fungsi kerugian yang menerima original_batch (atau transformed_batch jika Anda telah menyiapkan transformasi batch pada langkah 2) dan PyTorch model. Menggunakan fungsi kerugian yang ditentukan tanpa pengurangan, pen SageMaker yaringan cerdas menjalankan pass maju untuk setiap sampel data untuk mengevaluasi nilai kerugiannya.

    Kode berikut adalah contoh antarmuka yang diimplementasikan smart-sift Loss ing-bernama. SiftingImplementedLoss

    from typing import Any import torch import torch.nn as nn from torch import Tensor from smart_sifting.data_model.data_model_interface import SiftingBatch from smart_sifting.loss.abstract_sift_loss_module import Loss model=... # a PyTorch model based on torch.nn.Module class SiftingImplementedLoss(Loss): # You should add the following initializaztion function # to calculate loss per sample, not per batch. def __init__(self): self.loss_no_reduction = torch.nn.CrossEntropyLoss(reduction='none') def loss( self, model: torch.nn.Module, transformed_batch: SiftingBatch, original_batch: Any = None, ) -> torch.Tensor: device = next(model.parameters()).device batch = [t.to(device) for t in original_batch] # use this if you use original batch and skipped step 2 # batch = [t.to(device) for t in transformed_batch] # use this if you transformed batches in step 2 # compute loss outputs = model(batch) return self.loss_no_reduction(outputs.logits, batch[2])

    Sebelum loop pelatihan mencapai pass maju yang sebenarnya, perhitungan kerugian penyaringan ini dilakukan selama fase pemuatan data pengambilan batch di setiap iterasi. Nilai kerugian individu kemudian dibandingkan dengan nilai kerugian sebelumnya, dan persentil relatifnya diperkirakan per objek yang telah RelativeProbabilisticSiftConfig Anda atur pada langkah 1.

  4. Bungkus pemuat PyTroch data dengan SiftingDataloader kelas SageMaker AI.

    Terakhir, gunakan semua kelas yang diimplementasikan penyaringan SageMaker cerdas yang Anda konfigurasikan pada langkah-langkah sebelumnya ke kelas SiftingDataloder konfigurasi SageMaker AI. Kelas ini adalah pembungkus untuk PyTorch DataLoader. Dengan membungkus PyTorchDataLoader, SageMaker penyaringan cerdas terdaftar untuk dijalankan sebagai bagian dari pemuatan data di setiap iterasi pekerjaan PyTorch pelatihan. Contoh kode berikut menunjukkan penerapan pen SageMaker yaringan data AI ke a PyTorchDataLoader.

    from smart_sifting.dataloader.sift_dataloader import SiftingDataloader from torch.utils.data import DataLoader train_dataloader = DataLoader(...) # PyTorch data loader # Wrap the PyTorch data loader by SiftingDataloder train_dataloader = SiftingDataloader( sift_config=sift_config, # config object of RelativeProbabilisticSiftConfig orig_dataloader=train_dataloader, batch_transforms=ListBatchTransform(), # Optional, this is the custom class from step 2 loss_impl=SiftingImplementedLoss(), # PyTorch loss function wrapped by the Sifting Loss interface model=model, log_batch_data=False )