Xây dựng và triển khai mạng nơ-ron tích chập (CNN) phân loại hình ảnh bằng PyTorch

1. Xử lý và chuẩn bị dữ liệu hình ảnh

Trong quy trình huấn luyện mạng nơ-ron tích chập (CNN), bước tiền xử lý dữ liệu đóng vai trò quyết định. PyTorch cung cấp thư viện torchvision để hỗ trợ việc biến đổi và tải dữ liệu một cách hiệu quả.

Cấu hình các bước tiền xử lý (Transforms)

Để đảm bảo các hình ảnh đầu vào có cùng kích thước và định dạng mà mô hình yêu cầu, chúng ta sử dụng transforms.Compose. Dưới đây là cách thiết lập bộ biến đổi dữ liệu:

from torchvision import transforms

# Định nghĩa các bước xử lý cho tập huấn luyện và tập kiểm thử
img_pipelines = {
    'train': transforms.Compose([
        transforms.Resize((64, 64)), # Thay đổi kích thước ảnh về 64x64
        transforms.ToTensor(),       # Chuyển đổi ảnh sang Tensor và chuẩn hóa về [0, 1]
    ]),
    'valid': transforms.Compose([
        transforms.Resize((64, 64)),
        transforms.ToTensor(),
    ])
}
  • transforms.Resize: Đưa toàn bộ ảnh về một kích thước cố định để khớp với lớp đầu vào của mạng CNN.
  • transforms.ToTensor: Chuyển đổi định dạng ảnh (PIL hoặc Numpy) thành torch.FloatTensor.

Sử dụng ImageFolder để quản lý nhãn tự động

Lớp datasets.ImageFolder là một công cụ mạnh mẽ khi dữ liệu được tổ chức theo cấu trúc thư mục, trong đó tên thư mục chính là tên nhãn của lớp đó.

from torchvision import datasets
import os

data_path = './dataset'
# Khởi tạo dataset từ cấu trúc thư mục
data_storage = {
    x: datasets.ImageFolder(root=os.path.join(data_path, x),
                            transform=img_pipelines[x])
    for x in ['train', 'valid']
}

Cấu trúc thư mục chuẩn nên được thiết lập như sau:

dataset/
    train/
        class_A/
            img1.jpg
            img2.jpg
        class_B/
            img3.jpg
    valid/
        ...

Khởi tạo DataLoader để nạp dữ liệu theo Batch

Để tối ưu bộ nhớ và tốc độ huấn luyện, dữ liệu cần được chia thành các nhóm nhỏ (batches).

import torch

data_loaders = {
    x: torch.utils.data.DataLoader(dataset=data_storage[x],
                                   batch_size=8,
                                   shuffle=True)
    for x in ['train', 'valid']
}

2. Cấu trúc mạng CNN hai lớp cơ bản

Khi xây dựng mô hình bằng PyTorch, nn.Sequential thường được sử dụng để đóng gói các lớp có tính tuần tự, giúp mã nguồn gọn gàng hơn.

Mạng CNN thông thường sẽ bao gồm:

  • Lớp tích chập (Conv2d): Trích xuất các đặc trưng không gian từ ảnh.
  • Lớp kích hoạt (ReLU): Thêm tính phi tuyến cho mô hình.
  • Lớp gộp (MaxPool2d): Giảm kích thước bản đồ đặc trưng, giữ lại các thông tin quan trọng nhất.
  • Lớp Dropout: Ngăn chặn hiện tượng quá khớp (overfitting).
  • Lớp kết nối đầy đủ (Linear): Phân loại dựa trên các đặc trưng đã trích xuất.

3. Các lỗi thường gặp và giải pháp khắc phục

Lỗi tương thích phiên bản Python (collections.Iterable)

Trên các phiên bản Python 3.10 trở lên, thuộc tính Iterable đã được di chuyển vào collections.abc. Nếu gặp lỗi AttributeError: module 'collections' has no attribute 'Iterable', bạn có thể xử lý nhanh bằng cách:

import collections
if not hasattr(collections, 'Iterable'):
    import collections.abc
    collections.Iterable = collections.abc.Iterable

Lỗi không khớp kích thước Tensor (Dimension Mismatch)

Lỗi ValueError: Expected input batch_size to match target batch_size thường xảy ra do kích thước đầu ra của lớp tích chập cuối cùng không khớp với số lượng neuron đầu vào của lớp Linear (Fully Connected).

Cách khắc phục: Tính toán chính xác kích thước bản đồ đặc trưng sau các lớp Pooling hoặc sử dụng một bản in thử print(x.shape) ngay trước lớp Flatten để điều chỉnh tham số in_features của lớp Linear.

Hỗ trợ định dạng hình ảnh

PyTorch hỗ trợ hầu hết các định dạng phổ biến như .jpg, .png, .bmp. Tuy nhiên, với các định dạng đặc thù như .tiff, đôi khi thư viện Pillow (backend của torchvision) cần được cập nhật hoặc ảnh cần được chuyển đổi định dạng trước khi đưa vào luồng xử lý để tránh lỗi nạp tệp.

Thẻ: PyTorch CNN computer vision deep learning Image Classification

Đăng vào ngày 24 tháng 8 lúc 20:11