Machine Learning Academy · Lektion

Domänanpassning: Medicinsk bildanalys med få etiketter

Ni kommer att tillämpa transfer learning från ImageNet på en datamängd med lungröntgenbilder, implementera klassviktad förlust för obalanserade patologier och utvärdera AUC-ROC.

Lektion 4 av 412 steg

Domänanpassning: Medicinsk bildanalys med få etiketter är en gratis lektion i Machine Learning Academy på CoddyKit. Detta är lektion 4 av 4. Ni kan läsa hela lektionen gratis nedan och sedan öva praktiskt i webbläsaren med en inbyggd kodredigerare och en AI-handledare som är tillgänglig dygnet runt. Den ingår i lärvägen för Machine Learning Academy, och Era framsteg synkroniseras mellan webben och CoddyKit-appen. Kursen i Machine Learning Academy innehåller totalt 4 lektioner.

Utmaningen inom medicinsk bildanalys

Medicinsk bildanalys innebär unika utmaningar för transfer learning. Till skillnad från naturliga fotografier liknar lungröntgenbilder, MR-bilder och histologiska preparat inte alls bilderna i ImageNet: de är gråskalebilder eller har domänspecifika färgmönster, de viktiga egenskaperna (lesioner, noduli och förkalkningar) är subtila och domänspecifika, och märkning av data kräver expertradiologer – vilket gör stora märkta datamängder dyra och sällsynta.

Trots dessa utmaningar överträffar ImageNet-förtränade modeller konsekvent modeller som tränats från grunden i uppgifter inom medicinsk bildanalys, även när det visuella utseendet skiljer sig avsevärt. Universella lågnivåegenskaper (kantdetektorer och texturfilter) kan överföras mellan domäner och ger en stark initialisering som påskyndar konvergensen och förbättrar generaliseringen när antalet etiketter är begränsat.

CheXpert-datamängden: klassificering av röntgenbilder med flera etiketter

CheXpert är en benchmark-datamängd med lungröntgenbilder som innehåller 224 316 bilder och 14 etiketter (Cardiomegaly, Pleural Effusion, Pneumonia, Atelectasis med flera). I vårt scenario med få etiketter simulerar vi användning av endast en liten andel – säg 1 % (cirka 2 243 bilder) – för att efterlikna kliniska miljöer där budgeten för annotering är begränsad.

Detta är ett problem med klassificering med flera etiketter: varje bild kan samtidigt ha flera patologier, till skillnad från klassificering med en enda etikett. Målet är en vektor med 14 binära värden, och vi använder Binary Cross-Entropy with Logits (BCE) elementvis. Utvärderingen använder AUC-ROC per patologi, med medelvärde över alla 14 etiketter.

# Dataset setup (pseudo-code for illustration)
import torch
from torch.utils.data import Dataset
from PIL import Image
import pandas as pd

class CheXpertDataset(Dataset):
    def __init__(self, csv_path, img_dir, transform=None):
        self.df = pd.read_csv(csv_path)
        self.img_dir = img_dir
        self.transform = transform
        self.labels = ['Atelectasis', 'Cardiomegaly',
                       'Consolidation', 'Edema', 'Pleural Effusion']

    def __len__(self):
        return len(self.df)

    def __getitem__(self, idx):
        img_path = self.df.iloc[idx]['Path']
        img = Image.open(img_path).convert('RGB')  # Convert grayscale to 3-ch
        label = torch.tensor(self.df.iloc[idx][self.labels].values.astype(float))
        if self.transform:
            img = self.transform(img)
        return img, label

Konvertera gråskala till RGB för förtränade modeller

De flesta medicinska bilder (röntgenbilder och CT-bilder) är gråskalebilder (1 kanal), men ImageNet-förtränade modeller förväntar sig indata i RGB med 3 kanaler. Den enklaste lösningen är Image.convert('RGB'), som kopierar den enda kanalen tre gånger, eller transforms.Grayscale(num_output_channels=3) i transformationspipeline.

Detta är något ineffektivt – de tre kanalerna är identiska – men i praktiken fungerar det bra eftersom modellen helt enkelt lär sig att vikta alla tre kanaler lika. Ett alternativ är att ersätta det första konvolutionella lagret med ett nytt Conv2d(1, 64, kernel_size=7, ...) och initiera det genom att beräkna medelvärdet av vikterna för de tre indatakanalerna. Detta är mer principiellt korrekt men ökar träningens komplexitet.

from torchvision import transforms

# Method 1: Replicate channel at loading time (simplest)
# img = Image.open(path).convert('RGB')

