Databsae service refactoring. Warning, loses stars and listens stats
This commit is contained in:
@@ -19,12 +19,11 @@
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <functional>
|
||||
#include <memory>
|
||||
#include <string_view>
|
||||
#include "services/database/Types.hpp"
|
||||
#include "utils/EnumSet.hpp"
|
||||
#include "services/database/TrackListId.hpp"
|
||||
#include "services/recommendation/Types.hpp"
|
||||
#include "utils/EnumSet.hpp"
|
||||
|
||||
namespace Database
|
||||
{
|
||||
@@ -41,8 +40,8 @@ namespace Recommendation
|
||||
virtual void load(bool forceReload, const ProgressCallback& progressCallback = {}) = 0;
|
||||
virtual void requestCancelLoad() = 0;
|
||||
|
||||
virtual TrackContainer getSimilarTracksFromTrackList(Database::TrackListId tracklistId, std::size_t maxCount) const = 0;
|
||||
virtual TrackContainer getSimilarTracks(const std::vector<Database::TrackId>& tracksId, std::size_t maxCount) const = 0;
|
||||
virtual TrackContainer findSimilarTracksFromTrackList(Database::TrackListId tracklistId, std::size_t maxCount) const = 0;
|
||||
virtual TrackContainer findSimilarTracks(const std::vector<Database::TrackId>& tracksId, std::size_t maxCount) const = 0;
|
||||
virtual ReleaseContainer getSimilarReleases(Database::ReleaseId releaseId, std::size_t maxCount) const = 0;
|
||||
virtual ArtistContainer getSimilarArtists(Database::ArtistId artistId, EnumSet<Database::TrackArtistLinkType> linkTypes, std::size_t maxCount) const = 0;
|
||||
};
|
||||
|
||||
@@ -59,7 +59,7 @@ namespace Recommendation
|
||||
}
|
||||
|
||||
TrackContainer
|
||||
RecommendationService::getSimilarTracksFromTrackList(Database::TrackListId trackListId, std::size_t maxCount) const
|
||||
RecommendationService::findSimilarTracksFromTrackList(Database::TrackListId trackListId, std::size_t maxCount) const
|
||||
{
|
||||
TrackContainer res;
|
||||
|
||||
@@ -70,7 +70,7 @@ namespace Recommendation
|
||||
if (itEngine == std::cend(_engines))
|
||||
continue;
|
||||
|
||||
res = itEngine->second->getSimilarTracksFromTrackList(trackListId, maxCount);
|
||||
res = itEngine->second->findSimilarTracksFromTrackList(trackListId, maxCount);
|
||||
if (!res.empty())
|
||||
break;
|
||||
}
|
||||
@@ -79,7 +79,7 @@ namespace Recommendation
|
||||
}
|
||||
|
||||
TrackContainer
|
||||
RecommendationService::getSimilarTracks(const std::vector<Database::TrackId>& trackIds, std::size_t maxCount) const
|
||||
RecommendationService::findSimilarTracks(const std::vector<Database::TrackId>& trackIds, std::size_t maxCount) const
|
||||
{
|
||||
TrackContainer res;
|
||||
|
||||
@@ -91,7 +91,7 @@ namespace Recommendation
|
||||
continue;
|
||||
|
||||
const IEngine& engine {*itEngine->second};
|
||||
res = engine.getSimilarTracks(trackIds, maxCount);
|
||||
res = engine.findSimilarTracks(trackIds, maxCount);
|
||||
if (!res.empty())
|
||||
{
|
||||
LMS_LOG(RECOMMENDATION, DEBUG) << "Got " << res.size() << " similar tracks using engine '" << engineTypeToString(engineType) << "'";
|
||||
@@ -140,6 +140,8 @@ namespace Recommendation
|
||||
if (itEngine == std::cend(_engines))
|
||||
continue;
|
||||
|
||||
LMS_LOG(RECOMMENDATION, DEBUG) << "Trying engine '" << engineTypeToString(engineType) << "'";
|
||||
|
||||
const IEngine& engine {*itEngine->second};
|
||||
res = engine.getSimilarArtists(artistId, linkTypes, maxCount);
|
||||
if (!res.empty())
|
||||
|
||||
@@ -56,8 +56,8 @@ namespace Recommendation
|
||||
void load(bool forceReload, const ProgressCallback& progressCallback) override;
|
||||
void cancelLoad() override;
|
||||
|
||||
TrackContainer getSimilarTracksFromTrackList(Database::TrackListId tracklistId, std::size_t maxCount) const override;
|
||||
TrackContainer getSimilarTracks(const std::vector<Database::TrackId>& tracksId, std::size_t maxCount) const override;
|
||||
TrackContainer findSimilarTracksFromTrackList(Database::TrackListId tracklistId, std::size_t maxCount) const override;
|
||||
TrackContainer findSimilarTracks(const std::vector<Database::TrackId>& tracksId, std::size_t maxCount) const override;
|
||||
ReleaseContainer getSimilarReleases(Database::ReleaseId releaseId, std::size_t maxCount) const override;
|
||||
ArtistContainer getSimilarArtists(Database::ArtistId artistId, EnumSet<Database::TrackArtistLinkType> linkTypes, std::size_t maxCount) const override;
|
||||
|
||||
|
||||
@@ -29,40 +29,35 @@
|
||||
|
||||
namespace Recommendation {
|
||||
|
||||
std::unique_ptr<IEngine> createClustersEngine(Database::Db& db)
|
||||
using namespace Database;
|
||||
|
||||
std::unique_ptr<IEngine> createClustersEngine(Db& db)
|
||||
{
|
||||
return std::make_unique<ClusterEngine>(db);
|
||||
}
|
||||
|
||||
TrackContainer
|
||||
ClusterEngine::getSimilarTracks(const std::vector<Database::TrackId>& trackIds, std::size_t maxCount) const
|
||||
ClusterEngine::findSimilarTracks(const std::vector<TrackId>& trackIds, std::size_t maxCount) const
|
||||
{
|
||||
Database::Session& dbSession {_db.getTLSSession()};
|
||||
Session& dbSession {_db.getTLSSession()};
|
||||
|
||||
TrackContainer res;
|
||||
auto transaction {dbSession.createSharedTransaction()};
|
||||
|
||||
{
|
||||
auto transaction {dbSession.createSharedTransaction()};
|
||||
|
||||
const auto tracks {Database::Track::getSimilarTracks(dbSession, trackIds, 0, maxCount)};
|
||||
res.reserve(tracks.size());
|
||||
std::transform(std::cbegin(tracks), std::cend(tracks), std::back_inserter(res), [](const auto& track) { return track->getId(); });
|
||||
}
|
||||
|
||||
return res;
|
||||
const auto similarTrackIds {Track::findSimilarTracks(dbSession, trackIds, Range {0, maxCount})};
|
||||
return std::move(similarTrackIds.results);
|
||||
}
|
||||
|
||||
TrackContainer
|
||||
ClusterEngine::getSimilarTracksFromTrackList(Database::TrackListId tracklistId, std::size_t maxCount) const
|
||||
ClusterEngine::findSimilarTracksFromTrackList(TrackListId tracklistId, std::size_t maxCount) const
|
||||
{
|
||||
Database::Session& dbSession {_db.getTLSSession()};
|
||||
Session& dbSession {_db.getTLSSession()};
|
||||
|
||||
TrackContainer res;
|
||||
|
||||
{
|
||||
auto transaction {dbSession.createSharedTransaction()};
|
||||
|
||||
const Database::TrackList::pointer trackList {Database::TrackList::getById(dbSession, tracklistId)};
|
||||
const TrackList::pointer trackList {TrackList::find(dbSession, tracklistId)};
|
||||
if (!trackList)
|
||||
return res;
|
||||
|
||||
@@ -75,15 +70,15 @@ ClusterEngine::getSimilarTracksFromTrackList(Database::TrackListId tracklistId,
|
||||
}
|
||||
|
||||
ReleaseContainer
|
||||
ClusterEngine::getSimilarReleases(Database::ReleaseId releaseId, std::size_t maxCount) const
|
||||
ClusterEngine::getSimilarReleases(ReleaseId releaseId, std::size_t maxCount) const
|
||||
{
|
||||
Database::Session& dbSession {_db.getTLSSession()};
|
||||
Session& dbSession {_db.getTLSSession()};
|
||||
|
||||
ReleaseContainer res;
|
||||
{
|
||||
auto transaction {dbSession.createSharedTransaction()};
|
||||
|
||||
auto release {Database::Release::getById(dbSession, releaseId)};
|
||||
auto release {Release::find(dbSession, releaseId)};
|
||||
if (!release)
|
||||
return res;
|
||||
|
||||
@@ -96,24 +91,18 @@ ClusterEngine::getSimilarReleases(Database::ReleaseId releaseId, std::size_t max
|
||||
}
|
||||
|
||||
ArtistContainer
|
||||
ClusterEngine::getSimilarArtists(Database::ArtistId artistId, EnumSet<Database::TrackArtistLinkType> artistLinkTypes, std::size_t maxCount) const
|
||||
ClusterEngine::getSimilarArtists(ArtistId artistId, EnumSet<TrackArtistLinkType> artistLinkTypes, std::size_t maxCount) const
|
||||
{
|
||||
Database::Session& dbSession {_db.getTLSSession()};
|
||||
Session& dbSession {_db.getTLSSession()};
|
||||
|
||||
ResultContainer<Database::ArtistId> res;
|
||||
{
|
||||
auto transaction {dbSession.createSharedTransaction()};
|
||||
auto transaction {dbSession.createSharedTransaction()};
|
||||
|
||||
auto artist {Database::Artist::getById(dbSession, artistId)};
|
||||
if (!artist)
|
||||
return res;
|
||||
auto artist {Artist::find(dbSession, artistId)};
|
||||
if (!artist)
|
||||
return {};
|
||||
|
||||
const auto artists {artist->getSimilarArtists(artistLinkTypes, Database::Range {0, maxCount})};
|
||||
res.reserve(artists.size());
|
||||
std::transform(std::cbegin(artists), std::cend(artists), std::back_inserter(res), [](const auto& artist) { return artist->getId(); });
|
||||
}
|
||||
|
||||
return res;
|
||||
const auto similarArtistIds {artist->findSimilarArtists(artistLinkTypes, Range {0, maxCount})};
|
||||
return std::move(similarArtistIds.results);
|
||||
}
|
||||
|
||||
} // namespace Recommendation
|
||||
|
||||
@@ -38,8 +38,8 @@ namespace Recommendation
|
||||
void load(bool, const ProgressCallback&) override {}
|
||||
void requestCancelLoad() override {}
|
||||
|
||||
TrackContainer getSimilarTracksFromTrackList(Database::TrackListId tracklistId, std::size_t maxCount) const override;
|
||||
TrackContainer getSimilarTracks(const std::vector<Database::TrackId>& tracksId, std::size_t maxCount) const override;
|
||||
TrackContainer findSimilarTracksFromTrackList(Database::TrackListId tracklistId, std::size_t maxCount) const override;
|
||||
TrackContainer findSimilarTracks(const std::vector<Database::TrackId>& tracksId, std::size_t maxCount) const override;
|
||||
ReleaseContainer getSimilarReleases(Database::ReleaseId releaseId, std::size_t maxCount) const override;
|
||||
ArtistContainer getSimilarArtists(Database::ArtistId artistId, EnumSet<Database::TrackArtistLinkType> linkTypes, std::size_t maxCount) const override;
|
||||
|
||||
|
||||
@@ -36,7 +36,9 @@
|
||||
|
||||
namespace Recommendation {
|
||||
|
||||
std::unique_ptr<IEngine> createFeaturesEngine(Database::Db& db)
|
||||
using namespace Database;
|
||||
|
||||
std::unique_ptr<IEngine> createFeaturesEngine(Db& db)
|
||||
{
|
||||
return std::make_unique<FeaturesEngine>(db);
|
||||
}
|
||||
@@ -49,44 +51,13 @@ FeaturesEngine::getDefaultTrainFeatureSettings()
|
||||
{ "lowlevel.spectral_energyband_high.mean", {1}},
|
||||
{ "lowlevel.spectral_rolloff.median", {1}},
|
||||
{ "lowlevel.spectral_contrast_valleys.var", {1}},
|
||||
{ "lowlevel.erbbands.mean", {1}},
|
||||
{ "lowlevel.gfcc.mean", {1}},
|
||||
{ "lowlevel.erbbands.mean", {1}},
|
||||
{ "lowlevel.gfcc.mean", {1}},
|
||||
};
|
||||
|
||||
return defaultTrainFeatureSettings;
|
||||
}
|
||||
|
||||
static
|
||||
std::optional<FeatureValuesMap>
|
||||
getTrackFeatureValues(FeaturesEngine::FeaturesFetchFunc func, Database::TrackId trackId, const std::unordered_set<FeatureName>& featureNames)
|
||||
{
|
||||
return func(trackId, featureNames);
|
||||
}
|
||||
|
||||
static
|
||||
std::optional<FeatureValuesMap>
|
||||
getTrackFeatureValuesFromDb(Database::Session& session, Database::TrackId trackId, const std::unordered_set<FeatureName>& featureNames)
|
||||
{
|
||||
auto func = [&](Database::TrackId 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 res;
|
||||
|
||||
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)
|
||||
@@ -142,45 +113,46 @@ FeaturesEngine::loadFromTraining(const TrainSettings& trainSettings, const Progr
|
||||
|
||||
LMS_LOG(RECOMMENDATION, DEBUG) << "Features dimension = " << nbDimensions;
|
||||
|
||||
Database::Session& session {_db.getTLSSession()};
|
||||
Session& session {_db.getTLSSession()};
|
||||
|
||||
std::vector<Database::TrackId> trackIds;
|
||||
RangeResults<TrackFeaturesId> trackFeaturesIds;
|
||||
{
|
||||
auto transaction {session.createSharedTransaction()};
|
||||
|
||||
LMS_LOG(RECOMMENDATION, DEBUG) << "Getting Tracks with features...";
|
||||
trackIds = Database::Track::getAllIdsWithFeatures(session);
|
||||
LMS_LOG(RECOMMENDATION, DEBUG) << "Getting Tracks with features DONE (found " << trackIds.size() << " tracks)";
|
||||
LMS_LOG(RECOMMENDATION, DEBUG) << "Getting Track features...";
|
||||
trackFeaturesIds = TrackFeatures::find(session, Range {});
|
||||
LMS_LOG(RECOMMENDATION, DEBUG) << "Getting Track features DONE (found " << trackFeaturesIds.results.size() << " track features)";
|
||||
}
|
||||
|
||||
std::vector<SOM::InputVector> samples;
|
||||
std::vector<Database::TrackId> samplesTrackIds;
|
||||
std::vector<TrackId> samplesTrackIds;
|
||||
|
||||
samples.reserve(trackIds.size());
|
||||
samplesTrackIds.reserve(trackIds.size());
|
||||
samples.reserve(trackFeaturesIds.results.size());
|
||||
samplesTrackIds.reserve(trackFeaturesIds.results.size());
|
||||
|
||||
LMS_LOG(RECOMMENDATION, DEBUG) << "Extracting features...";
|
||||
for (Database::TrackId trackId : trackIds)
|
||||
// TODO handle errors using exceptions
|
||||
for (const TrackFeaturesId trackFeaturesId : trackFeaturesIds.results)
|
||||
{
|
||||
if (_loadCancelled)
|
||||
return;
|
||||
|
||||
std::optional<FeatureValuesMap> featureValuesMap;
|
||||
auto transaction {session.createSharedTransaction()};
|
||||
|
||||
if (_featuresFetchFunc)
|
||||
featureValuesMap = getTrackFeatureValues(_featuresFetchFunc, trackId, featureNames);
|
||||
else
|
||||
featureValuesMap = getTrackFeatureValuesFromDb(session, trackId, featureNames);
|
||||
|
||||
if (!featureValuesMap)
|
||||
TrackFeatures::pointer trackFeatures {TrackFeatures::find(session, trackFeaturesId)};
|
||||
if (!trackFeatures)
|
||||
continue;
|
||||
|
||||
std::optional<SOM::InputVector> inputVector {convertFeatureValuesMapToInputVector(*featureValuesMap, nbDimensions)};
|
||||
FeatureValuesMap featureValuesMap {trackFeatures->getFeatureValuesMap(featureNames)};
|
||||
if (featureValuesMap.empty())
|
||||
continue;
|
||||
|
||||
std::optional<SOM::InputVector> inputVector {convertFeatureValuesMapToInputVector(featureValuesMap, nbDimensions)};
|
||||
if (!inputVector)
|
||||
continue;
|
||||
|
||||
samples.emplace_back(std::move(*inputVector));
|
||||
samplesTrackIds.emplace_back(trackId);
|
||||
samplesTrackIds.emplace_back(trackFeatures->getTrack()->getId());
|
||||
}
|
||||
LMS_LOG(RECOMMENDATION, DEBUG) << "Extracting features DONE";
|
||||
|
||||
@@ -249,41 +221,41 @@ FeaturesEngine::loadFromCache(FeaturesEngineCache cache)
|
||||
}
|
||||
|
||||
TrackContainer
|
||||
FeaturesEngine::getSimilarTracksFromTrackList(Database::TrackListId trackListId, std::size_t maxCount) const
|
||||
FeaturesEngine::findSimilarTracksFromTrackList(TrackListId trackListId, std::size_t maxCount) const
|
||||
{
|
||||
const TrackContainer trackIds {[&]
|
||||
{
|
||||
TrackContainer res;
|
||||
|
||||
Database::Session& session {_db.getTLSSession()};
|
||||
Session& session {_db.getTLSSession()};
|
||||
|
||||
auto transaction {session.createSharedTransaction()};
|
||||
|
||||
const Database::TrackList::pointer trackList {Database::TrackList::getById(session, trackListId)};
|
||||
const TrackList::pointer trackList {TrackList::find(session, trackListId)};
|
||||
if (trackList)
|
||||
res = trackList->getTrackIds();
|
||||
|
||||
return res;
|
||||
}()};
|
||||
|
||||
return getSimilarTracks(trackIds, maxCount);
|
||||
return findSimilarTracks(trackIds, maxCount);
|
||||
}
|
||||
|
||||
TrackContainer
|
||||
FeaturesEngine::getSimilarTracks(const std::vector<Database::TrackId>& tracksIds, std::size_t maxCount) const
|
||||
FeaturesEngine::findSimilarTracks(const std::vector<TrackId>& tracksIds, std::size_t maxCount) const
|
||||
{
|
||||
auto similarTrackIds {getSimilarObjects(tracksIds, _trackMatrix, _trackPositions, maxCount)};
|
||||
|
||||
Database::Session& session {_db.getTLSSession()};
|
||||
Session& session {_db.getTLSSession()};
|
||||
|
||||
{
|
||||
// Report only existing ids, as tracks may have been removed a long time ago (refreshing the SOM takes some time)
|
||||
auto transaction {session.createSharedTransaction()};
|
||||
|
||||
similarTrackIds.erase(std::remove_if(std::begin(similarTrackIds), std::end(similarTrackIds),
|
||||
[&](Database::TrackId trackId)
|
||||
[&](TrackId trackId)
|
||||
{
|
||||
return !Database::Track::exists(session, trackId);
|
||||
return !Track::exists(session, trackId);
|
||||
}), std::end(similarTrackIds));
|
||||
}
|
||||
|
||||
@@ -291,11 +263,11 @@ FeaturesEngine::getSimilarTracks(const std::vector<Database::TrackId>& tracksIds
|
||||
}
|
||||
|
||||
ReleaseContainer
|
||||
FeaturesEngine::getSimilarReleases(Database::ReleaseId releaseId, std::size_t maxCount) const
|
||||
FeaturesEngine::getSimilarReleases(ReleaseId releaseId, std::size_t maxCount) const
|
||||
{
|
||||
auto similarReleaseIds {getSimilarObjects({releaseId}, _releaseMatrix, _releasePositions, maxCount)};
|
||||
|
||||
Database::Session& session {_db.getTLSSession()};
|
||||
Session& session {_db.getTLSSession()};
|
||||
|
||||
if (!similarReleaseIds.empty())
|
||||
{
|
||||
@@ -303,9 +275,9 @@ FeaturesEngine::getSimilarReleases(Database::ReleaseId releaseId, std::size_t ma
|
||||
auto transaction {session.createSharedTransaction()};
|
||||
|
||||
similarReleaseIds.erase(std::remove_if(std::begin(similarReleaseIds), std::end(similarReleaseIds),
|
||||
[&](Database::ReleaseId releaseId)
|
||||
[&](ReleaseId releaseId)
|
||||
{
|
||||
return !Database::Release::exists(session, releaseId);
|
||||
return !Release::exists(session, releaseId);
|
||||
}), std::end(similarReleaseIds));
|
||||
}
|
||||
|
||||
@@ -313,9 +285,9 @@ FeaturesEngine::getSimilarReleases(Database::ReleaseId releaseId, std::size_t ma
|
||||
}
|
||||
|
||||
ArtistContainer
|
||||
FeaturesEngine::getSimilarArtists(Database::ArtistId artistId, EnumSet<Database::TrackArtistLinkType> linkTypes, std::size_t maxCount) const
|
||||
FeaturesEngine::getSimilarArtists(ArtistId artistId, EnumSet<TrackArtistLinkType> linkTypes, std::size_t maxCount) const
|
||||
{
|
||||
auto getSimilarArtistIdsForLinkType {[&] (Database::TrackArtistLinkType linkType)
|
||||
auto getSimilarArtistIdsForLinkType {[&] (TrackArtistLinkType linkType)
|
||||
{
|
||||
ArtistContainer similarArtistIds;
|
||||
|
||||
@@ -328,9 +300,9 @@ FeaturesEngine::getSimilarArtists(Database::ArtistId artistId, EnumSet<Database:
|
||||
return getSimilarObjects({artistId}, itArtists->second, _artistPositions, maxCount);
|
||||
}};
|
||||
|
||||
std::unordered_set<Database::ArtistId> similarArtistIds;
|
||||
std::unordered_set<ArtistId> similarArtistIds;
|
||||
|
||||
for (Database::TrackArtistLinkType linkType : linkTypes)
|
||||
for (TrackArtistLinkType linkType : linkTypes)
|
||||
{
|
||||
const auto similarArtistIdsForLinkType {getSimilarArtistIdsForLinkType(linkType)};
|
||||
similarArtistIds.insert(std::begin(similarArtistIdsForLinkType), std::end(similarArtistIdsForLinkType));
|
||||
@@ -338,15 +310,15 @@ FeaturesEngine::getSimilarArtists(Database::ArtistId artistId, EnumSet<Database:
|
||||
|
||||
ArtistContainer res(std::cbegin(similarArtistIds), std::cend(similarArtistIds));
|
||||
|
||||
Database::Session& session {_db.getTLSSession()};
|
||||
Session& session {_db.getTLSSession()};
|
||||
{
|
||||
// Report only existing ids
|
||||
auto transaction {session.createSharedTransaction()};
|
||||
|
||||
res.erase(std::remove_if(std::begin(res), std::end(res),
|
||||
[&](Database::ArtistId artistId)
|
||||
[&](ArtistId artistId)
|
||||
{
|
||||
return !Database::Artist::exists(session, artistId);
|
||||
return !Artist::exists(session, artistId);
|
||||
}), std::end(res));
|
||||
}
|
||||
|
||||
@@ -406,7 +378,7 @@ FeaturesEngine::load(const SOM::Network& network, const TrackPositions& trackPos
|
||||
|
||||
LMS_LOG(RECOMMENDATION, DEBUG) << "Constructing maps...";
|
||||
|
||||
Database::Session& session {_db.getTLSSession()};
|
||||
Session& session {_db.getTLSSession()};
|
||||
|
||||
for (const auto& [trackId, positions] : trackPositions)
|
||||
{
|
||||
@@ -415,7 +387,7 @@ FeaturesEngine::load(const SOM::Network& network, const TrackPositions& trackPos
|
||||
|
||||
auto transaction {session.createSharedTransaction()};
|
||||
|
||||
const Track::pointer track {Database::Track::getById(session, trackId)};
|
||||
const Track::pointer track {Track::find(session, trackId)};
|
||||
if (!track)
|
||||
continue;
|
||||
|
||||
|
||||
@@ -52,19 +52,14 @@ class FeaturesEngine : public IEngine
|
||||
FeaturesEngine& operator=(const FeaturesEngine&) = delete;
|
||||
FeaturesEngine& operator=(FeaturesEngine&&) = delete;
|
||||
|
||||
using FeaturesFetchFunc = std::function<std::optional<std::unordered_map<std::string, std::vector<double>>>(Database::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 some cached data
|
||||
static void setFeaturesFetchFunc(FeaturesFetchFunc func) { _featuresFetchFunc = func; }
|
||||
|
||||
static const FeatureSettingsMap& getDefaultTrainFeatureSettings();
|
||||
|
||||
private:
|
||||
void load(bool forceReload, const ProgressCallback& progressCallback) override;
|
||||
void requestCancelLoad() override;
|
||||
|
||||
TrackContainer getSimilarTracksFromTrackList(Database::TrackListId tracklistId, std::size_t maxCount) const override;
|
||||
TrackContainer getSimilarTracks(const std::vector<Database::TrackId>& tracksId, std::size_t maxCount) const override;
|
||||
TrackContainer findSimilarTracksFromTrackList(Database::TrackListId tracklistId, std::size_t maxCount) const override;
|
||||
TrackContainer findSimilarTracks(const std::vector<Database::TrackId>& tracksId, std::size_t maxCount) const override;
|
||||
ReleaseContainer getSimilarReleases(Database::ReleaseId releaseId, std::size_t maxCount) const override;
|
||||
ArtistContainer getSimilarArtists(Database::ArtistId artistId, EnumSet<Database::TrackArtistLinkType> linkTypes, std::size_t maxCount) const override;
|
||||
|
||||
@@ -121,8 +116,6 @@ class FeaturesEngine : public IEngine
|
||||
|
||||
TrackPositions _trackPositions;
|
||||
TrackMatrix _trackMatrix;
|
||||
|
||||
static inline FeaturesFetchFunc _featuresFetchFunc;
|
||||
};
|
||||
|
||||
template <typename IdType>
|
||||
|
||||
@@ -21,9 +21,8 @@
|
||||
|
||||
#include <filesystem>
|
||||
#include <unordered_map>
|
||||
#include <unordered_set>
|
||||
|
||||
#include "services/database/Types.hpp"
|
||||
#include "services/database/TrackId.hpp"
|
||||
#include "som/Network.hpp"
|
||||
|
||||
namespace Recommendation {
|
||||
|
||||
+4
-3
@@ -20,8 +20,9 @@
|
||||
#pragma once
|
||||
|
||||
#include <memory>
|
||||
#include <string_view>
|
||||
#include "utils/EnumSet.hpp"
|
||||
#include "services/database/TrackListId.hpp"
|
||||
#include "services/database/Types.hpp"
|
||||
#include "services/recommendation/Types.hpp"
|
||||
|
||||
namespace Database
|
||||
@@ -39,8 +40,8 @@ namespace Recommendation
|
||||
virtual void load(bool forceReload, const ProgressCallback& progressCallback = {}) = 0;
|
||||
virtual void cancelLoad() = 0; // wait for cancel done
|
||||
|
||||
virtual TrackContainer getSimilarTracksFromTrackList(Database::TrackListId tracklistId, std::size_t maxCount) const = 0;
|
||||
virtual TrackContainer getSimilarTracks(const std::vector<Database::TrackId>& tracksId, std::size_t maxCount) const = 0;
|
||||
virtual TrackContainer findSimilarTracksFromTrackList(Database::TrackListId tracklistId, std::size_t maxCount) const = 0;
|
||||
virtual TrackContainer findSimilarTracks(const std::vector<Database::TrackId>& tracksId, std::size_t maxCount) const = 0;
|
||||
virtual ReleaseContainer getSimilarReleases(Database::ReleaseId releaseId, std::size_t maxCount) const = 0;
|
||||
virtual ArtistContainer getSimilarArtists(Database::ArtistId artistId, EnumSet<Database::TrackArtistLinkType> linkTypes, std::size_t maxCount) const = 0;
|
||||
};
|
||||
|
||||
@@ -1,7 +1,9 @@
|
||||
#pragma once
|
||||
|
||||
#include <functional>
|
||||
#include "services/database/Types.hpp"
|
||||
#include "services/database/ArtistId.hpp"
|
||||
#include "services/database/ReleaseId.hpp"
|
||||
#include "services/database/TrackId.hpp"
|
||||
|
||||
namespace Recommendation
|
||||
{
|
||||
|
||||
Reference in New Issue
Block a user