Analiza metricilor pe clasă
Metricile agregate sunt indicatori utili ai performanței modelului, însă de multe ori este informativ să le analizezi pe fiecare clasă în parte. Acest lucru poate scoate la iveală clasele pentru care modelul nu performează corespunzător.
În acest exercițiu, vei rula din nou bucla de evaluare pentru a obține precizia clasificatorului de nori, de această dată pe fiecare clasă. Apoi, vei asocia aceste scoruri cu numele claselor pentru a le interpreta. Ca de obicei, Precision a fost deja importat. Mult succes!
Acest exercițiu face parte din cursul
Deep Learning intermediar cu PyTorch
Instrucțiuni pentru exercițiu
- Definește o metrică de precizie potrivită pentru rezultate pe clasă.
- Calculează precizia per clasă finalizând comprehensiunea de dicționar, iterând peste
.items()ale atributului.class_to_idxaldataset_test.
Exercițiu interactiv practic
Încearcă acest exercițiu completând acest cod de exemplu.
# 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)