CoolFace
Modelpublic

GAD-Research-Lab/MedicalAI-Light-Weight

sourceHugging Faceapache-2.0updated 2mo agoView on Hugging Face
0likes14downloads
training.py350 linesDownload Raw Back to root
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