ResNet50 fine-tuned trên NIH ChestX-ray14 — 4 lớp đơn-nhãn

Fine-tune từ trọng số ImageNet1k, phân loại ĐƠN-NHÃN (single-label, softmax) giữa 4 bệnh lý lồng ngực phổ biến nhất trong bộ NIH ChestX-ray14 (chỉ dùng ảnh mang đúng 1 bệnh, loại ảnh "No Finding" và ảnh đa-bệnh). 4 lớp: Infiltration, Atelectasis, Effusion, Nodule.

Huấn luyện trên subset cân bằng 8000 ảnh (không phải toàn bộ 112.120 ảnh gốc), early stopping theo Val Macro-AUROC.

Kết quả

  • Val Macro-AUROC (OVR) tốt nhất: 0.7756
  • Test Macro-AUROC (OVR): 0.7919
  • Test Accuracy: 0.5381

Kiến trúc & vì sao phù hợp cho CAM/Grad-CAM

model.layer1 -> layer2 -> layer3 -> layer4 (các block conv) -> AdaptiveAvgPool2d -> 1 Linear duy nhất (model.fc, 4 lớp output, softmax). Không có lớp fully-connected trung gian nào khác, nên đây là kiến trúc kinh điển tương thích CAM (Zhou et al. dùng chính ResNet/GoogLeNet cho bài báo gốc) và Grad-CAM: hook vào model.layer4[-1] để lấy feature map 2048-kênh cuối cùng trước GAP.

Cách tải model

import torch, json
from torchvision import models
from safetensors.torch import load_file
from huggingface_hub import hf_hub_download

repo_id = "Purino/resnet50-nih-chestxray-4class"
config = json.load(open(hf_hub_download(repo_id, "config.json")))
weights_path = hf_hub_download(repo_id, "model.safetensors")

model = models.resnet50(weights=None)
model.fc = torch.nn.Linear(model.fc.in_features, config["num_labels"])
model.load_state_dict(load_file(weights_path))
model.eval()
# dự đoán: probs = torch.softmax(model(x), dim=1); class_idx = probs.argmax(dim=1)

Ví dụ Grad-CAM

Xem file inference_gradcam_example.py trong repo này — chạy được ngay, tạo heatmap Grad-CAM cho ảnh X-quang bất kỳ trong 1 trong 4 lớp trên. Có thể đối chiếu định tính với BBox_List_2017.csv của bộ dữ liệu gốc (bounding box tổn thương do bác sĩ khoanh trên ~1.000 ảnh) để kiểm tra vùng CAM có trùng vùng tổn thương thật hay không.

Giới hạn quan trọng (đọc trước khi dùng)

  • Chỉ huấn luyện trên 8000/112.120 ảnh, và CHỈ trên ảnh đơn-bệnh của 4 lớp trên — model KHÔNG xử lý được ảnh có nhiều bệnh đồng thời hoặc các bệnh lý khác trong 14 nhãn gốc.
  • Nhãn có nhiễu: nhãn gốc được trích xuất tự động từ báo cáo X-quang bằng NLP (NegBio/DNorm), độ chính xác ước tính ~90%, không phải do bác sĩ dán nhãn thủ công từng ảnh.
  • Không dùng cho chẩn đoán lâm sàng. Đây là model nghiên cứu/học thuật.
  • Chia tập train/val/test theo Patient ID (không rò rỉ dữ liệu), xem train_split.csv / val_split.csv / test_split.csv để biết chính xác ảnh nào thuộc tập nào.

Gợi ý mở rộng

  • Huấn luyện trên toàn bộ ảnh đơn-bệnh (bỏ giới hạn MAX_PER_CLASS).
  • Thêm các lớp bệnh khác trong 14 nhãn gốc (tăng NUM_SELECTED_CLASSES).
  • So sánh trực tiếp với bản MobileNetV2 (nhẹ hơn, nhanh hơn) để đánh đổi tốc độ vs độ chính xác.
Downloads last month
-
Safetensors
Model size
23.6M params
Tensor type
F32
·
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support