mirror of
https://github.com/dgirardeau/q3DMASC.git
synced 2026-08-29 16:40:49 +08:00
Display the confusion matrix at the end of the classification if any pre-existing 'Classification' field
+ syntax fix
This commit is contained in:
+55
-58
@@ -36,28 +36,28 @@ static QColor GetColor(double value, double r1, double g1, double b1)
|
||||
return QColor(r, g, b);
|
||||
}
|
||||
|
||||
ConfusionMatrix::ConfusionMatrix(const CCCoreLib::GenericDistribution::ScalarContainer& actual, const CCCoreLib::GenericDistribution::ScalarContainer& predicted)
|
||||
: nbClasses(0)
|
||||
, ui(new Ui::ConfusionMatrix)
|
||||
ConfusionMatrix::ConfusionMatrix( const CCCoreLib::GenericDistribution::ScalarContainer& actual,
|
||||
const CCCoreLib::GenericDistribution::ScalarContainer& predicted )
|
||||
: m_ui(new Ui::ConfusionMatrix)
|
||||
, m_overallAccuracy(0.0f)
|
||||
|
||||
{
|
||||
ui->setupUi(this);
|
||||
m_ui->setupUi(this);
|
||||
this->setWindowFlag(Qt::WindowStaysOnTopHint);
|
||||
|
||||
compute(actual, predicted);
|
||||
|
||||
this->ui->tableWidget->resizeColumnsToContents();
|
||||
this->ui->tableWidget->setSizeAdjustPolicy(QAbstractScrollArea::AdjustToContents);
|
||||
QSize tableSize = this->ui->tableWidget->sizeHint();
|
||||
this->m_ui->tableWidget->resizeColumnsToContents();
|
||||
this->m_ui->tableWidget->setSizeAdjustPolicy(QAbstractScrollArea::AdjustToContents);
|
||||
QSize tableSize = this->m_ui->tableWidget->sizeHint();
|
||||
QSize widgetSize = QSize(tableSize.width() + 30, tableSize.height() + 50);
|
||||
this->setMinimumSize(widgetSize);
|
||||
}
|
||||
|
||||
ConfusionMatrix::~ConfusionMatrix()
|
||||
{
|
||||
delete ui;
|
||||
ui = nullptr;
|
||||
delete m_ui;
|
||||
m_ui = nullptr;
|
||||
}
|
||||
|
||||
void ConfusionMatrix::computePrecisionRecallF1Score(cv::Mat& matrix, cv::Mat& precisionRecallF1Score, cv::Mat& vec_TP_FN)
|
||||
@@ -106,16 +106,13 @@ void ConfusionMatrix::computePrecisionRecallF1Score(cv::Mat& matrix, cv::Mat& pr
|
||||
// compute F1-score
|
||||
for (int realIdx = 0; realIdx < nbClasses; realIdx++)
|
||||
{
|
||||
float den = precisionRecallF1Score.at<float>(realIdx, PRECISION)
|
||||
+ precisionRecallF1Score.at<float>(realIdx, RECALL);
|
||||
float den = precisionRecallF1Score.at<float>(realIdx, PRECISION)
|
||||
+ precisionRecallF1Score.at<float>(realIdx, RECALL);
|
||||
if (den == 0)
|
||||
precisionRecallF1Score.at<float>(realIdx, F1_SCORE) = std::numeric_limits<float>::quiet_NaN();
|
||||
else
|
||||
precisionRecallF1Score.at<float>(realIdx, F1_SCORE) =
|
||||
2
|
||||
* precisionRecallF1Score.at<float>(realIdx, PRECISION)
|
||||
* precisionRecallF1Score.at<float>(realIdx, RECALL)
|
||||
/ den;
|
||||
precisionRecallF1Score.at<float>(realIdx, F1_SCORE) = ((2 * precisionRecallF1Score.at<float>(realIdx, PRECISION))
|
||||
* precisionRecallF1Score.at<float>(realIdx, RECALL))/ den;
|
||||
}
|
||||
|
||||
}
|
||||
@@ -155,8 +152,8 @@ void ConfusionMatrix::compute(const CCCoreLib::GenericDistribution::ScalarContai
|
||||
classes.insert(actual.getValue(i));
|
||||
}
|
||||
int nbClasses = static_cast<int>(classes.size());
|
||||
confusionMatrix = cv::Mat(nbClasses, nbClasses, CV_32S, cv::Scalar(0));
|
||||
precisionRecallF1Score = cv::Mat(nbClasses, 3, CV_32F, cv::Scalar(0));
|
||||
m_confusionMatrix = cv::Mat(nbClasses, nbClasses, CV_32S, cv::Scalar(0));
|
||||
m_precisionRecallF1Score = cv::Mat(nbClasses, 3, CV_32F, cv::Scalar(0));
|
||||
cv::Mat vec_TP_FN(nbClasses, 1, CV_32S, cv::Scalar(0));
|
||||
|
||||
// fill the confusion matrix
|
||||
@@ -166,68 +163,68 @@ void ConfusionMatrix::compute(const CCCoreLib::GenericDistribution::ScalarContai
|
||||
int idxActual = std::distance(classes.begin(), classes.find(actualClass));
|
||||
int predictedClass = static_cast<int>(predicted.getValue(i));
|
||||
int idxPredicted = std::distance(classes.begin(), classes.find(predictedClass));
|
||||
confusionMatrix.at<int>(idxActual, idxPredicted)++;
|
||||
m_confusionMatrix.at<int>(idxActual, idxPredicted)++;
|
||||
}
|
||||
|
||||
// compute precision recall F1-score
|
||||
computePrecisionRecallF1Score(confusionMatrix, precisionRecallF1Score, vec_TP_FN);
|
||||
float overallAccuracy = computeOverallAccuracy(confusionMatrix);
|
||||
computePrecisionRecallF1Score(m_confusionMatrix, m_precisionRecallF1Score, vec_TP_FN);
|
||||
float overallAccuracy = computeOverallAccuracy(m_confusionMatrix);
|
||||
|
||||
// display the overall accuracy
|
||||
this->ui->label_overallAccuracy->setText(QString::number(overallAccuracy, 'g', 2));
|
||||
this->m_ui->label_overallAccuracy->setText(QString::number(overallAccuracy, 'g', 2));
|
||||
|
||||
std::set<ScalarType>::iterator itB = classes.begin();
|
||||
std::set<ScalarType>::iterator itE = classes.end();
|
||||
class_numbers.assign(itB, itE);
|
||||
m_classNumbers.assign(itB, itE);
|
||||
|
||||
// BUILD THE QTABLEWIDGET
|
||||
|
||||
this->ui->tableWidget->setColumnCount(2+ nbClasses + 3); // +2 for titles, +3 for precision / recall / F1-score
|
||||
this->ui->tableWidget->setRowCount(2 + nbClasses);
|
||||
this->m_ui->tableWidget->setColumnCount(2+ nbClasses + 3); // +2 for titles, +3 for precision / recall / F1-score
|
||||
this->m_ui->tableWidget->setRowCount(2 + nbClasses);
|
||||
// create a font for the table widgets
|
||||
QFont font;
|
||||
font.setBold(true);
|
||||
QTableWidgetItem *newItem = nullptr;
|
||||
// set the row and column names
|
||||
this->ui->tableWidget->setSpan(0, 0, 2, 2); // empty area
|
||||
this->ui->tableWidget->setSpan(0, 2, 1, nbClasses); // 'Predicted' header
|
||||
this->ui->tableWidget->setSpan(2, 0, nbClasses, 1); // 'Actual' header
|
||||
this->ui->tableWidget->setSpan(0, 2 + nbClasses, 1, 3); // empty area
|
||||
this->m_ui->tableWidget->setSpan(0, 0, 2, 2); // empty area
|
||||
this->m_ui->tableWidget->setSpan(0, 2, 1, nbClasses); // 'Predicted' header
|
||||
this->m_ui->tableWidget->setSpan(2, 0, nbClasses, 1); // 'Actual' header
|
||||
this->m_ui->tableWidget->setSpan(0, 2 + nbClasses, 1, 3); // empty area
|
||||
// Predicted
|
||||
newItem = new QTableWidgetItem("Predicted");
|
||||
newItem->setFont(font);
|
||||
newItem->setBackground(Qt::lightGray);
|
||||
newItem->setTextAlignment(Qt::AlignCenter);
|
||||
this->ui->tableWidget->setItem(0, 2, newItem);
|
||||
this->m_ui->tableWidget->setItem(0, 2, newItem);
|
||||
// Real
|
||||
newItem = new QTableWidgetItem("Real");
|
||||
newItem->setFont(font);
|
||||
newItem->setBackground(Qt::lightGray);
|
||||
newItem->setTextAlignment(Qt::AlignCenter);
|
||||
this->ui->tableWidget->setItem(2, 0, newItem);
|
||||
this->m_ui->tableWidget->setItem(2, 0, newItem);
|
||||
// add precision / recall / F1-score headers
|
||||
newItem = new QTableWidgetItem("Precision");
|
||||
newItem->setToolTip("TP / (TP + FP)");
|
||||
newItem->setFont(font);
|
||||
this->ui->tableWidget->setItem(1, 2 + nbClasses + PRECISION, newItem);
|
||||
this->m_ui->tableWidget->setItem(1, 2 + nbClasses + PRECISION, newItem);
|
||||
newItem = new QTableWidgetItem("Recall");
|
||||
newItem->setToolTip("TP / (TP + FN)");
|
||||
newItem->setFont(font);
|
||||
this->ui->tableWidget->setItem(1, 2 + nbClasses + RECALL, newItem);
|
||||
this->m_ui->tableWidget->setItem(1, 2 + nbClasses + RECALL, newItem);
|
||||
newItem = new QTableWidgetItem("F1-score");
|
||||
newItem->setToolTip("Harmonic mean of precision and recall (the closer to 1 the better)\n2 x precision x recall / (precision + recall)");
|
||||
newItem->setFont(font);
|
||||
this->ui->tableWidget->setItem(1, 2 + nbClasses + F1_SCORE, newItem);
|
||||
this->m_ui->tableWidget->setItem(1, 2 + nbClasses + F1_SCORE, newItem);
|
||||
// add column names and row names
|
||||
for (int idx = 0; idx < class_numbers.size(); idx++)
|
||||
for (int idx = 0; idx < m_classNumbers.size(); idx++)
|
||||
{
|
||||
QString str = QString::number(class_numbers[idx]);
|
||||
QString str = QString::number(m_classNumbers[idx]);
|
||||
newItem = new QTableWidgetItem(str);
|
||||
newItem->setFont(font);
|
||||
this->ui->tableWidget->setItem(1, 2 + idx, newItem);
|
||||
this->m_ui->tableWidget->setItem(1, 2 + idx, newItem);
|
||||
newItem = new QTableWidgetItem(str);
|
||||
newItem->setFont(font);
|
||||
this->ui->tableWidget->setItem(2 + idx, 1, newItem);
|
||||
this->m_ui->tableWidget->setItem(2 + idx, 1, newItem);
|
||||
}
|
||||
|
||||
// FILL THE QTABLEWIDGET
|
||||
@@ -236,7 +233,7 @@ void ConfusionMatrix::compute(const CCCoreLib::GenericDistribution::ScalarContai
|
||||
for (int row = 0; row < nbClasses; row++)
|
||||
for (int column = 0; column < nbClasses; column++)
|
||||
{
|
||||
double val = confusionMatrix.at<int>(row, column);
|
||||
double val = m_confusionMatrix.at<int>(row, column);
|
||||
QTableWidgetItem *newItem = new QTableWidgetItem(QString::number(val));
|
||||
if (row == column)
|
||||
{
|
||||
@@ -246,18 +243,18 @@ void ConfusionMatrix::compute(const CCCoreLib::GenericDistribution::ScalarContai
|
||||
{
|
||||
newItem->setBackground(GetColor(val / vec_TP_FN.at<int>(row, 0), 200, 50, 50));
|
||||
}
|
||||
this->ui->tableWidget->setItem(2 + row, + 2 + column, newItem);
|
||||
this->m_ui->tableWidget->setItem(2 + row, + 2 + column, newItem);
|
||||
}
|
||||
|
||||
// set precision / recall / F1-score values
|
||||
for (int realIdx=0; realIdx < nbClasses; realIdx++)
|
||||
for (int realIdx = 0; realIdx < nbClasses; realIdx++)
|
||||
{
|
||||
newItem = new QTableWidgetItem(QString::number(precisionRecallF1Score.at<float>(realIdx, PRECISION), 'g', 2));
|
||||
this->ui->tableWidget->setItem(2 + realIdx, 2 + nbClasses + PRECISION, newItem);
|
||||
newItem = new QTableWidgetItem(QString::number(precisionRecallF1Score.at<float>(realIdx, RECALL), 'g', 2));
|
||||
this->ui->tableWidget->setItem(2 + realIdx, 2 + nbClasses + RECALL, newItem);
|
||||
newItem = new QTableWidgetItem(QString::number(precisionRecallF1Score.at<float>(realIdx, F1_SCORE), 'g', 2));
|
||||
this->ui->tableWidget->setItem(2 + realIdx, 2 + nbClasses + F1_SCORE, newItem);
|
||||
newItem = new QTableWidgetItem(QString::number(m_precisionRecallF1Score.at<float>(realIdx, PRECISION), 'g', 2));
|
||||
this->m_ui->tableWidget->setItem(2 + realIdx, 2 + nbClasses + PRECISION, newItem);
|
||||
newItem = new QTableWidgetItem(QString::number(m_precisionRecallF1Score.at<float>(realIdx, RECALL), 'g', 2));
|
||||
this->m_ui->tableWidget->setItem(2 + realIdx, 2 + nbClasses + RECALL, newItem);
|
||||
newItem = new QTableWidgetItem(QString::number(m_precisionRecallF1Score.at<float>(realIdx, F1_SCORE), 'g', 2));
|
||||
this->m_ui->tableWidget->setItem(2 + realIdx, 2 + nbClasses + F1_SCORE, newItem);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -267,7 +264,7 @@ void ConfusionMatrix::setSessionRun(QString session, int run)
|
||||
|
||||
label = session + " / " + QString::number(run);
|
||||
|
||||
this->ui->label_sessionRun->setText(label);
|
||||
this->m_ui->label_sessionRun->setText(label);
|
||||
}
|
||||
|
||||
bool ConfusionMatrix::save(QString filePath)
|
||||
@@ -283,21 +280,21 @@ bool ConfusionMatrix::save(QString filePath)
|
||||
QTextStream stream(&file);
|
||||
stream << "# columns: predicted classes\n# rows: actual classes\n";
|
||||
stream << "# last three colums: precision / recall / F1-score\n";
|
||||
for (auto class_number : class_numbers)
|
||||
for (auto classNumber : m_classNumbers)
|
||||
{
|
||||
stream << class_number << " ";
|
||||
stream << classNumber << " ";
|
||||
}
|
||||
stream << Qt::endl;
|
||||
for (int row = 0; row < confusionMatrix.rows; row++)
|
||||
for (int row = 0; row < m_confusionMatrix.rows; row++)
|
||||
{
|
||||
stream << class_numbers.at(row) << " ";
|
||||
for (int col = 0; col < confusionMatrix.cols; col++)
|
||||
stream << m_classNumbers.at(row) << " ";
|
||||
for (int col = 0; col < m_confusionMatrix.cols; col++)
|
||||
{
|
||||
stream << confusionMatrix.at<int>(row, col) << " ";
|
||||
stream << m_confusionMatrix.at<int>(row, col) << " ";
|
||||
}
|
||||
stream << precisionRecallF1Score.at<float>(row, PRECISION) << " ";
|
||||
stream << precisionRecallF1Score.at<float>(row, RECALL) << " ";
|
||||
stream << precisionRecallF1Score.at<float>(row, F1_SCORE) << Qt::endl;
|
||||
stream << m_precisionRecallF1Score.at<float>(row, PRECISION) << " ";
|
||||
stream << m_precisionRecallF1Score.at<float>(row, RECALL) << " ";
|
||||
stream << m_precisionRecallF1Score.at<float>(row, F1_SCORE) << Qt::endl;
|
||||
}
|
||||
|
||||
file.close();
|
||||
@@ -306,7 +303,7 @@ bool ConfusionMatrix::save(QString filePath)
|
||||
|
||||
}
|
||||
|
||||
float ConfusionMatrix::getOverallAccuracy()
|
||||
float ConfusionMatrix::getOverallAccuracy() const
|
||||
{
|
||||
return m_overallAccuracy;
|
||||
}
|
||||
|
||||
+9
-9
@@ -25,22 +25,22 @@ public:
|
||||
F1_SCORE = 2
|
||||
};
|
||||
|
||||
explicit ConfusionMatrix(const CCCoreLib::GenericDistribution::ScalarContainer& actual, const CCCoreLib::GenericDistribution::ScalarContainer& predicted);
|
||||
explicit ConfusionMatrix( const CCCoreLib::GenericDistribution::ScalarContainer& actual,
|
||||
const CCCoreLib::GenericDistribution::ScalarContainer& predicted );
|
||||
~ConfusionMatrix() override;
|
||||
|
||||
void computePrecisionRecallF1Score(cv::Mat& matrix, cv::Mat& precisionRecallF1Score, cv::Mat &vec_TP_FN);
|
||||
float computeOverallAccuracy(cv::Mat& matrix);
|
||||
void compute(const CCCoreLib::GenericDistribution::ScalarContainer& actual, const CCCoreLib::GenericDistribution::ScalarContainer& predicted);
|
||||
void compute( const CCCoreLib::GenericDistribution::ScalarContainer& actual,
|
||||
const CCCoreLib::GenericDistribution::ScalarContainer& predicted );
|
||||
void setSessionRun(QString session, int run);
|
||||
bool save(QString filePath);
|
||||
float getOverallAccuracy();
|
||||
float getOverallAccuracy() const;
|
||||
|
||||
private:
|
||||
std::set<ScalarType> classes;
|
||||
int nbClasses;
|
||||
Ui::ConfusionMatrix *ui;
|
||||
cv::Mat confusionMatrix;
|
||||
cv::Mat precisionRecallF1Score;
|
||||
std::vector<ScalarType> class_numbers;
|
||||
Ui::ConfusionMatrix* m_ui;
|
||||
cv::Mat m_confusionMatrix;
|
||||
cv::Mat m_precisionRecallF1Score;
|
||||
std::vector<ScalarType> m_classNumbers;
|
||||
float m_overallAccuracy;
|
||||
};
|
||||
|
||||
+1
-1
@@ -198,7 +198,7 @@ void q3DMASCPlugin::doClassifyAction()
|
||||
QString errorMessage;
|
||||
masc::Feature::Source::Set featureSources;
|
||||
masc::Feature::ExtractSources(features, featureSources);
|
||||
if (!classifier.classify(featureSources, corePoints.cloud, errorMessage, m_app->getMainWindow()))
|
||||
if (!classifier.classify(featureSources, corePoints.cloud, errorMessage, m_app->getMainWindow(), m_app))
|
||||
{
|
||||
m_app->dispToConsole(errorMessage, ccMainAppInterface::ERR_CONSOLE_MESSAGE);
|
||||
generatedScalarFields.releaseSFs(false);
|
||||
|
||||
@@ -288,8 +288,9 @@ bool Classifier::classify( const Feature::Source::Set& featureSources,
|
||||
{
|
||||
if (app)
|
||||
{
|
||||
ConfusionMatrix* confusionMatrix = new ConfusionMatrix(CCCoreLib::GenericDistribution::SFAsScalarContainer(*classifSFBackup),
|
||||
CCCoreLib::GenericDistribution::SFAsScalarContainer(*classificationSF));
|
||||
ConfusionMatrix* confusionMatrix = new ConfusionMatrix( CCCoreLib::GenericDistribution::SFAsScalarContainer(*classifSFBackup),
|
||||
CCCoreLib::GenericDistribution::SFAsScalarContainer(*classificationSF));
|
||||
confusionMatrix->show();
|
||||
}
|
||||
}
|
||||
|
||||
@@ -482,8 +483,8 @@ bool Classifier::evaluate(const Feature::Source::Set& featureSources,
|
||||
metrics.ratio = static_cast<float>(metrics.goodGuess) / metrics.sampleCount;
|
||||
}
|
||||
|
||||
ConfusionMatrix* confusionMatrix = new ConfusionMatrix(CCCoreLib::GenericDistribution::VectorAsScalarContainer(actualClass),
|
||||
CCCoreLib::GenericDistribution::VectorAsScalarContainer(predictectedClass));
|
||||
ConfusionMatrix* confusionMatrix = new ConfusionMatrix( CCCoreLib::GenericDistribution::VectorAsScalarContainer(actualClass),
|
||||
CCCoreLib::GenericDistribution::VectorAsScalarContainer(predictectedClass) );
|
||||
train3DMASCDialog.addConfusionMatrixAndSaveTraces(confusionMatrix);
|
||||
if (app)
|
||||
{
|
||||
|
||||
Reference in New Issue
Block a user