Multiclass classification curves

ROC and precision-recall curves evaluate a multiclass classifier one class at a time using one-vs-rest targets. Micro averages pool every class decision, while macro averages give each class equal weight. Precision-recall curves are especially informative when the class distribution is imbalanced.

import numpy as np
import pandas as pd
import plotly.io as pio

import qrt as q

pio.renderers.default = "notebook_connected"
rng = np.random.default_rng(7)

Example data

Create deterministic, imbalanced down, flat, and up targets, then turn noisy class scores into probabilities with a stable softmax.

classes = ["down", "flat", "up"]
y_true = np.repeat(classes, [80, 140, 60])
class_to_column = {label: index for index, label in enumerate(classes)}

raw_scores = rng.normal(size=(len(y_true), len(classes)))
true_columns = np.array([class_to_column[label] for label in y_true])
raw_scores[np.arange(len(y_true)), true_columns] += 2.0

shifted_scores = raw_scores - raw_scores.max(axis=1, keepdims=True)
probabilities = np.exp(shifted_scores)
probabilities /= probabilities.sum(axis=1, keepdims=True)
y_score = pd.DataFrame(probabilities, columns=classes)

pd.concat([pd.Series(y_true, name="actual"), y_score], axis=1).head()
actual down flat up
0 down 0.778217 0.141815 0.079969
1 down 0.750972 0.157164 0.091864
2 down 0.639106 0.311109 0.049785
3 down 0.564834 0.232026 0.203140
4 down 0.857400 0.041183 0.101417

Curve data

The q.stats functions return tidy DataFrames containing per-class, micro-average, and macro-average rows. The results can be filtered, exported, or analyzed independently of Plotly.

roc_curves = q.stats.multiclass_roc_curve(y_true, y_score)
pr_curves = q.stats.multiclass_precision_recall_curve(y_true, y_score)

roc_summary = roc_curves.assign(
    label=lambda frame: frame["class"].fillna(frame["curve"])
).groupby("label", sort=False)["auc"].first()
pr_summary = pr_curves.assign(
    label=lambda frame: frame["class"].fillna(frame["curve"])
).groupby("label", sort=False)["average_precision"].first()

pd.concat(
    [roc_summary.rename("roc_auc"), pr_summary.rename("average_precision")],
    axis=1,
)
roc_auc average_precision
label
down 0.969750 0.930457
flat 0.968776 0.969555
up 0.957197 0.897322
micro 0.967175 0.941861
macro 0.965787 0.932445

ROC curve

Each solid line is a one-vs-rest class curve. Its AUC measures ranking quality across all decision thresholds; the diagonal reference represents random ranking.

q.plot.roc(
    y_true,
    y_score,
    title="Three-class ROC curve",
).show()

Precision-recall curve

Average precision (AP) summarizes precision across recall levels. The horizontal reference is the pooled class prevalence; this view is particularly useful when good performance on minority classes matters.

q.plot.precision_recall(
    y_true,
    y_score,
    title="Three-class precision-recall curve",
).show()
Back to top