Bắt đầu ngayBắt đầu miễn phí

Tối ưu hóa quá trình huấn luyện với Lightning

Bằng cách triển khai các kỹ thuật tự động như ModelCheckpointEarlyStopping, bạn sẽ giúp mô hình chọn được bộ tham số hoạt động tốt nhất đồng thời tránh các phép tính không cần thiết.

Bộ dữ liệu, một tập con của Osmanya MNIST, là một trường hợp thực tế nơi các kỹ thuật huấn luyện AI có khả năng mở rộng có thể cải thiện đáng kể hiệu suất và độ chính xác.

OsmanyaDataModuleImageClassifier đã được định nghĩa sẵn cho bạn.

Bài tập này là một phần của khóa học

Mô hình AI có khả năng mở rộng với PyTorch Lightning

Xem khóa học

Hướng dẫn bài tập

  • Import các callback bạn sẽ dùng để lưu checkpoint mô hình và dừng sớm.
  • Huấn luyện mô hình với các callback ModelCheckpointEarlyStopping.

Bài tập tương tác thực hành trực tiếp

Hãy thử làm bài tập này bằng cách hoàn thành đoạn mã mẫu này.

# Import relevant checkpoints
from lightning.pytorch.callbacks import ____, ____

class EvaluatedImageClassifier(ImageClassifier):
    def validation_step(self, batch, batch_idx):
        x, y = batch
        y_hat = self(x)
        acc = (y_hat.argmax(dim=1) == y).float().mean()
        self.log("val_acc", acc)

model = EvaluatedImageClassifier()
data_module = OsmanyaDataModule()
# Train the model with ModelCheckpoint and EarlyStopping checkpoints
trainer = Trainer(____=[____(monitor="val_acc", save_top_k=1), ____(monitor="val_acc", patience=3)])
trainer.fit(model, datamodule=data_module)
Chỉnh sửa và Chạy Mã