Understanding tf.keras.metrics.Precision and Recall for multiclass classification

Viewed 201

I am building a model for a multiclass classification problem. So I want to evaluate the model performance using the Recall and Precision. I have 4 classes in the dataset and it is provided in one hot representation.

I was reading the Precision and Recall tf.keras documentation, and have some questions:

  1. When it calculating the Precision and Recall for the multi-class classification, how can we take the average of all of the labels, meaning the global precision & Recall? is it calculated with macro or micro since it is not specified in the documentation as in the Sikit learn.
  2. If I want to calculate the precision & Recall for each label separately, can I use the argument class_id for each label to do one_vs_rest or binary classification. Like what I have done in the code below?
  3. can I use the argument top_k with the value top_k=2 would be helpful here or it is not suitable for my classification of 4 classes only?
  4. While I am measuring the performance of each class, What could be the difference, when I set the top_k=1 and not setting top_koverall?
model.compile(
      optimizer='sgd',
      loss=tf.keras.losses.CategoricalCrossentropy(),
      metrics=[tf.keras.metrics.CategoricalAccuracy(),
               ##class 0
               tf.keras.metrics.Precision(class_id=0,top_k=2), 
               tf.keras.metrics.Recall(class_id=0,top_k=2),
              ##class 1
               tf.keras.metrics.Precision(class_id=1,top_k=2), 
               tf.keras.metrics.Recall(class_id=1,top_k=2),
              ##class 2
               tf.keras.metrics.Precision(class_id=2,top_k=2), 
               tf.keras.metrics.Recall(class_id=2,top_k=2),
              ##class 3
               tf.keras.metrics.Precision(class_id=3,top_k=2), 
               tf.keras.metrics.Recall(class_id=3,top_k=2),
])

Any clarification of this function will be appreciated. Thanks in advance

2 Answers

3. can I use the argument top_k with the value top_k=2 would be helpful here or it is not suitable for my classification of 4 classes only?

According to the description, it will only calculate top_k(with the function of _filter_top_k) predictions, and turn other predictions to False if you use this argument

The example from official document link:https://www.tensorflow.org/api_docs/python/tf/keras/metrics/Precision

You may also want to read the original code: https://github.com/keras-team/keras/blob/07e13740fd181fc3ddec7d9a594d8a08666645f6/keras/utils/metrics_utils.py#L487 With top_k=2, it will calculate precision over y_true[:2] and y_pred[:2]

m = tf.keras.metrics.Precision(top_k=2)
m.update_state([0, 0, 1, 1], [1, 1, 1, 1])
m.result().numpy()
0.0

As we can see the note posted in the example here, it will only calculate y_true[:2] and y_pred[:2], which means the precision will calculate only top 2 predictions (also turn the rest of y_pred to 0).

If you want to use 4 classes classification, the argument of class_id maybe enough.

4.While I am measuring the performance of each class, What could be the difference when I set the top_k=1 and not setting top_koverall? The function will calculate the precision across all the predictions your model make if you don't set top_k value. If you want to measure the perfromance.

Top k may works for other model, not for classification model

1. Is it macro or micro ?

To be precise, all the metrics are reset at the beginning of every epoch and at the beginning of every validation if there is. So I guess, we can call it macro.

2. Class specific precision and recall ?

You can take a look at tf.compat.v1.metrics.precision_at_k and tf.compat.v1.metrics.recall_at_k. It seems that it computes the respectivly the precision at the recall for a specific class k.

https://www.tensorflow.org/api_docs/python/tf/compat/v1/metrics/precision_at_k

https://www.tensorflow.org/api_docs/python/tf/compat/v1/metrics/recall_at_k

Related