# Method 2: Use transforms
medical_transform = transforms.Compose([
    transforms.Resize((224, 224)),
    transforms.Grayscale(num_output_channels=3),  # 1 -> 3 channels
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.485, 0.456, 0.406],
                         std=[0.229, 0.224, 0.225])
])

# Method 3: Modify first conv layer for single-channel input
import torchvision.models as models
model = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V1)
# Average the 3-channel weights into 1 channel
w = model.conv1.weight.mean(dim=1, keepdim=True)
model.conv1 = torch.nn.Conv2d(1, 64, kernel_size=7, stride=2, padding=3, bias=False)
model.conv1.weight.data = w

Förlust för flera etiketter: BCE with Logits

Vid klassificering med flera etiketter är varje etikett en oberoende binär prediktion. Vi använder nn.BCEWithLogitsLoss() på en vektor med 14 logits. Förlusten beräknas elementvis och medelvärdesbildas över både de 14 klasserna och batchstorleken.

Obalanserade etiketter är en stor utmaning inom medicinsk bildanalys: endast 5–10 % av bilderna visar Pneumonia eller Consolidation, medan 40–60 % visar Pleural Effusion. Skicka pos_weight till BCEWithLogitsLoss för att vikta upp den sällsynta positiva klassen: ett pos_weight på 10 gör att modellen lägger 10× mer vikt vid positiva exempel på den patologin.

import torch
import torch.nn as nn

# Multi-label BCE loss
criterion = nn.BCEWithLogitsLoss()

# With class weighting for imbalanced pathologies
# Compute pos_weight from training data frequencies
pos_counts = train_labels.sum(dim=0)  # Positives per class
neg_counts = len(train_labels) - pos_counts
pos_weight = (neg_counts / pos_counts.clamp(min=1)).clamp(max=20)  # Cap at 20

criterion = nn.BCEWithLogitsLoss(pos_weight=pos_weight)

print('pos_weight per label:', pos_weight)
# High values mean that class is rare and needs more emphasis

Modellarkitektur för utdata med flera etiketter

Ersätt klassificeringshuvudet i ResNet-50 med ett linjärt lager som ger 14 logits (en per patologi), i stället för standardvärdet 1000 för ImageNet. Vi använder inte sigmoid i forward-passet – BCEWithLogitsLoss tillämpar det internt för numerisk stabilitet. Vid inferens tillämpar ni sigmoid manuellt för att få fram sannolikheter.

Att lägga till ett dropout-lager före det slutliga linjära lagret är särskilt viktigt när datamängden är liten, eftersom regularisering hindrar det lilla huvudet från att överanpassa. En dropout-sannolikhet på 0,3–0,5 är vanlig vid finjustering för medicinsk bildanalys.

import torchvision.models as models
import torch.nn as nn

NUM_CLASSES = 14  # One per CheXpert pathology label

model = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V1)

# Replace with multi-label head
model.fc = nn.Sequential(
    nn.Dropout(p=0.4),
    nn.Linear(model.fc.in_features, NUM_CLASSES)
)
# No sigmoid here - BCEWithLogitsLoss handles it internally

print('Model output shape for batch of 16:',
      model(torch.randn(16, 3, 224, 224)).shape)  # (16, 14)

AUC-ROC: rätt mått för medicinska uppgifter

Noggrannhet är oanvändbart för obalanserade medicinska datamängder. Om endast 5 % av patienterna har Pneumonia uppnår en modell som alltid förutsäger ”ingen pneumoni” 95 % noggrannhet, samtidigt som den är helt oanvändbar kliniskt. AUC-ROC (Area Under the ROC Curve) mäter förmågan att skilja mellan klasser över alla tröskelvärden.

Vi beräknar AUC-ROC separat för var och en av de 14 patologierna och rapporterar genomsnittlig AUC över alla etiketter. AUC på 0,5 motsvarar slumpen, 0,7 är acceptabelt, 0,85 eller högre motsvarar klinisk nivå och 0,9 eller högre närmar sig ofta radiologprestanda. roc_auc_score i scikit-learn beräknar detta effektivt.

import numpy as np
from sklearn.metrics import roc_auc_score
import torch

def evaluate_auc(model, loader, device):
    model.eval()
    all_logits, all_labels = [], []
    with torch.no_grad():
        for images, labels in loader:
            logits = model(images.to(device))
            all_logits.append(torch.sigmoid(logits).cpu().numpy())
            all_labels.append(labels.numpy())
    probs = np.vstack(all_logits)    # (N, 14)
    targets = np.vstack(all_labels)  # (N, 14)
    # AUC per class, then average
    aucs = [roc_auc_score(targets[:, i], probs[:, i])
            for i in range(targets.shape[1])]
    return np.mean(aucs), aucs

