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