Začněte nyníZačněte zdarma

Vyhodnocení klasifikačních modelů RNN

Tým PyBooks teď chce, abys vyhodnotil/a model RNN, který jsi vytvořil/a a spustil/a na datasetu Newsgroup. Cílem bylo klasifikovat články do jedné ze tří kategorií:

rec.autos, sci.med a comp.graphics.

Model byl natrénován a ty jsi pro každou epochu vypsal/a její číslo a hodnotu ztrátové funkce.

Pomocí torchmetrics vyhodnoť různé metriky svého modelu. Pro tebe jsou již načteny tyto metriky: Accuracy, Precision, Recall, F1Score.

K dispozici máš také instanci rnn_model natrénovanou v předchozím cvičení.

Toto cvičení je součástí kurzu

Deep Learning for Text with PyTorch

Zobrazit kurz

Pokyny k cvičení

  • Vytvoř instanci každé metriky pro víceřídní klasifikaci a nastav num_classes na počet kategorií.
  • Vygeneruj predikce pro rnn_model pomocí testovacích dat X_test_seq.
  • Vypočítej metriky na základě predikovaných tříd a skutečných štítků.

Interaktivní cvičení na vyzkoušení si v praxi

Vyzkoušejte si toto cvičení dokončením tohoto ukázkového kódu.

# Create an instance of the metrics
accuracy = Accuracy(task="multiclass", ____)
precision = Precision(____, num_classes=____)
recall = Recall(task=____, num_classes=____)
f1 = F1Score(____, ____)

# Generate the predictions
outputs = ____(X_test_seq)
_, predicted = ____.____(outputs, 1)

# Calculate the metrics
accuracy_score = accuracy(____, y_test_seq)
precision_score = precision(____, y_test_seq)
recall_score = recall(____, y_test_seq)
f1_score = f1(____, y_test_seq)
print("RNN Model - Accuracy: {}, Precision: {}, Recall: {}, F1 Score: {}".format(accuracy_score, precision_score, recall_score, f1_score))
Upravit a spustit kód