Phân loại ảnh cảnh vật bằng Graph Neural Network
Phân loại sáu class cảnh vật của bộ Intel Image Classification bằng GCN trên graph superpixel, so sánh ba cách dựng graph từ cùng một ảnh.
- PyTorch Geometric
- Python
- SLIC
- Delaunay
Một đồ án nhóm: phân loại sáu class cảnh vật (buildings, forest, glacier, mountain, sea, street) của bộ Intel Image Classification mà không dùng CNN. Mỗi ảnh được chuyển thành một graph rồi đưa qua Graph Convolutional Network. Tôi phụ trách phần xử lý dữ liệu, model, huấn luyện và trực quan hoá.
Vì sao lại là graph#
CNN mặc định ảnh là một lưới pixel đều nhau và quan hệ lân cận luôn cố định. Tôi muốn thử hướng ngược lại: gom pixel thành các vùng đồng nhất bằng SLIC, coi mỗi vùng là một node, và để quan hệ giữa các vùng trở thành edge của graph. Với cách nhìn đó, câu hỏi thật sự của đồ án không phải "GNN có thắng CNN không", mà là cách dựng graph ảnh hưởng đến kết quả ra sao.
Ba cách dựng graph từ cùng một ảnh#
Tôi làm ba biến thể, mỗi biến thể một notebook, dùng chung dataset 14.034 ảnh train và 3.000 ảnh test:
gcn-dt.ipynb: resize ảnh về 32×32, mỗi pixel là một node, feature là RGB, edge nối bằng Delaunay triangulation trên toạ độ pixel.gcn-combine.ipynb: SLIC cắt ảnh thành khoảng 100 superpixel, feature của node là màu trung bình cộng toạ độ tâm (5 chiều), edge vẫn là Delaunay nhưng chạy trên tâm vùng.gcn_slic.ipynb: SLIC với 400 superpixel trên ảnh 100×100, edge dựng theo region adjacency — hai superpixel chạm nhau thì có edge — và feature của node mở rộng lên 9 chiều.
Biến thể thứ hai là bước trung gian gọn nhất, vì Delaunay trên tâm vùng cho ngay danh sách edge mà không phải duyệt lại ảnh:
def create_combined_graph(image, n_segments=100, compactness=10):
image = resize(image, (100, 100), anti_aliasing=True)
segments = slic(image, n_segments=n_segments, compactness=compactness, sigma=1)
regions = regionprops(segments + 1)
# Mỗi superpixel thành một node: toạ độ tâm và màu trung bình của vùng.
centers = np.array([r.centroid for r in regions])
colors = np.array([image[segments == i].mean(axis=0) for i in np.unique(segments)])
# Delaunay trên tâm vùng; mỗi tam giác góp ba cạnh, dùng set để khử trùng lặp.
edges = set()
for simplex in Delaunay(centers).simplices:
for i in range(3):
edges.add(tuple(sorted((simplex[i], simplex[(i + 1) % 3]))))
x = torch.tensor(np.hstack([colors, centers]), dtype=torch.float)
edge_index = torch.tensor(np.array(list(edges)).T, dtype=torch.long)
return Data(x=x, edge_index=edge_index), segmentsFeature của node mới là thứ tạo khác biệt#
Hai biến thể đầu chỉ mô tả node bằng màu và vị trí, và cả hai đều dừng lại quanh mức 60%. Ở biến thể cuối tôi thêm hình học của vùng: eccentricity, aspect ratio, solidity, toạ độ tâm và chu vi, cộng với màu trung bình — tổng cộng 9 chiều. Đây là thay đổi làm điểm số nhảy nhiều nhất, và cũng hợp lý: rừng, biển hay phố khác nhau ở hình dạng các mảng chứ không chỉ ở màu.
Kiến trúc thì tôi cố tình giữ đơn giản để so sánh còn công bằng: ba layer GCNConv
128 → 64 → 32 kèm BatchNorm, gộp bằng global mean pool, rồi hai layer Linear và dropout
0.3.
class GCNModel(torch.nn.Module):
def __init__(self, input_dim, hidden_dims, num_classes, dropout=0.3):
super().__init__()
self.conv1 = GCNConv(input_dim, hidden_dims[0])
self.bn1 = BatchNorm(hidden_dims[0])
self.conv2 = GCNConv(hidden_dims[0], hidden_dims[1])
self.bn2 = BatchNorm(hidden_dims[1])
self.conv3 = GCNConv(hidden_dims[1], hidden_dims[2])
self.bn3 = BatchNorm(hidden_dims[2])
self.fc1 = Linear(hidden_dims[2], hidden_dims[3])
self.fc2 = Linear(hidden_dims[3], num_classes)
self.relu = ReLU()
self.dropout = Dropout(dropout)
def forward(self, data):
x, edge_index, batch = data.x, data.edge_index, data.batch
x = self.relu(self.bn1(self.conv1(x, edge_index)))
x = self.relu(self.bn2(self.conv2(x, edge_index)))
x = self.relu(self.bn3(self.conv3(x, edge_index)))
# Số superpixel mỗi ảnh một khác, mean pool đưa về vector cố định.
x = global_mean_pool(x, batch)
return self.fc2(self.dropout(self.relu(self.fc1(x))))Tách tiền xử lý khỏi huấn luyện#
Dựng region adjacency phải duyệt từng pixel để biết hai superpixel nào chạm nhau. Với
14.034 ảnh thì bước này chậm hơn hẳn việc huấn luyện, và ban đầu tôi chạy lại nó mỗi lần
sửa model. Nên tôi cắt đôi pipeline: chuyển ảnh thành graph một lần, lưu mỗi graph
thành một file .pt trong graphs/train_graphs/, và bỏ qua những file đã tồn tại
khi chạy lại. Sau đó vòng train chỉ còn việc nạp .pt vào DataLoader — đổi model
không còn kéo theo đổi dữ liệu.
Kết quả#
- Biến thể superpixel + region adjacency + 9 feature đạt 78,17% accuracy trên 3.000 ảnh test, dừng sớm ở epoch 87 với patience 10
- Ba cách dựng graph cách biệt rõ rệt: pixel + Delaunay 59,6%, superpixel + Delaunay 62,4%, superpixel + region adjacency 78,2%
buildingsluôn là class yếu nhất: precision 0,85 nhưng recall chỉ 0,47, phần lớn bị nhầm sangstreet— hai cảnh có kết cấu vùng gần giống nhau- Cả ba model được đóng gói thành app Streamlit nhận ảnh upload hoặc URL và trả về nhãn dự đoán