クラス別メトリクスの分析
集計されたメトリクスはモデル性能の有用な指標ですが、クラスごとの値を見るとさらに多くの示唆が得られます。モデルが苦手としているクラスが見つかるかもしれません。
この演習では、評価ループをもう一度実行して、クラウド分類器の precision をクラス別に取得します。次に、そのスコアをクラス名に対応づけて解釈します。いつも通り、Precision はすでにインポート済みです。頑張ってください!
この演習はコースの一部です
PyTorchによる中級ディープラーニング
演習の手順
- クラス別の結果に適した precision メトリクスを定義します。
dataset_testの.class_to_idx属性の.items()を反復し、辞書内包表記を完成させてクラスごとの precision を計算します。
実践的なインタラクティブ演習
このサンプルコードを完成させて、この演習に挑戦してみましょう。
# 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)