Added a cache for track features, made the crossover ratio configurable, now using a fitness proportional selection

This commit is contained in:
emeric
2019-12-03 13:47:48 +01:00
parent 93cd3b6020
commit 74a5f6389d
6 changed files with 181 additions and 46 deletions
@@ -36,43 +36,52 @@
namespace Similarity {
static
std::optional<SOM::InputVector>
getInputVectorFromTrack(Database::Session& session, Database::IdType trackId, const std::unordered_set<FeatureName>& featureNames, std::size_t nbDimensions)
std::optional<FeatureValuesMap>
getTrackFeatureValues(FeaturesSearcher::FeaturesFetchFunc func, Database::IdType trackId, const std::unordered_set<FeatureName>& featureNames)
{
FeatureValuesMap featureValuesMap;
return func(trackId, featureNames);
}
static
std::optional<FeatureValuesMap>
getTrackFeatureValuesFromDb(Database::Session& session, Database::IdType trackId, const std::unordered_set<FeatureName>& featureNames)
{
auto func = [&](Database::IdType trackId, const std::unordered_set<FeatureName>& featureNames)
{
std::optional<FeatureValuesMap> res;
auto transaction {session.createSharedTransaction()};
Database::Track::pointer track {Database::Track::getById(session, trackId)};
if (!track)
return std::nullopt;
return res;
featureValuesMap = track->getTrackFeatures()->getFeatureValuesMap(featureNames);
if (featureValuesMap.empty())
return std::nullopt;
}
res = track->getTrackFeatures()->getFeatureValuesMap(featureNames);
if (res->empty())
res.reset();
return res;
};
return getTrackFeatureValues(func, trackId, featureNames);
}
static
std::optional<SOM::InputVector>
convertFeatureValuesMapToInputVector(const FeatureValuesMap& featureValuesMap, std::size_t nbDimensions)
{
std::size_t i {};
std::optional<SOM::InputVector> res {SOM::InputVector {nbDimensions}};
for (const auto& featureName : featureNames)
for (const auto& [featureName, values] : featureValuesMap)
{
const auto it {featureValuesMap.find(featureName)};
if (it == std::cend(featureValuesMap))
if (values.size() != getFeatureDef(featureName).nbDimensions)
{
LMS_LOG(SIMILARITY, WARNING) << "Cannot find feature '" << featureName << "' for track id'" << trackId << "'";
LMS_LOG(SIMILARITY, WARNING) << "Dimension mismatch for feature '" << featureName << "'. Expected " << getFeatureDef(featureName).nbDimensions << ", got " << values.size();
res.reset();
break;
}
if (it->second.size() != getFeatureDef(featureName).nbDimensions)
{
LMS_LOG(SIMILARITY, WARNING) << "Dimension mismatch for feature '" << featureName << "'. Expected " << getFeatureDef(featureName).nbDimensions << ", got " << it->second.size() << ", trackId = " << trackId;
res.reset();
break;
}
for (double val : it->second)
for (double val : values)
(*res)[i++] = val;
}
@@ -134,7 +143,17 @@ FeaturesSearcher::FeaturesSearcher(Database::Session& session,
if (stopRequested && stopRequested())
return;
std::optional<SOM::InputVector> inputVector {getInputVectorFromTrack(session, trackId, featureNames, nbDimensions)};
std::optional<FeatureValuesMap> featureValuesMap;
if (_featuresFetchFunc)
featureValuesMap = getTrackFeatureValues(_featuresFetchFunc, trackId, featureNames);
else
featureValuesMap = getTrackFeatureValuesFromDb(session, trackId, featureNames);
if (!featureValuesMap)
continue;
std::optional<SOM::InputVector> inputVector {convertFeatureValuesMapToInputVector(*featureValuesMap, nbDimensions)};
if (!inputVector)
continue;
@@ -20,6 +20,7 @@
#pragma once
#include <map>
#include <optional>
#include <set>
#include <string>
@@ -70,6 +71,11 @@ class FeaturesSearcher
FeaturesCache toCache() const;
using FeaturesFetchFunc = std::function<std::optional<std::unordered_map<std::string, std::vector<double>>>(Database::IdType /*trackId*/, const std::unordered_set<std::string>& /*features*/)>;
// Default is to retrieve the features from the database (may be slow).
// Use this only if you want to train different searchers with the same data
static void setFeaturesFetchFunc(FeaturesFetchFunc func) { _featuresFetchFunc = func; }
private:
using ObjectPositions = std::map<Database::IdType, std::set<SOM::Position>>;
@@ -96,6 +102,7 @@ class FeaturesSearcher
SOM::Matrix<std::set<Database::IdType>> _tracksMap;
ObjectPositions _trackPositions;
static inline FeaturesFetchFunc _featuresFetchFunc;
};
} // ns Similarity
-7
View File
@@ -176,10 +176,3 @@ RandGenerator& getRandGenerator()
return randGenerator;
}
int
getRandom(int min, int max)
{
std::uniform_int_distribution<> dist {min, max};
return dist (getRandGenerator());
}
+15 -2
View File
@@ -113,8 +113,21 @@ constexpr T clamp(T v, T lo, T hi, Compare comp = {})
using RandGenerator = std::mt19937;
RandGenerator& getRandGenerator();
int
getRandom(int min, int max);
template <typename T>
T
getRandom(T min, T max)
{
std::uniform_int_distribution<> dist {min, max};
return dist (getRandGenerator());
}
template <typename T>
T
getRealRandom(T min, T max)
{
std::uniform_real_distribution<> dist {min, max};
return dist (getRandGenerator());
}
template <typename Container>
void