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
PinMemoryLoadervànum_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
stdcó bằng 0 không trước khi chia — luôn dùngclamp(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()