So sánh nn.Module và nn.functional trong PyTorch: Hướng dẫn sử dụng và Best Practices

Tổng quan về nn.Module và nn.functional

Trong PyTorch, việc xây dựng các mạng neural thường liên quan đến hai API chính: nn.Module (thông qua các lớp như nn.Conv2d) và nn.functional (thông qua các hàm như F.conv2d). Mặc dù cả hai đều thực hiện các phép toán cơ bản giống nhau, nhưng cách thức tổ chức và quản lý trạng thái lại có những khác biệt quan trọng.

1. Điểm tương đồng giữa hai API

  • Chức năng tính toán cốt lõi là tương đương. Ví dụ, nn.Conv2d và F.conv2d đều thực hiện phép tích chập.
  • Hiệu suất thực thi và tốc độ tính toán là gần như không có sự khác biệt.

Về bản chất, nn.functional.xxx là các hàm thuần túy, trong khi nn.Xxx là các lớp được đóng gói kế thừa từ nn.Module. Do kế thừa nn.Module, các lớp nn.Xxx sở hữu thêm các phương thức quản lý trạng thái như train(), eval(), state_dict(), và load_state_dict().

2. Những khác biệt cốt lõi

Cơ chế gọi và truyền tham số

nn.Xxx yêu cầu khởi tạo đối tượng và truyền các siêu tham số (như số kênh, kích thước kernel) trước. Sau đó, dữ liệu đầu vào được truyền qua đối tượng đã khởi tạo. Ngược lại, nn.functional.xxx yêu cầu truyền cả dữ liệu đầu vào lẫn các trọng số (weights) và độ chệch (bias) ngay trong lời gọi hàm.

import torch
import torch.nn as nn
import torch.nn.functional as F

# Sử dụng nn.Module
input_tensor = torch.randn(16, 1, 64, 64)
conv_layer = nn.Conv2d(in_channels=1, out_channels=32, kernel_size=5, stride=2)
output_nn = conv_layer(input_tensor)

# Sử dụng nn.functional
kernel_weights = torch.randn(32, 1, 5, 5)
layer_bias = torch.randn(32)
output_f = F.conv2d(input_tensor, kernel_weights, layer_bias, stride=2)

Tích hợp với nn.Sequential

Các lớp nn.Xxx có thể được kết hợp trực tiếp và mượt mà trong nn.Sequential để xây dựng luồng dữ liệu. nn.functional không hỗ trợ điều này vì nó là các hàm độc lập.

feature_extractor = nn.Sequential(
    nn.Conv2d(in_channels=3, out_channels=16, kernel_size=3, padding=1),
    nn.ReLU(inplace=True),
    nn.AvgPool2d(kernel_size=2, stride=2),
    nn.Conv2d(in_channels=16, out_channels=32, kernel_size=3, padding=1),
    nn.ReLU(inplace=True)
)

Quản lý trọng số (Weights) và tham số

Khi sử dụng nn.Xxx, PyTorch tự động khởi tạo và quản lý các trọng số. Với nn.functional, lập trình viên phải tự định nghĩa các trọng số dưới dạng nn.Parameter và truyền chúng thủ công, gây khó khăn cho việc tái sử dụng mã nguồn.

# Định nghĩa mạng sử dụng nn.Module
class VisionNetwork(nn.Module):
    def __init__(self):
        super(VisionNetwork, self).__init__()
        self.block1 = nn.Sequential(
            nn.Conv2d(3, 32, kernel_size=3, padding=1),
            nn.ReLU(),
            nn.MaxPool2d(2)
        )
        self.block2 = nn.Sequential(
            nn.Conv2d(32, 64, kernel_size=3, padding=1),
            nn.ReLU(),
            nn.MaxPool2d(2)
        )
        self.classifier = nn.Linear(64 * 7 * 7, 10)

    def forward(self, pixel_data):
        out = self.block1(pixel_data)
        out = self.block2(out)
        out = out.view(out.size(0), -1)
        return self.classifier(out)

# Định nghĩa mạng tương đương sử dụng nn.functional
class VisionNetworkFunctional(nn.Module):
    def __init__(self):
        super(VisionNetworkFunctional, self).__init__()
        self.w1 = nn.Parameter(torch.randn(32, 3, 3, 3))
        self.b1 = nn.Parameter(torch.randn(32))
        self.w2 = nn.Parameter(torch.randn(64, 32, 3, 3))
        self.b2 = nn.Parameter(torch.randn(64))
        self.w3 = nn.Parameter(torch.randn(64 * 7 * 7, 10))
        self.b3 = nn.Parameter(torch.randn(10))

    def forward(self, pixel_data):
        out = F.conv2d(pixel_data, self.w1, self.b1, padding=1)
        out = F.relu(out)
        out = F.max_pool2d(out, 2)
        
        out = F.conv2d(out, self.w2, self.b2, padding=1)
        out = F.relu(out)
        out = F.max_pool2d(out, 2)
        
        out = out.view(out.size(0), -1)
        return F.linear(out, self.w3, self.b3)

3. Hành vi đặc biệt của Dropout trong hai API

Một điểm khác biệt cực kỳ quan trọng liên quan đến Dropout. Các lớp như nn.Dropout tự động nhận biết trạng thái của mô hình (training hay evaluation). Khi gọi model.eval(), dropout sẽ tự động bị vô hiệu hóa. Ngược lại, F.dropout sẽ luôn hoạt động trừ khi bạn truyền thủ công cờ training.

class DropoutModelA(nn.Module):
    def __init__(self):
        super().__init__()
        self.drop_layer = nn.Dropout(p=0.3)
        
    def forward(self, x):
        return self.drop_layer(x)

class DropoutModelB(nn.Module):
    def __init__(self):
        super().__init__()
        
    def forward(self, x):
        return F.dropout(x, p=0.3) # Luôn active nếu không truyền training=False

net_a = DropoutModelA()
net_b = DropoutModelB()
dummy_input = torch.ones(5)

net_a.eval()
net_b.eval()

# net_a sẽ trả về giá trị giữ nguyên (dropout disabled)
# net_b vẫn sẽ loại bỏ ngẫu nhiên các phần tử (dropout still active)

Nếu bắt buộc phải sử dụng F.dropout, cần phải truyền biến self.training để đồng bộ trạng thái với mô hình:

class DropoutModelC(nn.Module):
    def __init__(self):
        super().__init__()
        
    def forward(self, x):
        return F.dropout(x, p=0.3, training=self.training)

4. Hướng dẫn lựa chọn API phù hợp

PyTorch khuyến nghị sử dụng nn.Xxx cho các lớp có tham số cần học (như Conv2d, Linear, BatchNorm2d). Đối với các hàm không có tham số học (như hàm kích hoạt ReLU, các phép pooling, hay hàm mất mát), việc sử dụng nn.functional hoặc nn.Xxx đều được, nhưng nn.functional thường gọn gàng hơn.

Tuy nhiên, đối với Dropout, luôn luôn ưu tiên sử dụng nn.Dropout để tránh các lỗi ngầm định khi chuyển đổi giữa chế độ training và evaluation.

Việc lựa chọn cuối cùng phụ thuộc vào độ phức tạp của kiến trúc. Nếu nn.Module không đáp ứng được tính linh hoạt cần thiết (ví dụ: các phép toán động, kiến trúc neural đặc thù), nn.functional là công cụ mạnh mẽ để can thiệp sâu vào luồng tính toán.

Thẻ: PyTorch nn-module functional-api model-evaluation tensor-operations

Đăng vào ngày 2 tháng 10 lúc 00:30