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)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.
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()