Kom igångKom igång gratis

Analysera mätvärden per klass

Aggregerade mätvärden ger en bra övergripande bild av modellens prestanda, men det är ofta mer informativt att titta på mätvärdena per klass. Det kan avslöja klasser där modellen presterar sämre.

I den här övningen kör du utvärderingsloopen igen för att beräkna precisionen för vår molnklassificerare – den här gången per klass. Sedan kopplar du dessa värden till klassnamnen för att kunna tolka resultaten. Som vanligt har Precision redan importerats åt dig. Lycka till!

Den här övningen är en del av kursen

Fördjupad djupinlärning med PyTorch

Visa kurs

Övningsinstruktioner

  • Definiera ett precisionsmått som passar för resultat per klass.
  • Beräkna precisionen per klass genom att komplettera dict-comprehensionen, och iterera över .items() för attributet .class_to_idx hos dataset_test.

Interaktiv övning med praktiskt arbete

Testa den här övningen genom att slutföra den här exempelkoden.

# Define precision metric
metric_precision = Precision(
    ____, ____, ____
)

net.eval()
with torch.no_grad():
    for images, labels in dataloader_test:
        outputs = net(images)
        _, preds = torch.max(outputs, 1)
        metric_precision(preds, labels)
precision = metric_precision.compute()

# Get precision per class
precision_per_class = {
    k: ____[____].____
    for k, v 
    in ____
}
print(precision_per_class)
Redigera och kör kod