使用卷積神經網路進行影像分類

利用 PyTorch 與 MNIST 資料集,在 AI 執行時訓練卷積神經網路(CNN)用於影像分類。 MNIST 包含 70,000 張手寫數字(0-9)的灰階影像,非常適合學習影像分類技術。

您將瞭解如何:

  • 用 A10G GPU 將你的筆電連接到無伺服器 GPU 運算
  • 定義一個簡單的卷積神經網路架構
  • 用單一 GPU 訓練模型,並將指標記錄到 MLflow
  • 將模型檢查點儲存到 Unity 目錄卷
  • 載入並評估訓練好的模型

注意

此範例需要 Databricks AI 環境版本 5 或以上。

連接到無伺服器的 GPU 運算

這台筆記型電腦需要 GPU 來有效訓練神經網路。 請依照以下步驟連接無伺服器 GPU 運算:

  1. 點選筆記本頂端的 「連接 」下拉選單。
  2. 選擇 無伺服器 GPU。
  3. 打開筆記本右側的 環境 側面板。
  4. 在此示範中,將 Accelerator 設為 1xA10。
  5. 從環境下拉選單選擇 AI v6。
  6. 選擇 「應用」 並點擊 「確認 」,將此環境套用到你的筆記本。

欲了解更多資訊,請參閱 無伺服器 GPU 運算。

配置檢查點儲存位置

以下儲存格建立 元件參數 ,以指定模型檢查點在 Unity 目錄中儲存的位置。 這些參數定義:

  • uc_catalog:Unity Catalog 目錄名稱
  • uc_schema:目錄中的結構(資料庫)
  • uc_volume:用來儲存檢查點檔案的磁碟卷
  • uc_model_name:該磁碟區內此特定型號的子目錄

這些值在筆記本中被用來構建檢查點路徑: /Volumes/{uc_catalog}/{uc_schema}/{uc_volume}/{uc_model_name}

以下儲存格使用佔位值作為預設值。 用筆記本頂端的小工具更新數值。 或者,直接更新下一個儲存格的預設值。

dbutils.widgets.text("uc_catalog", "main")
dbutils.widgets.text("uc_schema", "default")
dbutils.widgets.text("uc_volume", "checkpoints")
dbutils.widgets.text("uc_model_name", "cnn_mnist")

定義卷積神經網路

以下單元定義了一種簡單的 CNN 影像分類架構。 該網絡由以下組成:

  • 兩層卷積層,並使用最大池化來從影像中提取特徵
  • 兩層完全連接的結構用於分類擷取的特徵
  • 使用dropout層以防止過擬合

程式碼還定義了輔助類別,用於在 Unity Catalog 卷中檢查模型和優化器狀態,以及為分散式訓練(多 GPU 情境)建立的函式。

此實作改編自 Horovod PyTorch MNIST 範例。

import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.distributed as dist
import torch.distributed.checkpoint as dcp
from datetime import timedelta
import os

from torch.distributed.checkpoint.state_dict import get_state_dict, set_state_dict
from torch.distributed.checkpoint.stateful import Stateful

class Net(nn.Module):
    def __init__(self):
        super(Net, self).__init__()
        self.conv1 = nn.Conv2d(1, 10, kernel_size=5)
        self.conv2 = nn.Conv2d(10, 20, kernel_size=5)
        self.conv2_drop = nn.Dropout2d()
        self.fc1 = nn.Linear(320, 50)
        self.fc2 = nn.Linear(50, 10)

    def forward(self, x):
        x = F.relu(F.max_pool2d(self.conv1(x), 2))
        x = F.relu(F.max_pool2d(self.conv2_drop(self.conv2(x)), 2))
        x = x.view(-1, 320)
        x = F.relu(self.fc1(x))
        x = F.dropout(x, training=self.training)
        x = self.fc2(x)
        return F.log_softmax(x, dim=1)

UC_CATALOG = dbutils.widgets.get("uc_catalog")
UC_SCHEMA = dbutils.widgets.get("uc_schema")
UC_VOLUME = dbutils.widgets.get("uc_volume")
UC_MODEL_NAME = dbutils.widgets.get("uc_model_name")

# Ensure that the UC Volume directory exists first
CHECKPOINT_DIR = f"/Volumes/{UC_CATALOG}/{UC_SCHEMA}/{UC_VOLUME}/{UC_MODEL_NAME}"

class AppState(Stateful):
    """This is a useful wrapper for checkpointing the Application State. Since this object is compliant
    with the Stateful protocol, DCP will automatically call state_dict/load_stat_dict as needed in the
    dcp.save/load APIs.

    Note: We take advantage of this wrapper to hande calling distributed state dict methods on the model
    and optimizer.
    """

    def __init__(self, model, optimizer=None):
        self.model = model
        self.optimizer = optimizer

    def state_dict(self):
        # this line automatically manages FSDP FQN's, as well as sets the default state dict type to FSDP.SHARDED_STATE_DICT
        model_state_dict, optimizer_state_dict = get_state_dict(self.model, self.optimizer)
        return {
            "model": model_state_dict,
            "optim": optimizer_state_dict
        }

    def load_state_dict(self, state_dict):
        # sets our state dicts on the model and optimizer, now that we've loaded
        set_state_dict(
            self.model,
            self.optimizer,
            model_state_dict=state_dict["model"],
            optim_state_dict=state_dict["optim"]
        )

def setup():
    rank = int(os.environ["RANK"])
    world_size = int(os.environ["WORLD_SIZE"])
    # Shorter timeouts help surface failures quickly instead of hanging
    dist.init_process_group(
        backend="nccl",
        timeout=timedelta(seconds=120),
        init_method="env://",
        rank=rank,
        world_size=world_size,
    )
    torch.cuda.set_device(int(os.environ.get("LOCAL_RANK", 0)))
    dist.barrier()
    if rank == 0:
        print("PG up; all ranks reached barrier")


