Hiệu suất và Ổn định khi Làm Việc với Bộ Dữ Liệu QM9 trong PyTorch Geometric

Bộ dữ liệu QM9 là một tiêu chuẩn vàng trong học máy phân tử, chứa thông tin cấu trúc và 19 thuộc tính lượng tử của hơn 130.000 phân tử nhỏ. Tuy nhiên, việc tích hợp mượt mà vào các pipeline huấn luyện GNN thường gặp trở ngại do ba yếu tố chính: quản lý đường dẫn không nhất quán, phụ thuộc vào thư viện hóa học bên ngoài, và cách xử lý mục tiêu đa chiều thiếu rõ ràng. Bài viết này trình bày một hướng tiếp cận thực tiễn, tập trung vào tính ổn định, khả năng mở rộng và hiệu suất — thay vì chỉ "làm cho chạy được".

Cấu hình đường dẫn an toàn — không phụ thuộc vào môi trường thực thi

Việc sử dụng __file__ hoặc os.getcwd() để xác định thư mục lưu trữ dữ liệu dễ gây lỗi khi chạy trên Jupyter, Colab hoặc hệ thống phân tán. Thay vào đó, nên tách biệt hoàn toàn vị trí lưu trữ khỏi logic tải dữ liệu:

from torch_geometric.datasets import QM9
import tempfile

# Sử dụng thư mục tạm hoặc thư mục được kiểm soát rõ ràng
with tempfile.TemporaryDirectory() as tmp_dir:
    dataset = QM9(root=tmp_dir)
    print(f"Đã tải thành công {len(dataset)} mẫu vào {tmp_dir}")

Xử lý RDKit — chủ động kiểm soát luồng tiền xử lý

PyTorch Geometric tự động chuyển sang phiên bản đã xử lý sẵn nếu RDKit không khả dụng — nhưng điều này làm mất đi khả năng trích xuất đặc trưng từ cấu trúc phân tử (ví dụ: SMILES, topological fingerprints). Giải pháp tối ưu là kiểm tra và thiết lập rõ ràng luồng xử lý:

import subprocess
import sys

def ensure_rdkit():
    try:
        import rdkit
        return True
    except ImportError:
        subprocess.check_call([sys.executable, "-m", "pip", "install", "rdkit-pypi"])
        return True

ensure_rdkit()

# Bắt buộc sử dụng luồng xử lý gốc (không dùng preprocessed_url)
from torch_geometric.datasets import QM9
dataset = QM9(
    root="./qm9_full",
    force_reload=True,  # Đảm bảo tái xử lý nếu cần
)

Chuẩn hóa mục tiêu — tách biệt logic tiền xử lý và mô hình

Thay vì sửa trực tiếp dataset.data.y, hãy xây dựng lớp biến đổi có thể tái sử dụng và kiểm soát được:

import torch
from torch_geometric.transforms import BaseTransform

class StandardizeTarget(BaseTransform):
    def __init__(self, target_idx: int, eps: float = 1e-8):
        self.target_idx = target_idx
        self.eps = eps
        self.mean = None
        self.std = None

    def __call__(self, data):
        if self.mean is None:
            y_full = data.y[:, self.target_idx]
            self.mean = y_full.mean()
            self.std = y_full.std().clamp(min=self.eps)
        data.y = (data.y[:, self.target_idx] - self.mean) / self.std
        return data

# Áp dụng trong pipeline
transform = StandardizeTarget(target_idx=7)  # Năng lượng nội tại ở 0K
dataset = QM9(root="./qm9", transform=transform)

Tối ưu hóa hiệu suất — từ bộ nhớ đến GPU

QM9 có thể chiếm hơn 2 GB RAM khi tải đầy đủ. Để giảm áp lực bộ nhớ và tăng tốc độ huấn luyện:

  • Sử dụng pre_filter để loại bỏ mẫu không cần thiết ngay từ đầu:
def filter_small_molecules(data):
    return data.num_nodes >= 5 and data.num_nodes <= 29

dataset = QM9(
    root="./qm9",
    pre_filter=filter_small_molecules,
)
  • Kết hợp PinMemoryLoadernum_workers > 0 để tăng băng thông I/O:
from torch_geometric.loader import DataLoader
from torch.utils.data import default_collate

class PinMemoryLoader(DataLoader):
    def __init__(self, *args, **kwargs):
        super().__init__(*args, **kwargs)
        self.pin_memory = True

loader = PinMemoryLoader(
    dataset,
    batch_size=64,
    num_workers=6,
    shuffle=True,
    collate_fn=default_collate,
)

Mô hình mẫu — kiến trúc nhẹ, dễ kiểm soát

Dưới đây là một mạng GNN đơn giản nhưng đầy đủ chức năng, được thiết kế để minh họa rõ luồng dữ liệu và tránh các anti-pattern phổ biến (như GRU lồng sâu hay Set2Set không cần thiết):

import torch
import torch.nn as nn
from torch_geometric.nn import GCNConv, global_mean_pool

class QM9Predictor(nn.Module):
    def __init__(self, in_channels: int, hidden_dim: int = 128, out_channels: int = 1):
        super().__init__()
        self.embed = nn.Linear(in_channels, hidden_dim)
        self.conv1 = GCNConv(hidden_dim, hidden_dim)
        self.conv2 = GCNConv(hidden_dim, hidden_dim)
        self.head = nn.Sequential(
            nn.Linear(hidden_dim, hidden_dim // 2),
            nn.ReLU(),
            nn.Linear(hidden_dim // 2, out_channels),
        )

    def forward(self, x, edge_index, batch):
        x = self.embed(x)
        x = self.conv1(x, edge_index).relu()
        x = self.conv2(x, edge_index)
        x = global_mean_pool(x, batch)
        return self.head(x).squeeze(-1)

# Khởi tạo và huấn luyện
model = QM9Predictor(dataset.num_features)
optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4, weight_decay=1e-2)

Chẩn đoán và khắc phục sự cố thường gặp

Một số lỗi phổ biến và cách xử lý:

  • Lỗi "OSError: [Errno 24] Too many open files": Thiết lập giới hạn file descriptor cao hơn trước khi khởi động Python (trên Linux/macOS):
    ulimit -n 4096
  • Dữ liệu bị "nan" sau chuẩn hóa: Kiểm tra xem std có bằng 0 không trước khi chia — luôn dùng clamp(min=eps).
  • GPU out-of-memory khi batch_size=32: Giảm kích thước batch và tăng gradient_accumulation_steps:
accum_steps = 4
for i, batch in enumerate(train_loader):
    loss = compute_loss(model, batch)
    loss.backward()
    if (i + 1) % accum_steps == 0:
        optimizer.step()
        optimizer.zero_grad()

Thẻ: pytorch-geometric qm9 graph-neural-networks rdkit molecular-machine-learning

Đăng vào ngày 5 tháng 9 lúc 16:07