GAD-Research-Lab/MedicalAI-Light-Weight
014
1import argparse2import csv3import os4import random5 6import torch7import torch.nn as nn8from PIL import Image9from torch.utils.data import Dataset, DataLoader, random_split10 11DATA_DIR = "./data"12CSV_PATH = os.path.join(DATA_DIR, "dataset.csv")13IMAGES_DIR = os.path.join(DATA_DIR, "images")14CHECKPOINT_DIR = "./checkpoints"15CHECKPOINT_PATH = os.path.join(CHECKPOINT_DIR, "fusion_model.pth")16CONFIDENCE_THRESHOLD = 0.7517 18CSV_COLUMNS = ["image_path", "source", "symptoms", "diagnosis", "labels"]19 20def load_label_list():21 if not os.path.exists(CSV_PATH):22 return []23 labels = set()24 with open(CSV_PATH, newline="", encoding="utf-8") as f:25 for row in csv.DictReader(f):26 d = row.get("diagnosis", "").strip().lower()27 if d:28 labels.add(d)29 return sorted(labels)30 31def prepare_data():32 from datasets import load_dataset33 os.makedirs(IMAGES_DIR, exist_ok=True)34 35 file_exists = os.path.exists(CSV_PATH)36 if not file_exists:37 with open(CSV_PATH, "w", newline="", encoding="utf-8") as f:38 writer = csv.writer(f)39 writer.writerow(CSV_COLUMNS)40 41 rows_written = 042 sources_done = []43 44 # ── IU-Xray (image + question + report) ──45 print("Downloading IU-Xray from Hugging Face...")46 iuxray = load_dataset("ayyuce/Indiana_University_Chest_X-ray_Collection", split="train")47 written = 048 for i, example in enumerate(iuxray):49 symptoms = (example.get("question") or "").strip()50 diagnosis = (example.get("report") or "").strip()51 image = example.get("image")52 if not symptoms or not diagnosis or image is None:53 continue54 image_path = os.path.join(IMAGES_DIR, f"iu_xray_{i}.jpg")55 image.convert("RGB").save(image_path)56 with open(CSV_PATH, "a", newline="", encoding="utf-8") as f:57 writer = csv.writer(f)58 writer.writerow([image_path, "iu_xray", symptoms, diagnosis, ""])59 written += 160 print(f" IU-Xray: {written} rows")61 rows_written += written62 sources_done.append(f"iu_xray ({written})")63 64 # ── NIH Chest X-ray (image + disease labels) ──65 print("Downloading NIH Chest X-ray from Hugging Face...")66 nih = load_dataset("g-ronimo/NIH-Chest-X-ray-dataset_resized300px", split="train", streaming=True)67 label_names = [68 "No Finding", "Atelectasis", "Cardiomegaly", "Effusion", "Infiltration",69 "Mass", "Nodule", "Pneumonia", "Pneumothorax", "Consolidation",70 "Edema", "Emphysema", "Fibrosis", "Pleural_Thickening", "Hernia"71 ]72 written = 073 for i, example in enumerate(nih):74 if written >= 3000:75 break76 image = example.get("image")77 label_indices = example.get("labels", [])78 if image is None or not label_indices:79 continue80 label_str = "|".join(label_names[idx] for idx in label_indices)81 primary_diagnosis = label_names[label_indices[0]]82 image_path = os.path.join(IMAGES_DIR, f"nih_{i}.jpg")83 image.convert("RGB").save(image_path)84 with open(CSV_PATH, "a", newline="", encoding="utf-8") as f:85 writer = csv.writer(f)86 writer.writerow([image_path, "nih", "", primary_diagnosis, label_str])87 written += 188 if written % 500 == 0:89 print(f" NIH progress: {written}...")90 print(f" NIH: {written} rows")91 rows_written += written92 sources_done.append(f"nih ({written})")93 94 print(f"Done. Total: {rows_written} rows written to {CSV_PATH}")95 print(f"Sources: {', '.join(sources_done)}")96 97def add_data(image_path, symptoms, diagnosis, labels=""):98 os.makedirs(DATA_DIR, exist_ok=True)99 file_exists = os.path.exists(CSV_PATH)100 with open(CSV_PATH, "a", newline="", encoding="utf-8") as f:101 writer = csv.writer(f)102 if not file_exists:103 writer.writerow(CSV_COLUMNS)104 writer.writerow([image_path, "user", symptoms, diagnosis, labels])105 print(f"Added 1 row to {CSV_PATH}: diagnosis='{diagnosis}'")106 107class FusionDataset(Dataset):108 def __init__(self, csv_path, label_list):109 self.rows = []110 with open(csv_path, newline="", encoding="utf-8") as f:111 for row in csv.DictReader(f):112 d = row.get("diagnosis", "").strip().lower()113 if d:114 self.rows.append(row)115 self.label_list = label_list116 117 def __len__(self):118 return len(self.rows)119 120 def __getitem__(self, idx):121 row = self.rows[idx]122 image = Image.open(row["image_path"]).convert("RGB")123 symptoms = row.get("symptoms", "").strip()124 label_idx = self.label_list.index(row["diagnosis"].strip().lower())125 return image, symptoms, label_idx126 127def collate_fn(batch):128 images = [item[0] for item in batch]129 symptoms = [item[1] for item in batch]130 labels = torch.tensor([item[2] for item in batch], dtype=torch.long)131 return images, symptoms, labels132 133class DiagnosisFusionModel(nn.Module):134 def __init__(self, num_conditions):135 super().__init__()136 from transformers import CLIPModel, CLIPProcessor, AutoTokenizer, AutoModel137 self.image_processor = CLIPProcessor.from_pretrained("openai/clip-vit-base-patch32")138 self.image_encoder = CLIPModel.from_pretrained("openai/clip-vit-base-patch32")139 self.symptom_tokenizer = AutoTokenizer.from_pretrained("emilyalsentzer/Bio_ClinicalBERT")140 self.symptom_encoder = AutoModel.from_pretrained("emilyalsentzer/Bio_ClinicalBERT")141 for param in self.image_encoder.parameters():142 param.requires_grad = False143 for param in self.symptom_encoder.parameters():144 param.requires_grad = False145 self.classifier = nn.Sequential(146 nn.Linear(512 + 768, 256),147 nn.ReLU(),148 nn.Dropout(0.2),149 nn.Linear(256, num_conditions),150 )151 152 def encode_images(self, images):153 inputs = self.image_processor(images=images, return_tensors="pt")154 with torch.no_grad():155 return self.image_encoder.get_image_features(**inputs)156 157 def encode_symptoms(self, symptom_texts):158 inputs = self.symptom_tokenizer(159 symptom_texts, return_tensors="pt", padding=True, truncation=True, max_length=64160 )161 with torch.no_grad():162 outputs = self.symptom_encoder(**inputs)163 return outputs.last_hidden_state.mean(dim=1)164 165 def forward(self, images, symptom_texts):166 image_vecs = self.encode_images(images)167 symptom_vecs = self.encode_symptoms(symptom_texts)168 combined = torch.cat([image_vecs, symptom_vecs], dim=-1)169 return self.classifier(combined)170 171def train(epochs, batch_size, lr, val_split, use_amp, grad_accum):172 from rich.console import Console173 from rich.table import Table174 from rich.progress import Progress, BarColumn, TextColumn, TimeElapsedColumn175 _console = Console()176 177 has_gpu = torch.cuda.is_available()178 use_amp = use_amp and has_gpu179 scaler = torch.cuda.amp.GradScaler() if use_amp else None180 181 label_list = load_label_list()182 if not label_list:183 _console.print("[red]No data found. Run --mode prepare-data or --mode add-data first.[/red]")184 return185 186 dataset = FusionDataset(CSV_PATH, label_list)187 val_size = max(int(val_split * len(dataset)), 1)188 train_size = len(dataset) - val_size189 train_subset, val_subset = random_split(dataset, [train_size, val_size])190 191 train_loader = DataLoader(train_subset, batch_size=batch_size, shuffle=True, collate_fn=collate_fn)192 val_loader = DataLoader(val_subset, batch_size=batch_size, shuffle=False, collate_fn=collate_fn)193 194 _console.print(f"[bold cyan]Training Setup[/bold cyan]")195 _console.print(f" Classes: {len(label_list)}")196 _console.print(f" Train/Val: {len(train_subset)}/{len(val_subset)}")197 _console.print(f" Batch size: {batch_size} Grad accum: {grad_accum}")198 _console.print(f" Device: {'GPU' if has_gpu else 'CPU'} AMP: {'ON' if use_amp else 'OFF'}")199 200 model = DiagnosisFusionModel(num_conditions=len(label_list))201 if has_gpu:202 model = model.cuda()203 optimizer = torch.optim.AdamW(model.classifier.parameters(), lr=lr)204 loss_fn = nn.CrossEntropyLoss()205 206 for epoch in range(epochs):207 _console.print(f"\n[bold yellow]Epoch {epoch + 1}/{epochs}[/bold yellow]")208 _console.print("-" * 40)209 210 # ── Train ──211 model.train()212 train_loss = 0.0213 optimizer.zero_grad()214 train_progress = Progress(215 TextColumn("[cyan] Train[/cyan]"),216 BarColumn(),217 TextColumn("{task.completed}/{task.total}"),218 TextColumn("[green]{task.fields[loss]:.4f}[/green]"),219 TimeElapsedColumn(),220 transient=True,221 )222 with train_progress:223 task = train_progress.add_task("", total=len(train_loader), loss=0.0)224 for i, (images, symptoms, labels) in enumerate(train_loader):225 if has_gpu:226 labels = labels.cuda()227 with torch.amp.autocast("cuda", enabled=use_amp):228 logits = model(images, symptoms)229 loss = loss_fn(logits, labels)230 loss = loss / grad_accum231 if use_amp:232 scaler.scale(loss).backward()233 else:234 loss.backward()235 236 if (i + 1) % grad_accum == 0 or (i + 1) == len(train_loader):237 if use_amp:238 scaler.step(optimizer)239 scaler.update()240 else:241 optimizer.step()242 optimizer.zero_grad()243 244 train_loss += loss.item() * grad_accum245 train_progress.update(task, advance=1, loss=loss.item() * grad_accum)246 247 avg_train_loss = train_loss / len(train_loader)248 249 # ── Validation ──250 model.eval()251 val_loss = 0.0252 with torch.no_grad():253 for images, symptoms, labels in val_loader:254 if has_gpu:255 labels = labels.cuda()256 logits = model(images, symptoms)257 loss = loss_fn(logits, labels)258 val_loss += loss.item()259 260 avg_val_loss = val_loss / len(val_loader)261 262 table = Table(show_header=False, box=None)263 table.add_column("Metric", style="cyan")264 table.add_column("Value", style="green")265 table.add_row("Train loss", f"{avg_train_loss:.4f}")266 table.add_row("Val loss", f"{avg_val_loss:.4f}")267 _console.print(table)268 269 os.makedirs(CHECKPOINT_DIR, exist_ok=True)270 torch.save({"model_state": model.classifier.state_dict(), "label_list": label_list}, CHECKPOINT_PATH)271 _console.print(f"[green]Saved checkpoint to {CHECKPOINT_PATH}[/green]")272 273def test(batch_size):274 if not os.path.exists(CHECKPOINT_PATH):275 print("No checkpoint found. Run --mode train first.")276 return277 checkpoint = torch.load(CHECKPOINT_PATH, weights_only=False)278 label_list = checkpoint["label_list"]279 model = DiagnosisFusionModel(num_conditions=len(label_list))280 model.classifier.load_state_dict(checkpoint["model_state"])281 model.eval()282 dataset = FusionDataset(CSV_PATH, label_list)283 test_size = max(int(0.2 * len(dataset)), 1)284 _, test_subset = random_split(dataset, [len(dataset) - test_size, test_size])285 loader = DataLoader(test_subset, batch_size=batch_size, shuffle=False, collate_fn=collate_fn)286 correct = 0287 inconclusive = 0288 total = 0289 with torch.no_grad():290 for images, symptoms, labels in loader:291 logits = model(images, symptoms)292 probs = torch.softmax(logits, dim=-1)293 confidence, predicted = torch.max(probs, dim=-1)294 for i in range(len(labels)):295 total += 1296 if confidence[i].item() < CONFIDENCE_THRESHOLD:297 inconclusive += 1298 elif predicted[i].item() == labels[i].item():299 correct += 1300 print(f"Tested on {total} held-out examples")301 print(f"Correct (above confidence threshold): {correct} ({100 * correct / total:.1f}%)")302 print(f"Flagged as inconclusive / needs follow-up: {inconclusive} ({100 * inconclusive / total:.1f}%)")303 304def info():305 if not os.path.exists(CSV_PATH):306 print("No dataset.csv found. Run --mode prepare-data first.")307 return308 sources = {}309 total = 0310 with open(CSV_PATH, newline="", encoding="utf-8") as f:311 for row in csv.DictReader(f):312 src = row.get("source", "unknown")313 sources[src] = sources.get(src, 0) + 1314 total += 1315 print(f"Dataset: {CSV_PATH}")316 print(f"Total rows: {total}")317 for src, count in sorted(sources.items()):318 print(f" {src}: {count}")319 print(f"Images dir: {IMAGES_DIR}")320 img_count = len([x for x in os.listdir(IMAGES_DIR) if os.path.isfile(os.path.join(IMAGES_DIR, x))]) if os.path.exists(IMAGES_DIR) else 0321 print(f"Images: {img_count}")322 323if __name__ == "__main__":324 parser = argparse.ArgumentParser(description="Train/test the medical image+symptom fusion model")325 parser.add_argument("--mode", required=True, choices=["prepare-data", "add-data", "train", "test", "info"])326 parser.add_argument("--image", help="Path to an image file (for --mode add-data)")327 parser.add_argument("--symptoms", help="Symptom description text (for --mode add-data)")328 parser.add_argument("--diagnosis", help="Diagnosis label (for --mode add-data)")329 parser.add_argument("--epochs", type=int, default=5)330 parser.add_argument("--batch_size", type=int, default=8)331 parser.add_argument("--lr", type=float, default=1e-3)332 parser.add_argument("--val_split", type=float, default=0.15, help="Fraction of data for validation")333 parser.add_argument("--use_amp", action="store_true", help="Enable mixed precision (GPU only)")334 parser.add_argument("--grad_accum", type=int, default=1, help="Gradient accumulation steps")335 args = parser.parse_args()336 337 if args.mode == "prepare-data":338 prepare_data()339 elif args.mode == "add-data":340 if not (args.image and args.symptoms and args.diagnosis):341 print("--mode add-data requires --image, --symptoms, and --diagnosis")342 else:343 add_data(args.image, args.symptoms, args.diagnosis)344 elif args.mode == "train":345 train(args.epochs, args.batch_size, args.lr, args.val_split, args.use_amp, args.grad_accum)346 elif args.mode == "test":347 test(args.batch_size)348 elif args.mode == "info":349 info()350 