Träning med få etiketter: viktiga tekniker

När ni endast har hundratals eller några tusen märkta medicinska bilder finns det flera tekniker som hjälper er att utnyttja de tillgängliga data maximalt. Kraftig dataaugmentering är viktigast: slumpmässiga vändningar, rotationer samt variationer i kontrast och ljusstyrka hjälper, samtidigt som augmenteringarna måste vara kliniskt rimliga (en lungröntgenbild bör inte vändas vertikalt – det skulle aldrig förekomma i klinisk praxis).

Progressiv skalning (att först träna med lägre upplösning och sedan öka den) är en annan effektiv teknik. Börja med 128×128 för att snabbt kunna iterera och finjustera sedan med 224×224 eller till och med 320×320 för slutlig noggrannhet. Detta går mycket snabbare än att alltid träna med full upplösning och ger ofta samma noggrannhet.

from torchvision import transforms

# Clinically appropriate augmentations for chest X-rays
medical_aug = transforms.Compose([
    transforms.RandomResizedCrop(224, scale=(0.85, 1.0)),  # Mild crop
    transforms.RandomHorizontalFlip(p=0.5),                 # OK: X-rays can be flipped
    # transforms.RandomVerticalFlip(p=0.5),                 # NOT OK: clinically invalid
    transforms.ColorJitter(brightness=0.2, contrast=0.3),  # Simulate scan variation
    transforms.RandomRotation(degrees=10),                  # Slight mis-alignment
    transforms.ToTensor(),
    transforms.Normalize([0.485, 0.456, 0.406],
                         [0.229, 0.224, 0.225])
])

Självövervakad förträning för medicinsk bildanalys

När ImageNet-egenskaper inte överförs tillräckligt väl är självövervakad förträning på omärkta medicinska bilder ett kraftfullt alternativ. Metoder som SimCLR, MoCo och DINO lär sig representationer genom att träna modellen att känna igen att två augmenterade versioner av samma bild liknar varandra, utan att några etiketter behövs.

Arbetsflödet är följande: (1) förträna på stora omärkta medicinska datamängder (CheXpert har 224 000 omärkta bilder), (2) finjustera med den lilla mängden märkta data. Detta överträffar konsekvent ImageNet-transfer i medicinska uppgifter, eftersom modellen lär sig egenskaper som är specifika för den medicinska domänen i stället för egenskaper för allmän objektigenkänning.

# Self-supervised pre-training concept (SimCLR-style)
# No labels needed during this phase

class SimCLRLoss(torch.nn.Module):
    def __init__(self, temperature=0.07):
        super().__init__()
        self.temperature = temperature

    def forward(self, z_i, z_j):
        # z_i, z_j: augmented views of the same images
        # Maximise agreement between paired views
        # Minimise agreement between all other pairs in batch
        z = torch.cat([z_i, z_j], dim=0)
        z = torch.nn.functional.normalize(z, dim=1)
        sim = torch.matmul(z, z.T) / self.temperature
        # Contrastive loss computation...
        return sim  # Simplified illustration

Osäkerhetskvantifiering i medicinsk AI

Inom medicinskt beslutsstöd är det lika viktigt att veta hur säker en modell är som att känna till själva prediktionen. En modell som säger ”Pneumonia: 95 % sannolikhet” bör vara mer tillförlitlig än en som säger ”52 % sannolikhet”. Standardiserade softmax-sannolikheter är ofta överdrivet säkra och representerar inte den verkliga osäkerheten.

Monte Carlo Dropout (MC Dropout) approximerar bayesiansk osäkerhet genom att behålla dropout aktiverat vid inferens och köra forward-passet flera gånger. Variansen i prediktionerna mellan körningarna skattar osäkerheten. Prediktioner med hög varians bör flaggas för mänsklig granskning i stället för att hanteras automatiskt.

import torch

def mc_dropout_predict(model, x, n_samples=30):
    model.train()  # Keep dropout active during inference
    predictions = []
    with torch.no_grad():
        for _ in range(n_samples):
            logits = model(x)
            probs = torch.sigmoid(logits)
            predictions.append(probs)
    preds = torch.stack(predictions)  # (n_samples, batch, num_classes)
    mean_pred = preds.mean(dim=0)     # Average prediction
    uncertainty = preds.std(dim=0)    # Std = uncertainty estimate
    return mean_pred, uncertainty

