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
18
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