def cleanup():
    try:
        dist.barrier()
    finally:
        dist.destroy_process_group()

設定訓練參數

以下單元設定訓練的超參數:

  • batch_size:每次訓練迭代處理的影像數量
  • num_epochs: 訓練資料集中完成的次數
  • momentum:SGD 優化器的動量因子
  • log_interval:訓練進度記錄頻率
# Specify training parameters
batch_size = 100
num_epochs = 5
momentum = 0.5
log_interval = 100

定義訓練迴路

以下單元定義 train_one_epoch 函數,該函數:

  • 迭代處理批次訓練資料
  • 執行前向與反向傳播
  • 利用優化器更新模型權重
  • 定期將訓練損失記錄到 MLflow
def train_one_epoch(model, device, data_loader, optimizer, epoch):
    model.train()
    for batch_idx, (data, target) in enumerate(data_loader):
        data, target = data.to(device), target.to(device)
        optimizer.zero_grad()
        output = model(data)
        loss = F.nll_loss(output, target)
        loss.backward()
        optimizer.step()
        if batch_idx % log_interval == 0:
            print('Train Epoch: {} [{}/{} ({:.0f}%)]\tLoss: {:.6f}'.format(
                epoch, batch_idx * len(data), len(data_loader) * len(data),
                100. * batch_idx / len(data_loader), loss.item()))
            # Log metrics
            mlflow.log_metric('loss', loss.item(), step=epoch * len(data_loader) + batch_idx)

在單一 GPU 上訓練模型

以下儲存格定義了主要的訓練函數:

  • 載入 MNIST 訓練資料集
  • 初始化模型與優化器
  • 訓練指定紀元數量的模型
  • 每個紀元後,檢查點會儲存到 Unity 目錄卷
  • 將指標記錄到 MLflow 以進行實驗追蹤
import mlflow
import torch.optim as optim
from torchvision import datasets, transforms

def train(learning_rate):

  with mlflow.start_run() as run:
    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')

    train_dataset = datasets.MNIST(
      'data',
      train=True,
      download=True,
      transform=transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,))]))
    data_loader = torch.utils.data.DataLoader(train_dataset, batch_size=batch_size, shuffle=True)

    model = Net().to(device)
    optimizer = optim.SGD(model.parameters(), lr=learning_rate, momentum=momentum)
    with torch.no_grad():
      input_example, _ = next(iter(data_loader))
      output_example = model(input_example.to(device))

    for epoch in range(1, num_epochs + 1):
      train_one_epoch(model, device, data_loader, optimizer, epoch)

      state_dict = { "app": AppState(model, optimizer) }
      dcp.save(state_dict, checkpoint_id=CHECKPOINT_DIR)
      print(f"saved checkpoint to {CHECKPOINT_DIR}")

執行訓練函數

以下儲存格以學習率 0.001 執行 train 函數。 培訓流程將:

  • 下載MNIST資料集(如果尚未存在於緩存中)
  • 訓練模型針對五個紀元
  • 顯示訓練進度與損失值
  • 將模型檢查點儲存到 Unity 目錄卷
  • 將指標記錄到 MLflow

訓練通常只需幾分鐘,在 A10G GPU 上進行。

train(learning_rate = 0.001)

載入並評估訓練好的模型

訓練完成後,你可以從檢查點載入模型,並評估其在測試資料集上的表現。

以下儲存格定義了一個 test 函數:

  • 從 Unity 目錄卷檢查點載入模型狀態
  • 下載 MNIST 測試資料集
  • 根據測試資料評估模型
  • 計算並顯示平均測試損耗
def test():
  # Load model state from checkpoint using dcp
  model = Net()
  optimizer = optim.SGD(model.parameters(), lr=0.001, momentum=momentum)
  app_state = AppState(model, optimizer)
  state_dict = { "app": app_state }
  dcp.load(state_dict, checkpoint_id=CHECKPOINT_DIR)
  model.eval()

  device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
  model.to(device)
  test_dataset = datasets.MNIST(
    'data',
    train=False,
    download=True,
    transform=transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,))]))
  data_loader = torch.utils.data.DataLoader(test_dataset)

  test_loss = 0
  for data, target in data_loader:
      data, target = data.to(device), target.to(device)
      output = model(data)
      test_loss += F.nll_loss(output, target)

  test_loss /= len(data_loader.dataset)
  print("Average test loss: {}".format(test_loss.item()))

執行評估

以下儲存格執行 test 函數,以評估 MNIST 測試資料集上的訓練模型。 較低的測試損耗表示模型效能較佳。

test()

結論

祝賀! 你已經成功用無伺服器 GPU 運算訓練出一個影像分類模型。 您已學到如何做到以下幾點:

  • 配置並連接無伺服器 GPU 運算
  • 定義卷積神經網路架構
  • 用 PyTorch 訓練模型,並用 MLflow 記錄指標
  • 將模型檢查點儲存到 Unity 目錄卷中
  • 載入並評估已訓練好的模型

斷開 GPU 運算

為了避免不必要的 GPU 使用,請手動斷開 GPU:

  1. 在筆記本頂端選擇已連接
  2. 將滑鼠移至 無伺服器架構
  3. 從下拉選單選擇終止
  4. 選擇 確認 以終止

注意:如果你沒有手動斷線,連線會在 60 分鐘不使用後自動終止。

下一步

探索這些資源,了解更多關於 Databricks 機器學習的資訊:

範例筆記本

使用卷積神經網路進行影像分類

拿筆記本