# High uncertainty cases should be reviewed by a radiologist
MEAN_THRESHOLD = 0.5
UNCERTAINTY_THRESHOLD = 0.15

Etiska överväganden inom medicinsk ML

Att införa ML-modeller i medicinska miljöer innebär ett stort etiskt ansvar. Partiskhet i datamängden är en kritisk fråga: modeller som huvudsakligen tränats på bilder från en viss sjukhuskedjas skannrar, patientgrupper eller skannertillverkare kan misslyckas på bilder från andra miljöer – ett välkänt problem inom radiologisk AI.

Före driftsättning ska ni alltid utvärdera modellens prestanda för olika demografiska undergrupper (ålder, kön och etnicitet) samt med olika typer av skanningsutrustning. Regelverk som FDA:s riktlinjer för AI/ML-baserad Software as a Medical Device (SaMD) i USA och EU:s förordning om medicintekniska produkter (MDR) kräver noggrann klinisk validering före driftsättning. Modeller bör stödja radiologer, inte ersätta dem – särskilt vid allvarliga patologier där fel kan få livsavgörande konsekvenser.

Snabb kontroll

Testa vad ni har förstått av domänanpassning för medicinsk bildanalys i den här lektionen.

Sammanfattning av lektionen

I den här lektionen har ni lärt er följande: transfer learning för medicinsk bildanalys tillämpar ImageNet-förtränade modeller på domäner med gråskalebilder och få etiketter genom att konvertera kanaler och ersätta klassificeringshuvudet, BCE-förlust med flera etiketter och pos_weight hanterar den extrema klassobalans som är typisk för datamängder med patologier, och AUC-ROC är det korrekta utvärderingsmåttet eftersom det mäter förmågan att skilja mellan klasser oberoende av klassbalansen. När ni nu behärskar dessa färdigheter inom transfer learning är ni redo att utforska NLP med BERT i nästa kurs.

Gratis att börja

Lär dig Python med en AI-lärare – gratis

Skriv och kör riktig kod i webbläsaren, få omedelbar hjälp av en AI-lärare dygnet runt och fortsätt där du slutade – på webben eller i appen.

Kurser
30
Lektioner
120

Vanliga frågor

Är lektionen ”Domänanpassning: Medicinsk bildanalys med få etiketter” gratis?

Ja – hela texten till ”Domänanpassning: Medicinsk bildanalys med få etiketter” kan läsas gratis här på webben. Om Ni vill öva interaktivt med en inbyggd kodredigerare och en AI-handledare som är tillgänglig dygnet runt och låsa upp resten av kursen i Machine Learning Academy, kan Ni uppgradera till CoddyKit PRO. Kursen i Machine Learning Academy innehåller totalt 4 lektioner.

Vad lär jag mig i ”Domänanpassning: Medicinsk bildanalys med få etiketter”?

Ni kommer att tillämpa transfer learning från ImageNet på en datamängd med lungröntgenbilder, implementera klassviktad förlust för obalanserade patologier och utvärdera AUC-ROC. Ni övar på Machine Learning Academy med praktisk kod som körs direkt i webbläsaren, medan en AI-handledare som är tillgänglig dygnet runt svarar på Era frågor under lektionen.

Behöver jag någon erfarenhet för att börja lära mig Machine Learning Academy?

Du behöver inga förkunskaper. Utbildningen i Machine Learning Academy på CoddyKit är upplagd för allt från nybörjare till avancerade elever, så att du kan börja här eller från början och gå fram i din egen takt. Detta är lektion 4 av 4.

Hur lång tid tar lektionen ”Domänanpassning: Medicinsk bildanalys med få etiketter”?

De flesta CoddyKit-lektioner tar cirka 5–10 minuter. Varje lektion är kort och interaktiv, så att du gör stadiga framsteg och kan fortsätta precis där du slutade – på webben eller i appen.

Kan jag skriva och köra kod i den här Machine Learning Academy-lektionen?

Ja. Varje Machine Learning Academy-lektion innehåller en inbyggd kodredigerare, så att du kan skriva och köra riktig kod direkt i webbläsaren och få omedelbar AI-feedback – utan lokal installation.

Alla lektioner i den här kursen

  1. Förtränade modeller i torchvision: ResNet, EfficientNet och ViT
  2. Egenskapsextraktion: Frys backbone
  3. Finjustering: Avfrysning och låga inlärningshastigheter
  4. Domänanpassning: Medicinsk bildanalys med få etiketter
← Tillbaka till Machine Learning Academy