This commit is contained in:
Daniel Girardeau-Montaut
2019-03-26 14:28:27 +01:00
parent e9c145ee47
commit da6293e522
14 changed files with 420 additions and 225 deletions
+95 -69
View File
@@ -48,50 +48,45 @@ bool Classifier::isValid() const
return (m_rtrees && m_rtrees->isTrained());
}
static IScalarFieldWrapper::Shared GetSource(const Feature::Shared& f, const ccPointCloud* cloud)
static IScalarFieldWrapper::Shared GetSource(const Feature::Source& fs, const ccPointCloud* cloud)
{
IScalarFieldWrapper::Shared source(nullptr);
if (!f)
{
assert(false);
ccLog::Warning(QObject::tr("Internal error: invalid feature (nullptr)"));
}
switch (f->source)
switch (fs.type)
{
case Feature::ScalarField:
case Feature::Source::ScalarField:
{
assert(!f->sourceName.isEmpty());
int sfIdx = cloud->getScalarFieldIndexByName(qPrintable(f->sourceName));
assert(fs.name.isEmpty());
int sfIdx = cloud->getScalarFieldIndexByName(qPrintable(fs.name));
if (sfIdx >= 0)
{
source.reset(new ScalarFieldWrapper(cloud->getScalarField(sfIdx)));
}
else
{
ccLog::Warning(QObject::tr("Internal error: unknwon scalar field '%1'").arg(f->sourceName));
ccLog::Warning(QObject::tr("Internal error: unknwon scalar field '%1'").arg(fs.name));
return IScalarFieldWrapper::Shared(nullptr);
}
}
break;
case Feature::DimX:
case Feature::Source::DimX:
source.reset(new DimScalarFieldWrapper(cloud, DimScalarFieldWrapper::DimX));
break;
case Feature::DimY:
case Feature::Source::DimY:
source.reset(new DimScalarFieldWrapper(cloud, DimScalarFieldWrapper::DimY));
break;
case Feature::DimZ:
case Feature::Source::DimZ:
source.reset(new DimScalarFieldWrapper(cloud, DimScalarFieldWrapper::DimZ));
break;
case Feature::Red:
case Feature::Source::Red:
source.reset(new ColorScalarFieldWrapper(cloud, ColorScalarFieldWrapper::Red));
break;
case Feature::Green:
case Feature::Source::Green:
source.reset(new ColorScalarFieldWrapper(cloud, ColorScalarFieldWrapper::Green));
break;
case Feature::Blue:
case Feature::Source::Blue:
source.reset(new ColorScalarFieldWrapper(cloud, ColorScalarFieldWrapper::Blue));
break;
}
@@ -99,7 +94,11 @@ static IScalarFieldWrapper::Shared GetSource(const Feature::Shared& f, const ccP
return source;
}
bool Classifier::classify(const Feature::Set& features, ccPointCloud* cloud, QString& errorMessage, QWidget* parentWidget/*=nullptr*/)
bool Classifier::classify( const Feature::Source::Set& featureSources,
ccPointCloud* cloud,
QString& errorMessage,
QWidget* parentWidget/*=nullptr*/
)
{
if (!cloud)
{
@@ -114,9 +113,9 @@ bool Classifier::classify(const Feature::Set& features, ccPointCloud* cloud, QSt
return false;
}
if (features.empty())
if (featureSources.empty())
{
errorMessage = QObject::tr("Training method called without any feature?!");
errorMessage = QObject::tr("Training method called without any feature (source)?!");
return false;
}
@@ -138,7 +137,7 @@ bool Classifier::classify(const Feature::Set& features, ccPointCloud* cloud, QSt
classificationSF->fill(0); //0 = no classification?
int sampleCount = static_cast<int>(cloud->size());
int attributesPerSample = static_cast<int>(features.size());
int attributesPerSample = static_cast<int>(featureSources.size());
ccLog::Print(QObject::tr("[3DMASC] Classifying %1 points with %2 feature(s)").arg(sampleCount).arg(attributesPerSample));
@@ -160,18 +159,13 @@ bool Classifier::classify(const Feature::Set& features, ccPointCloud* cloud, QSt
wrappers.reserve(attributesPerSample);
for (int fIndex = 0; fIndex < attributesPerSample; ++fIndex)
{
const Feature::Shared &f = features[fIndex];
if (!f)
{
assert(false);
return false;
}
const Feature::Source& fs = featureSources[fIndex];
IScalarFieldWrapper::Shared source = GetSource(f, cloud);
IScalarFieldWrapper::Shared source = GetSource(fs, cloud);
if (!source || !source->isValid())
{
assert(false);
errorMessage = QObject::tr("Internal error: invalid source '%1'").arg(f->sourceName);
errorMessage = QObject::tr("Internal error: invalid source '%1'").arg(fs.name);
return false;
}
@@ -226,8 +220,21 @@ bool Classifier::classify(const Feature::Set& features, ccPointCloud* cloud, QSt
return success;
}
bool Classifier::evaluate(const Feature::Set& features, CCLib::ReferenceCloud* testSubset, AccuracyMetrics& metrics, QString& errorMessage, QWidget* parentWidget/*=nullptr*/)
bool Classifier::evaluate(const Feature::Source::Set& featureSources,
ccPointCloud* testCloud,
AccuracyMetrics& metrics,
QString& errorMessage,
CCLib::ReferenceCloud* testSubset/*=nullptr=*/,
QString outputSFName/*=QString()*/,
QWidget* parentWidget/*=nullptr*/)
{
if (!testCloud)
{
//invalid input
assert(false);
errorMessage = QObject::tr("Invalid input cloud");
return false;
}
metrics.sampleCount = metrics.goodGuess = 0;
metrics.ratio = 0.0f;
@@ -237,35 +244,52 @@ bool Classifier::evaluate(const Feature::Set& features, CCLib::ReferenceCloud* t
return false;
}
if (features.empty())
if (featureSources.empty())
{
errorMessage = QObject::tr("Training method called without any feature?!");
errorMessage = QObject::tr("Training method called without any feature (source)?!");
return false;
}
if (!testSubset)
if (testSubset && testSubset->getAssociatedCloud() != testCloud)
{
assert(false);
errorMessage = QObject::tr("No test subset provided");
return false;
}
ccPointCloud* cloud = dynamic_cast<ccPointCloud*>(testSubset->getAssociatedCloud());
if (!cloud)
{
errorMessage = QObject::tr("Invalid test subset (associated point cloud is not a ccPointCloud)");
errorMessage = QObject::tr("Invalid test subset (associated point cloud is different)");
return false;
}
//look for the classification field
CCLib::ScalarField* classifSF = GetClassificationSF(cloud);
if (!classifSF || classifSF->size() < cloud->size())
CCLib::ScalarField* classifSF = GetClassificationSF(testCloud);
if (!classifSF || classifSF->size() < testCloud->size())
{
assert(false);
errorMessage = QObject::tr("Missing/Invalid 'Classification' field on input cloud");
errorMessage = QObject::tr("Missing/invalid 'Classification' field on input cloud");
return false;
}
int testSampleCount = static_cast<int>(testSubset->size());
int attributesPerSample = static_cast<int>(features.size());
CCLib::ScalarField* outputSF = nullptr;
if (!outputSFName.isEmpty())
{
int outSFIndex = testCloud->getScalarFieldIndexByName(qPrintable(outputSFName));
if (outSFIndex < 0)
{
ccScalarField* _outputSF = new ccScalarField(qPrintable(outputSFName));
if (!_outputSF->resizeSafe(testCloud->size()))
{
errorMessage = QObject::tr("Not enough memory to create output scalar field");
_outputSF->release();
return false;
}
testCloud->addScalarField(_outputSF);
outputSF = _outputSF;
}
else
{
outputSF = testCloud->getScalarField(outSFIndex);
}
outputSF->fill(NAN_VALUE);
outputSF->computeMinAndMax();
}
unsigned testSampleCount = (testSubset ? testSubset->size() : testCloud->size());
int attributesPerSample = static_cast<int>(featureSources.size());
ccLog::Print(QObject::tr("[3DMASC] Testing data: %1 samples with %2 feature(s)").arg(testSampleCount).arg(attributesPerSample));
@@ -273,7 +297,7 @@ bool Classifier::evaluate(const Feature::Set& features, CCLib::ReferenceCloud* t
cv::Mat test_data;
try
{
test_data.create(testSampleCount, attributesPerSample, CV_32FC1);
test_data.create(static_cast<int>(testSampleCount), attributesPerSample, CV_32FC1);
}
catch (const cv::Exception& cvex)
{
@@ -294,24 +318,18 @@ bool Classifier::evaluate(const Feature::Set& features, CCLib::ReferenceCloud* t
//fill the data matrix
for (int fIndex = 0; fIndex < attributesPerSample; ++fIndex)
{
const Feature::Shared &f = features[fIndex];
if (!f)
{
assert(false);
return false;
}
IScalarFieldWrapper::Shared source = GetSource(f, cloud);
const Feature::Source& fs = featureSources[fIndex];
IScalarFieldWrapper::Shared source = GetSource(fs, testCloud);
if (!source || !source->isValid())
{
assert(false);
errorMessage = QObject::tr("Internal error: invalid source '%1'").arg(f->sourceName);
errorMessage = QObject::tr("Internal error: invalid source '%1'").arg(fs.name);
return false;
}
for (unsigned i = 0; i < testSubset->size(); ++i)
for (unsigned i = 0; i < testSampleCount; ++i)
{
unsigned pointIndex = testSubset->getPointGlobalIndex(i);
unsigned pointIndex = (testSubset ? testSubset->getPointGlobalIndex(i) : i);
double value = source->pointValue(pointIndex);
test_data.at<float>(i, fIndex) = static_cast<float>(value);
}
@@ -319,12 +337,12 @@ bool Classifier::evaluate(const Feature::Set& features, CCLib::ReferenceCloud* t
//estimate the efficiency of the classifier
{
metrics.sampleCount = testSubset->size();
metrics.sampleCount = testSampleCount;
metrics.goodGuess = 0;
for (unsigned i = 0; i < testSubset->size(); ++i)
for (unsigned i = 0; i < testSampleCount; ++i)
{
unsigned pointIndex = testSubset->getPointGlobalIndex(i);
unsigned pointIndex = (testSubset ? testSubset->getPointGlobalIndex(i) : i);
ScalarType pointClass = classifSF->getValue(pointIndex);
int iClass = static_cast<int>(pointClass);
//if (iClass < 0 || iClass > 255)
@@ -334,10 +352,15 @@ bool Classifier::evaluate(const Feature::Set& features, CCLib::ReferenceCloud* t
//}
float predictedClass = m_rtrees->predict(test_data.row(i));
if (static_cast<int>(predictedClass) == iClass)
int iPredictedClass = static_cast<int>(predictedClass);
if (iPredictedClass == iClass)
{
++metrics.goodGuess;
}
if (outputSF)
{
outputSF->setValue(pointIndex, static_cast<ScalarType>(iPredictedClass));
}
if (pDlg && !nProgress.oneStep())
{
@@ -346,6 +369,9 @@ bool Classifier::evaluate(const Feature::Set& features, CCLib::ReferenceCloud* t
}
}
if (outputSF)
outputSF->computeMinAndMax();
metrics.ratio = static_cast<float>(metrics.goodGuess) / metrics.sampleCount;
}
@@ -354,15 +380,15 @@ bool Classifier::evaluate(const Feature::Set& features, CCLib::ReferenceCloud* t
bool Classifier::train( const ccPointCloud* cloud,
const RandomTreesParams& params,
const Feature::Set& features,
const Feature::Source::Set& featureSources,
QString& errorMessage,
CCLib::ReferenceCloud* trainSubset/*=nullptr*/,
ccMainAppInterface* app/*=nullptr*/,
QWidget* parentWidget/*=nullptr*/)
{
if (features.empty())
if (featureSources.empty())
{
errorMessage = QObject::tr("Training method called without any feature?!");
errorMessage = QObject::tr("Training method called without any feature (source)?!");
return false;
}
if (!cloud)
@@ -387,7 +413,7 @@ bool Classifier::train( const ccPointCloud* cloud,
}
int sampleCount = static_cast<int>(trainSubset ? trainSubset->size() : cloud->size());
int attributesPerSample = static_cast<int>(features.size());
int attributesPerSample = static_cast<int>(featureSources.size());
if (app)
{
@@ -426,13 +452,13 @@ bool Classifier::train( const ccPointCloud* cloud,
//fill the training data matrix
for (int fIndex = 0; fIndex < attributesPerSample; ++fIndex)
{
const Feature::Shared &f = features[fIndex];
const Feature::Source& fs = featureSources[fIndex];
IScalarFieldWrapper::Shared source = GetSource(f, cloud);
IScalarFieldWrapper::Shared source = GetSource(fs, cloud);
if (!source || !source->isValid())
{
assert(false);
errorMessage = QObject::tr("Internal error: invalid source '%1'").arg(f->sourceName);
errorMessage = QObject::tr("Internal error: invalid source '%1'").arg(fs.name);
return false;
}