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
Ö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_idxhosdataset_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)