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ư ModelCheckpoint và EarlyStopping, 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.
OsmanyaDataModule và ImageClassifier đã đượ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
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
ModelCheckpointvàEarlyStopping.
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)