11float mse(
const VectorXf& y,
const VectorXf& yhat, VectorXf& loss,
12 const vector<float>& class_weights)
14 loss = (yhat - y).array().pow(2);
19VectorXf
log_loss(
const VectorXf& y,
const VectorXf& predict_proba,
20 const vector<float>& class_weights)
28 loss.resize(y.rows());
29 for (
unsigned i = 0; i < y.rows(); ++i)
31 if (predict_proba(i) < eps || 1 - predict_proba(i) < eps)
33 loss(i) = -(y(i)*log(eps) + (1-y(i))*log(1-eps));
35 loss(i) = -(y(i)*log(predict_proba(i)) + (1-y(i))*log(1-predict_proba(i)));
38 std::runtime_error(
"loss(i)= " +
to_string(loss(i))
39 +
". y = " +
to_string(y(i)) +
", predict_proba(i) = "
48 const VectorXf& predict_proba, VectorXf& loss,
49 const vector<float>& class_weights)
51 loss =
log_loss(y,predict_proba,class_weights);
53 if (!class_weights.empty())
55 float sum_weights = 0;
58 VectorXf weighted_loss;
59 weighted_loss.resize(y.rows());
60 for (
unsigned i = 0; i < y.rows(); ++i)
62 weighted_loss(i) = loss(i) * class_weights.at(y(i));
63 sum_weights += class_weights.at(y(i));
68 return weighted_loss.sum() / sum_weights;
76 const VectorXf& predict_proba, VectorXf& loss,
77 const vector<float>& class_weights )
79 VectorXi yhat = (predict_proba.array() > 0.5).cast<
int>();
82 loss = (yhat.array() != y.cast<
int>().array()).cast<
float>();
86 if (!class_weights.empty()) {
87 for (
int i = 0; i < y.rows(); ++i) {
88 loss(i) *= class_weights.at(y(i));
89 scale += class_weights.at(y(i));
94 scale =
static_cast<float>(loss.size());
98 return 1.0 - (loss.sum() / scale);
103 const VectorXf& predict_proba, VectorXf& loss,
104 const vector<float>& class_weights )
106 VectorXi yhat = (predict_proba.array() > 0.5).cast<
int>();
108 loss = (yhat.array() != y.cast<
int>().array()).cast<
float>();
115 int num_instances = y.rows();
116 for (
int i = 0; i < num_instances; ++i) {
120 if (yhat(i) == 1.0 && y(i) == 1.0) TP += weight;
121 else if (yhat(i) == 1.0 && y(i) == 0.0) FP += weight;
122 else if (yhat(i) == 0.0 && y(i) == 0.0) TN += weight;
128 float TPR = (TP + eps) / (TP + FN + eps);
129 float TNR = (TN + eps) / (TN + FP + eps);
131 return (TPR + TNR) / 2.0;
138vector<float> sample_weights(
const VectorXf& y,
const vector<float>& class_weights)
140 vector<float> w(y.size(), 1.0f);
141 if (!class_weights.empty())
142 for (
int i = 0; i < y.size(); ++i)
143 w[i] = class_weights.at(
static_cast<int>(y(i)));
148float binary_average_precision(
const VectorXf& y,
const VectorXf& predict_proba,
149 const vector<float>& w)
151 int num_instances = y.size();
155 vector<int> order(num_instances);
156 iota(order.begin(), order.end(), 0);
157 stable_sort(order.begin(), order.end(), [&](
int i,
int j) {
158 return predict_proba(i) > predict_proba(j);
162 vector<float> y_sorted(num_instances);
163 vector<float> p_sorted(num_instances);
164 vector<float> w_sorted(num_instances);
165 for (
int i = 0; i < num_instances; ++i) {
168 y_sorted[i] = y(idx);
169 p_sorted[i] = predict_proba(idx);
170 w_sorted[i] = w[idx];
172 ysum += y_sorted[i] * w_sorted[i];
183 if (fabs(p_sorted.back() - p_sorted.front()) <= eps) {
185 float total_weight = std::accumulate(w_sorted.begin(), w_sorted.end(), 0.0f);
189 return total_weight == 0.0f ? 0.0f : ysum / total_weight;
193 vector<int> unique_indices = {};
194 set<float> unique_probas = {};
196 for (
int i=0; i<p_sorted.size(); ++i)
197 if (unique_probas.insert(p_sorted.at(i)).second)
198 unique_indices.push_back(i);
200 unique_indices.push_back(num_instances);
204 vector<float> precision = {1.0};
205 vector<float> recall = {0.0};
207 for (
size_t i = 0; i < unique_indices.size() - 1; ++i) {
208 int start = unique_indices[i];
209 int end = unique_indices[i+1];
212 for (
int j = start; j < end; ++j) {
213 tp += y_sorted.at(j) * w_sorted.at(j);
214 fp += (1.0f - y_sorted.at(j)) * w_sorted.at(j);
216 float relevant = tp + fp;
217 precision.push_back(relevant == 0.0f ? 0.0f : tp / relevant);
218 recall.push_back(ysum == 0.0f ? 1.0f : tp / ysum);
223 float average_precision = 0.0f;
224 for (
size_t i = 0; i < precision.size() - 1; ++i) {
225 average_precision += (recall[i+1] - recall[i]) * precision[i+1];
228 return average_precision;
234float binary_roc_auc(
const VectorXf& y,
const VectorXf& predict_proba,
235 const vector<float>& w)
237 int num_instances = y.size();
239 vector<int> order(num_instances);
240 iota(order.begin(), order.end(), 0);
241 stable_sort(order.begin(), order.end(), [&](
int i,
int j) {
242 return predict_proba(i) > predict_proba(j);
247 for (
int i = 0; i < num_instances; ++i) {
250 neg += (1.0f - y(i)) * w[i];
254 if (pos == 0.0f || neg == 0.0f)
257 float tp = 0.0f, fp = 0.0f;
258 float tp_prev = 0.0f, fp_prev = 0.0f;
260 for (
int i = 0; i < num_instances; ++i) {
262 tp += y(idx) * w[idx];
263 fp += (1.0f - y(idx)) * w[idx];
266 bool last_of_block = (i == num_instances - 1)
267 || (predict_proba(order[i+1]) != predict_proba(idx));
270 area += (fp - fp_prev) * (tp + tp_prev) / 2.0f;
276 return area / (pos * neg);
280void confusion(
const VectorXf& y,
const ArrayXi& yhat,
int label,
281 const vector<float>& w,
float& TP,
float& FP,
float& FN)
287 for (
int i = 0; i < y.size(); ++i) {
288 bool is_true =
static_cast<int>(y(i)) == label;
289 bool is_pred = yhat(i) == label;
291 if ( is_true && is_pred) TP += w[i];
292 else if (!is_true && is_pred) FP += w[i];
293 else if ( is_true && !is_pred) FN += w[i];
297ArrayXi argmax_rows(
const ArrayXXf& predict_proba)
301 ArrayXi yhat(predict_proba.rows());
302 for (
int i = 0; i < predict_proba.rows(); ++i)
303 predict_proba.row(i).maxCoeff(&yhat(i));
310float multi_macro_precision_recall(
const VectorXf& y,
const ArrayXXf& predict_proba,
311 VectorXf& loss,
const vector<float>& class_weights,
314 if (predict_proba.rows() != y.rows())
315 HANDLE_ERROR_THROW(
"Multiclass probabilities and labels have different numbers of rows");
317 ArrayXi yhat = argmax_rows(predict_proba);
320 loss = (yhat != y.cast<
int>().array()).cast<
float>();
322 vector<float> w = sample_weights(y, class_weights);
326 for (
int label = 0; label < predict_proba.cols(); ++label) {
327 bool present = (y.cast<
int>().array() == label).any() || (yhat == label).any();
332 confusion(y, yhat, label, w, TP, FP, FN);
334 float denom = precision ? TP + FP : TP + FN;
335 sum += denom == 0.0f ? 0.0f : TP / denom;
338 return n_labels == 0 ? 0.0f : sum / n_labels;
345 const vector<float>& class_weights) {
353 loss =
log_loss(y, predict_proba, class_weights);
355 return binary_average_precision(y, predict_proba, sample_weights(y, class_weights));
361 VectorXf& loss,
const vector<float>& class_weights)
363 ArrayXi yhat = (predict_proba.array() > 0.5).cast<
int>();
366 loss = (yhat != y.cast<
int>().array()).cast<
float>();
369 confusion(y, yhat, 1, sample_weights(y, class_weights), TP, FP, FN);
371 return (TP + FP) == 0.0f ? 0.0f : TP / (TP + FP);
375 VectorXf& loss,
const vector<float>& class_weights)
377 ArrayXi yhat = (predict_proba.array() > 0.5).cast<
int>();
378 loss = (yhat != y.cast<
int>().array()).cast<
float>();
381 confusion(y, yhat, 1, sample_weights(y, class_weights), TP, FP, FN);
383 return (TP + FN) == 0.0f ? 0.0f : TP / (TP + FN);
387 VectorXf& loss,
const vector<float>& class_weights)
390 loss =
log_loss(y, predict_proba, class_weights);
392 return binary_roc_auc(y, predict_proba, sample_weights(y, class_weights));
397 const vector<float>& class_weights)
399 if (predict_proba.rows() != y.rows())
400 HANDLE_ERROR_THROW(
"Multiclass probabilities and labels have different numbers of rows");
402 constexpr float eps = 1e-6f;
403 VectorXf loss(y.rows());
404 for (
int i = 0; i < y.rows(); ++i)
406 const int label =
static_cast<int>(y(i));
412 loss(i) = -std::log(std::clamp(predict_proba(i, label), eps, 1.0f - eps));
418 const ArrayXXf& predict_proba, VectorXf& loss,
419 const vector<float>& class_weights)
423 if (class_weights.empty())
427 float sum_weights = 0.0f;
428 float weighted_loss = 0.0f;
429 for (
int i = 0; i < y.rows(); ++i)
431 const float weight = class_weights.at(
static_cast<int>(y(i)));
432 weighted_loss += loss(i) * weight;
433 sum_weights += weight;
435 return sum_weights == 0.0f ? 0.0f : weighted_loss / sum_weights;
439 const ArrayXXf& predict_proba, VectorXf& loss,
440 const vector<float>& class_weights )
442 if (predict_proba.rows() != y.rows())
443 HANDLE_ERROR_THROW(
"Multiclass probabilities and labels have different numbers of rows");
446 for (
int i = 0; i < predict_proba.rows(); ++i)
447 predict_proba.row(i).maxCoeff(&yhat(i));
449 loss = (yhat.array() != y.cast<
int>().array()).cast<
float>();
451 if (class_weights.empty())
452 return 1.0f - loss.mean();
454 float weighted_errors = 0.0f;
455 float sum_weights = 0.0f;
456 for (
int i = 0; i < y.rows(); ++i)
458 const float weight = class_weights.at(
static_cast<int>(y(i)));
459 weighted_errors += loss(i) * weight;
460 sum_weights += weight;
462 return sum_weights == 0.0f ? 0.0f : 1.0f - weighted_errors / sum_weights;
466 const ArrayXXf& predict_proba, VectorXf& loss,
467 const vector<float>& class_weights)
469 if (predict_proba.rows() != y.rows())
470 HANDLE_ERROR_THROW(
"Multiclass probabilities and labels have different numbers of rows");
473 for (
int i = 0; i < predict_proba.rows(); ++i)
474 predict_proba.row(i).maxCoeff(&yhat(i));
475 loss = (yhat.array() != y.cast<
int>().array()).cast<
float>();
477 VectorXf correct = VectorXf::Zero(predict_proba.cols());
478 VectorXf support = VectorXf::Zero(predict_proba.cols());
479 for (
int i = 0; i < y.rows(); ++i)
481 const int label =
static_cast<int>(y(i));
487 support(label) += 1.0f;
488 if (yhat(i) == label)
489 correct(label) += 1.0f;
492 float recall_sum = 0.0f;
493 int present_classes = 0;
494 for (
int label = 0; label < support.size(); ++label)
495 if (support(label) > 0.0f)
497 recall_sum += correct(label) / support(label);
500 return present_classes == 0 ? 0.0f : recall_sum / present_classes;
504 VectorXf& loss,
const vector<float>& class_weights)
506 return multi_macro_precision_recall(y, predict_proba, loss, class_weights,
true);
510 VectorXf& loss,
const vector<float>& class_weights)
512 return multi_macro_precision_recall(y, predict_proba, loss, class_weights,
false);
516 VectorXf& loss,
const vector<float>& class_weights)
520 vector<float> w = sample_weights(y, class_weights);
524 for (
int label = 0; label < predict_proba.cols(); ++label) {
525 VectorXf y_bin = (y.cast<
int>().array() == label).cast<
float>();
528 if (y_bin.sum() == 0.0f || y_bin.sum() == y_bin.size())
531 sum += binary_roc_auc(y_bin, predict_proba.col(label).matrix(), w);
534 return n_labels == 0 ? 0.5f : sum / n_labels;
538 VectorXf& loss,
const vector<float>& class_weights)
542 vector<float> w = sample_weights(y, class_weights);
546 for (
int label = 0; label < predict_proba.cols(); ++label) {
547 VectorXf y_bin = (y.cast<
int>().array() == label).cast<
float>();
550 if (y_bin.sum() == 0.0f)
553 sum += binary_average_precision(y_bin, predict_proba.col(label).matrix(), w);
556 return n_labels == 0 ? 0.0f : sum / n_labels;
#define HANDLE_ERROR_THROW(err)
float multi_zero_one_loss(const VectorXf &y, const ArrayXXf &predict_proba, VectorXf &loss, const vector< float > &class_weights)
Accuracy for multi-classification.
float precision_score(const VectorXf &y, const VectorXf &predict_proba, VectorXf &loss, const vector< float > &class_weights)
Precision for binary classification (threshold 0.5, positive label 1).
float zero_one_loss(const VectorXf &y, const VectorXf &predict_proba, VectorXf &loss, const vector< float > &class_weights)
Accuracy for binary classification.
float mean_log_loss(const VectorXf &y, const VectorXf &predict_proba, VectorXf &loss, const vector< float > &class_weights)
log loss
float multi_recall_score(const VectorXf &y, const ArrayXXf &predict_proba, VectorXf &loss, const vector< float > &class_weights)
Macro-averaged recall for multi-classification.
float mean_multi_log_loss(const VectorXf &y, const ArrayXXf &predict_proba, VectorXf &loss, const vector< float > &class_weights)
Calculates the mean multinomial log loss between the predicted probabilities and the true labels.
float average_precision_score(const VectorXf &y, const VectorXf &predict_proba, VectorXf &loss, const vector< float > &class_weights)
Calculates the average precision score between the predicted probabilities and the true labels.
float mse(const VectorXf &y, const VectorXf &yhat, VectorXf &loss, const vector< float > &class_weights)
mean squared error
VectorXf multi_log_loss(const VectorXf &y, const ArrayXXf &predict_proba, const vector< float > &class_weights)
Calculates the multinomial log loss between the predicted probabilities and the true labels.
float multi_bal_zero_one_loss(const VectorXf &y, const ArrayXXf &predict_proba, VectorXf &loss, const vector< float > &class_weights)
Balanced accuracy for multi-classification.
float recall_score(const VectorXf &y, const VectorXf &predict_proba, VectorXf &loss, const vector< float > &class_weights)
Recall for binary classification (threshold 0.5, positive label 1).
float multi_roc_auc_score(const VectorXf &y, const ArrayXXf &predict_proba, VectorXf &loss, const vector< float > &class_weights)
Macro-averaged one-vs-rest AUROC for multi-classification.
VectorXf log_loss(const VectorXf &y, const VectorXf &predict_proba, const vector< float > &class_weights)
Calculates the log loss between the predicted probabilities and the true labels.
float multi_precision_score(const VectorXf &y, const ArrayXXf &predict_proba, VectorXf &loss, const vector< float > &class_weights)
Macro-averaged precision for multi-classification.
float multi_average_precision_score(const VectorXf &y, const ArrayXXf &predict_proba, VectorXf &loss, const vector< float > &class_weights)
Macro-averaged one-vs-rest average precision for multi-classification.
float bal_zero_one_loss(const VectorXf &y, const VectorXf &predict_proba, VectorXf &loss, const vector< float > &class_weights)
Balanced accuracy for binary classification.
float roc_auc_score(const VectorXf &y, const VectorXf &predict_proba, VectorXf &loss, const vector< float > &class_weights)
Area under the ROC curve for binary classification.
string to_string(const T &value)
template function to convert objects to string for logging
< nsga2 selection operator for getting the front
Eigen::Array< int, Eigen::Dynamic, 1 > ArrayXi