From d1bdaf62033ebd126af1f9b5d7e1630af4f8e155 Mon Sep 17 00:00:00 2001 From: emeric Date: Sat, 15 Feb 2020 15:52:39 +0100 Subject: [PATCH] Started to restore clusters based recommendation engine --- README.md | 2 +- approot/messages.xml | 6 +- approot/messages_fr.xml | 6 +- src/libs/database/impl/Artist.cpp | 15 +++ src/libs/database/impl/Release.cpp | 15 +++ src/libs/database/impl/Session.cpp | 2 +- src/libs/database/impl/Track.cpp | 15 ++- src/libs/database/include/database/Artist.hpp | 1 + .../database/include/database/Release.hpp | 1 + .../include/database/ScanSettings.hpp | 10 +- src/libs/database/include/database/Track.hpp | 6 +- src/libs/recommendation/CMakeLists.txt | 1 + .../recommendation/impl/ClassifierCreator.cpp | 9 +- src/libs/recommendation/impl/Engine.cpp | 105 ++++++++++-------- src/libs/recommendation/impl/Engine.hpp | 13 ++- .../cluster/SimilarityClusterSearcher.hpp | 40 ------- .../ClustersClassifier.cpp} | 66 +++++++++-- .../impl/clusters/ClustersClassifier.hpp | 56 ++++++++++ .../ClustersClassifierCreator.hpp | 9 +- .../FeaturesClassifierCreator.hpp | 9 +- .../{Classifier.hpp => IClassifier.hpp} | 12 +- .../include/recommendation/IEngine.hpp | 6 +- src/lms/main.cpp | 5 +- src/lms/ui/admin/DatabaseSettingsView.cpp | 46 ++++---- src/test/database/DatabaseTest.cpp | 36 ++++++ .../LmsRecommendationFeatures.cpp | 11 +- 26 files changed, 326 insertions(+), 177 deletions(-) delete mode 100644 src/libs/recommendation/impl/cluster/SimilarityClusterSearcher.hpp rename src/libs/recommendation/impl/{cluster/SimilarityClusterSearcher.cpp => clusters/ClustersClassifier.cpp} (57%) create mode 100644 src/libs/recommendation/impl/clusters/ClustersClassifier.hpp rename src/libs/recommendation/include/recommendation/{Classifier.hpp => IClassifier.hpp} (85%) diff --git a/README.md b/README.md index a083a22c..25b903a4 100644 --- a/README.md +++ b/README.md @@ -17,7 +17,7 @@ A [demo](http://lms.demo.poupon.io) instance is available, with the following li * Audio transcode for maximum interoperability and low bandwith requirements * Persistent play queue across sessions * Subsonic API -* Album artist +* Compilation support * Multi-value tags: artists, genres, ... * Custom tags (ex: _mood_, _genre_, _albummood_, _albumgrouping_, ...) * MusicBrainzID support to handle duplicated artist and release names diff --git a/approot/messages.xml b/approot/messages.xml index 17d5ded8..20677bc8 100644 --- a/approot/messages.xml +++ b/approot/messages.xml @@ -39,9 +39,9 @@ Monthly Never Media root directory -Recommendation engine -Tags based -Audio analysis based +Recommendation engine +Tags based +Audio analysis based Scan complete: {1} total files, {2} additions, {3} updates, {4} deletions, {5} duplicates, {6} errors Scan launched! Scan options diff --git a/approot/messages_fr.xml b/approot/messages_fr.xml index 2d61dc1d..934cd400 100644 --- a/approot/messages_fr.xml +++ b/approot/messages_fr.xml @@ -39,9 +39,9 @@ Tous les mois Jamais Dossier racine des fichiers de musique -Moteur de recommandation -Basé sur les tags -Basé sur l'analyse audio +Moteur de recommandation +Basé sur les tags +Basé sur l'analyse audio Scan terminé : {1} fichiers, {2} ajouts, {3} mises à jour, {4} suppressions, {5} duplicatas, {6} erreurs Scan lancé ! Options diff --git a/src/libs/database/impl/Artist.cpp b/src/libs/database/impl/Artist.cpp index 5d17df1a..b5c6d221 100644 --- a/src/libs/database/impl/Artist.cpp +++ b/src/libs/database/impl/Artist.cpp @@ -104,6 +104,21 @@ Artist::getAllOrphans(Session& session) return std::vector(res.begin(), res.end()); } +std::vector +Artist::getAllIdsWithClusters(Session& session, std::optional limit) +{ + session.checkSharedLocked(); + + Wt::Dbo::collection res = session.getDboSession().query + ("SELECT DISTINCT a.id FROM artist a" + " INNER JOIN track t ON t.id = t_a_l.track_id INNER JOIN track_artist_link t_a_l ON t_a_l.artist_id = a.id" + " INNER JOIN track_cluster t_c ON t_c.track_id = t.id") + .limit(limit ? static_cast(*limit) : -1); + + return std::vector(res.begin(), res.end()); +} + + static Wt::Dbo::Query getQuery(Session& session, diff --git a/src/libs/database/impl/Release.cpp b/src/libs/database/impl/Release.cpp index 343a7ce1..7fdd639a 100644 --- a/src/libs/database/impl/Release.cpp +++ b/src/libs/database/impl/Release.cpp @@ -252,6 +252,21 @@ Release::getByFilter(Session& session, return res; } +std::vector +Release::getAllIdsWithClusters(Session& session, std::optional limit) +{ + session.checkSharedLocked(); + + Wt::Dbo::collection res = session.getDboSession().query + ("SELECT DISTINCT r.id FROM release r" + " INNER JOIN track t ON t.release_id = r.id" + " INNER JOIN track_cluster t_c ON t_c.track_id = t.id") + .limit(limit ? static_cast(*limit) : -1); + + return std::vector(res.begin(), res.end()); +} + + std::optional Release::getTotalTrackNumber(void) const { diff --git a/src/libs/database/impl/Session.cpp b/src/libs/database/impl/Session.cpp index 36750ec7..081fe8a8 100644 --- a/src/libs/database/impl/Session.cpp +++ b/src/libs/database/impl/Session.cpp @@ -118,7 +118,7 @@ Session::doDatabaseMigrationIfNeeded() { _session.execute("DROP TABLE similarity_settings"); _session.execute("DROP TABLE similarity_settings_feature"); - _session.execute("ALTER TABLE scan_settings ADD similarity_engine_type INTEGER NOT NULL DEFAULT(" + std::to_string(static_cast(ScanSettings::SimilarityEngineType::Clusters)) + ")"); + _session.execute("ALTER TABLE scan_settings ADD similarity_engine_type INTEGER NOT NULL DEFAULT(" + std::to_string(static_cast(ScanSettings::RecommendationEngineType::Clusters)) + ")"); } else if (version == 8) { diff --git a/src/libs/database/impl/Track.cpp b/src/libs/database/impl/Track.cpp index a84af8e1..2acff171 100644 --- a/src/libs/database/impl/Track.cpp +++ b/src/libs/database/impl/Track.cpp @@ -163,6 +163,19 @@ Track::getAllIdsWithFeatures(Session& session, std::optional limit) return std::vector(res.begin(), res.end()); } +std::vector +Track::getAllIdsWithClusters(Session& session, std::optional limit) +{ + session.checkSharedLocked(); + + Wt::Dbo::collection res = session.getDboSession().query + ("SELECT DISTINCT t.id FROM track t" + " INNER JOIN track_cluster t_c ON t_c.track_id = t.id") + .limit(limit ? static_cast(*limit) : -1); + + return std::vector(res.begin(), res.end()); +} + std::vector Track::getClusters(void) const { @@ -263,7 +276,7 @@ Track::getByFilter(Session& session, std::vector Track::getSimilarTracks(Session& session, - const std::set& tracks, + const std::unordered_set& tracks, std::optional offset, std::optional size) { diff --git a/src/libs/database/include/database/Artist.hpp b/src/libs/database/include/database/Artist.hpp index 01d261b3..f53082c1 100644 --- a/src/libs/database/include/database/Artist.hpp +++ b/src/libs/database/include/database/Artist.hpp @@ -68,6 +68,7 @@ class Artist : public Wt::Dbo::Dbo static std::vector getAllIds(Session& session); static std::vector getAllOrphans(Session& session); // No track related static std::vector getLastAdded(Session& session, Wt::WDateTime after, std::optional size = {}); + static std::vector getAllIdsWithClusters(Session& session, std::optional limit = {}); // Accessors const std::string& getName(void) const { return _name; } diff --git a/src/libs/database/include/database/Release.hpp b/src/libs/database/include/database/Release.hpp index c0a638f3..695bcdc5 100644 --- a/src/libs/database/include/database/Release.hpp +++ b/src/libs/database/include/database/Release.hpp @@ -66,6 +66,7 @@ class Release : public Wt::Dbo::Dbo std::optional offset, std::optional size, bool& moreExpected); + static std::vector getAllIdsWithClusters(Session& session, std::optional limit = {}); std::vector> getTracks(const std::set& clusters = std::set()) const; std::size_t getTracksCount() const; diff --git a/src/libs/database/include/database/ScanSettings.hpp b/src/libs/database/include/database/ScanSettings.hpp index 0ce37a49..7151df23 100644 --- a/src/libs/database/include/database/ScanSettings.hpp +++ b/src/libs/database/include/database/ScanSettings.hpp @@ -43,7 +43,7 @@ class ScanSettings : public Wt::Dbo::Dbo }; // Do not modify values (just add) - enum class SimilarityEngineType + enum class RecommendationEngineType { Clusters = 0, Features, @@ -60,7 +60,7 @@ class ScanSettings : public Wt::Dbo::Dbo UpdatePeriod getUpdatePeriod() const { return _updatePeriod; } std::vector> getClusterTypes() const; std::set getAudioFileExtensions() const; - SimilarityEngineType getSimilarityEngineType() const { return _similarityEngineType; } + RecommendationEngineType getRecommendationEngineType() const { return _recommendationEngineType; } // Setters void addAudioFileExtension(const std::filesystem::path& ext); @@ -68,7 +68,7 @@ class ScanSettings : public Wt::Dbo::Dbo void setUpdateStartTime(Wt::WTime t) { _startTime = t; } void setUpdatePeriod(UpdatePeriod p) { _updatePeriod = p; } void setClusterTypes(Session& session, const std::set& clusterTypeNames); - void setSimilarityEngineType(SimilarityEngineType type) { _similarityEngineType = type; } + void setRecommendationEngineType(RecommendationEngineType type) { _recommendationEngineType = type; } void incScanVersion(); template @@ -79,7 +79,7 @@ class ScanSettings : public Wt::Dbo::Dbo Wt::Dbo::field(a, _startTime, "start_time"); Wt::Dbo::field(a, _updatePeriod, "update_period"); Wt::Dbo::field(a, _audioFileExtensions, "audio_file_extensions"); - Wt::Dbo::field(a, _similarityEngineType,"similarity_engine_type"); + Wt::Dbo::field(a, _recommendationEngineType,"similarity_engine_type"); Wt::Dbo::hasMany(a, _clusterTypes, Wt::Dbo::ManyToOne, "scan_settings"); } @@ -89,7 +89,7 @@ class ScanSettings : public Wt::Dbo::Dbo std::string _mediaDirectory; Wt::WTime _startTime = Wt::WTime {0,0,0}; UpdatePeriod _updatePeriod {UpdatePeriod::Never}; - SimilarityEngineType _similarityEngineType {SimilarityEngineType::Clusters}; + RecommendationEngineType _recommendationEngineType {RecommendationEngineType::Clusters}; std::string _audioFileExtensions {".alac .mp3 .ogg .oga .aac .m4a .m4b .flac .wav .wma .aif .aiff .ape .mpc .shn .opus"}; Wt::Dbo::collection> _clusterTypes; }; diff --git a/src/libs/database/include/database/Track.hpp b/src/libs/database/include/database/Track.hpp index d3bcd824..24113f21 100644 --- a/src/libs/database/include/database/Track.hpp +++ b/src/libs/database/include/database/Track.hpp @@ -22,8 +22,9 @@ #include #include #include -#include #include +#include +#include #include #include @@ -58,7 +59,7 @@ class Track : public Wt::Dbo::Dbo static pointer getById(Session& session, IdType id); static pointer getByMBID(Session& session, const UUID& MBID); static std::vector getSimilarTracks(Session& session, - const std::set& trackIds, + const std::unordered_set& trackIds, std::optional offset = {}, std::optional size = {}); static std::vector getByClusters(Session& session, @@ -78,6 +79,7 @@ class Track : public Wt::Dbo::Dbo static std::vector getLastAdded(Session& session, const Wt::WDateTime& after, std::optional size = 1); static std::vector getAllWithMBIDAndMissingFeatures(Session& session); static std::vector getAllIdsWithFeatures(Session& session, std::optional limit = {}); + static std::vector getAllIdsWithClusters(Session& session, std::optional limit = {}); // Create utility static pointer create(Session& session, const std::filesystem::path& p); diff --git a/src/libs/recommendation/CMakeLists.txt b/src/libs/recommendation/CMakeLists.txt index b2161dc2..260f137b 100644 --- a/src/libs/recommendation/CMakeLists.txt +++ b/src/libs/recommendation/CMakeLists.txt @@ -1,5 +1,6 @@ add_library(lmsrecommendation SHARED + impl/clusters/ClustersClassifier.cpp impl/Engine.cpp impl/ClassifierCreator.cpp ) diff --git a/src/libs/recommendation/impl/ClassifierCreator.cpp b/src/libs/recommendation/impl/ClassifierCreator.cpp index 865dd1c9..90dd651b 100644 --- a/src/libs/recommendation/impl/ClassifierCreator.cpp +++ b/src/libs/recommendation/impl/ClassifierCreator.cpp @@ -17,22 +17,17 @@ * along with LMS. If not, see . */ -#include "recommendation/ClustersClassifierCreator.hpp" #include "recommendation/FeaturesClassifierCreator.hpp" -#include "recommendation/Classifier.hpp" +#include "recommendation/IClassifier.hpp" namespace Recommendation { - std::unique_ptr createClustersClassifier() + std::unique_ptr createFeaturesClassifier() { return {}; } - std::unique_ptr createFeaturesClassifier() - { - return {}; - } } diff --git a/src/libs/recommendation/impl/Engine.cpp b/src/libs/recommendation/impl/Engine.cpp index 384f6862..a5cf2a63 100644 --- a/src/libs/recommendation/impl/Engine.cpp +++ b/src/libs/recommendation/impl/Engine.cpp @@ -19,30 +19,49 @@ #include "Engine.hpp" -//#include "features/SimilarityFeaturesScannerAddon.hpp" -//#include "cluster/SimilarityClusterSearcher.hpp" +#include "recommendation/ClustersClassifierCreator.hpp" +#include "recommendation/FeaturesClassifierCreator.hpp" #include "database/ScanSettings.hpp" +#include "database/Session.hpp" #include "database/TrackList.hpp" namespace Recommendation { std::unique_ptr -createEngine() +createEngine(Database::Session& session) { - return std::make_unique(); + return std::make_unique(session); +} + +Engine::Engine(Database::Session& session) +{ + reloadSettings(session); } void -Engine::clearClassifiers() +Engine::reloadSettings(Database::Session& session) { - _classifiers.clear(); -} + using namespace Database; -void -Engine::addClassifier(std::unique_ptr classifier, unsigned priority) -{ - _classifiers.emplace(priority, std::move(classifier)); + const ScanSettings::RecommendationEngineType engineType {[&]() + { + auto transaction {session.createSharedTransaction()}; + return ScanSettings::get(session)->getRecommendationEngineType(); + }()}; + + clearClassifiers(); + + switch (engineType) + { + case ScanSettings::RecommendationEngineType::Features: +// _classifiers.emplace(0, createFeaturesClassifier()); // higher priority +// [[fallthrough]]; + + case ScanSettings::RecommendationEngineType::Clusters: + _classifiers.emplace(1, createClustersClassifier(session)); // lower priority + break; + } } std::vector @@ -79,60 +98,52 @@ Engine::getSimilarTracksFromTrackList(Database::Session& /*session*/, Database:: } std::vector -Engine::getSimilarTracks(Database::Session& /*dbSession*/, const std::unordered_set& /*trackIds*/, std::size_t /*maxCount*/) +Engine::getSimilarTracks(Database::Session& dbSession, const std::unordered_set& trackIds, std::size_t maxCount) { -#if 0 - auto engineType {getEngineType(dbSession)}; - auto somSearcher {_somAddon.getSearcher()}; - - if (engineType == Database::ScanSettings::SimilarityEngineType::Features - && somSearcher - && std::any_of(std::cbegin(trackIds), std::cend(trackIds), [&](Database::IdType trackId) { return somSearcher->isTrackClassified(trackId); } )) + for (const auto& [priority, classifier] : _classifiers) { - return somSearcher->getSimilarTracks(trackIds, maxCount); + if (std::any_of(std::cbegin(trackIds), std::cend(trackIds), [&](Database::IdType trackId) { return classifier->isTrackClassified(trackId); } )) + return classifier->getSimilarTracks(dbSession, trackIds, maxCount); } - else - return ClusterEngine::getSimilarTracks(dbSession, trackIds, maxCount); -#endif + return {}; } std::vector -Engine::getSimilarReleases(Database::Session& /*dbSession*/, Database::IdType /*releaseId*/, std::size_t /*maxCount*/) +Engine::getSimilarReleases(Database::Session& dbSession, Database::IdType releaseId, std::size_t maxCount) { -#if 0 - auto engineType {getEngineType(dbSession)}; - auto somSearcher {_somAddon.getSearcher()}; - - if (engineType == Database::ScanSettings::SimilarityEngineType::Features - && somSearcher - && somSearcher->isReleaseClassified(releaseId)) + for (const auto& [priority, classifier] : _classifiers) { - return somSearcher->getSimilarReleases(releaseId, maxCount); + if (classifier->isReleaseClassified(releaseId)) + return classifier->getSimilarReleases(dbSession, releaseId, maxCount); } - else - return ClusterEngine::getSimilarReleases(dbSession, releaseId, maxCount); -#endif + return {}; } std::vector -Engine::getSimilarArtists(Database::Session& /*dbSession*/, Database::IdType /*artistId*/, std::size_t /*maxCount*/) +Engine::getSimilarArtists(Database::Session& dbSession, Database::IdType artistId, std::size_t maxCount) { -#if 0 - auto engineType {getEngineType(dbSession)}; - auto somSearcher {_somAddon.getSearcher()}; - - if (engineType == Database::ScanSettings::SimilarityEngineType::Features - && somSearcher - && somSearcher->isArtistClassified(artistId)) + for (const auto& [priority, classifier] : _classifiers) { - return somSearcher->getSimilarArtists(artistId, maxCount); + if (classifier->isArtistClassified(artistId)) + return classifier->getSimilarArtists(dbSession, artistId, maxCount); } - else - return ClusterEngine::getSimilarArtists(dbSession, artistId, maxCount); -#endif + return {}; } +void +Engine::clearClassifiers() +{ + _classifiers.clear(); +} + +void +Engine::addClassifier(std::unique_ptr classifier, unsigned priority) +{ + _classifiers.emplace(priority, std::move(classifier)); +} + + } // ns Similarity diff --git a/src/libs/recommendation/impl/Engine.hpp b/src/libs/recommendation/impl/Engine.hpp index cbf66a2b..2fdace41 100644 --- a/src/libs/recommendation/impl/Engine.hpp +++ b/src/libs/recommendation/impl/Engine.hpp @@ -22,7 +22,7 @@ #include #include "recommendation/IEngine.hpp" -#include "recommendation/Classifier.hpp" +#include "recommendation/IClassifier.hpp" namespace Database { @@ -34,17 +34,22 @@ namespace Recommendation class Engine : public IEngine { public: + Engine(Database::Session& session); + + private: + + void reloadSettings(Database::Session& session) override; + // Closest results first std::vector getSimilarTracksFromTrackList(Database::Session& session, Database::IdType tracklistId, std::size_t maxCount) override; std::vector getSimilarTracks(Database::Session& session, const std::unordered_set& tracksId, std::size_t maxCount) override; std::vector getSimilarReleases(Database::Session& session, Database::IdType releaseId, std::size_t maxCount) override; std::vector getSimilarArtists(Database::Session& session, Database::IdType artistId, std::size_t maxCount) override; - private: void clearClassifiers(); - void addClassifier(std::unique_ptr classifier, unsigned priority); + void addClassifier(std::unique_ptr classifier, unsigned priority); - std::map> _classifiers; + std::map> _classifiers; }; } // ns Recommendation diff --git a/src/libs/recommendation/impl/cluster/SimilarityClusterSearcher.hpp b/src/libs/recommendation/impl/cluster/SimilarityClusterSearcher.hpp deleted file mode 100644 index 90ef0db8..00000000 --- a/src/libs/recommendation/impl/cluster/SimilarityClusterSearcher.hpp +++ /dev/null @@ -1,40 +0,0 @@ -/* - * Copyright (C) 2018 Emeric Poupon - * - * This file is part of LMS. - * - * LMS is free software: you can redistribute it and/or modify - * it under the terms of the GNU General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * LMS is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU General Public License for more details. - * - * You should have received a copy of the GNU General Public License - * along with LMS. If not, see . - */ - -#pragma once - -#include - -#include "database/Types.hpp" - -namespace Database { - class Session; -} - -namespace Similarity { - -namespace ClusterSearcher -{ - std::vector getSimilarTracksFromTrackList(Database::Session& session, Database::IdType tracklistId, std::size_t maxCount); - std::vector getSimilarTracks(Database::Session& session, const std::set& tracksId, std::size_t maxCount); - std::vector getSimilarReleases(Database::Session& session, Database::IdType releaseId, std::size_t maxCount); - std::vector getSimilarArtists(Database::Session& session, Database::IdType artistId, std::size_t maxCount); -}; - -} // namespace Similarity diff --git a/src/libs/recommendation/impl/cluster/SimilarityClusterSearcher.cpp b/src/libs/recommendation/impl/clusters/ClustersClassifier.cpp similarity index 57% rename from src/libs/recommendation/impl/cluster/SimilarityClusterSearcher.cpp rename to src/libs/recommendation/impl/clusters/ClustersClassifier.cpp index 7b24b5e4..f0aee487 100644 --- a/src/libs/recommendation/impl/cluster/SimilarityClusterSearcher.cpp +++ b/src/libs/recommendation/impl/clusters/ClustersClassifier.cpp @@ -17,7 +17,7 @@ * along with LMS. If not, see . */ -#include "SimilarityClusterSearcher.hpp" +#include "ClustersClassifier.hpp" #include "database/Artist.hpp" #include "database/Cluster.hpp" @@ -26,11 +26,60 @@ #include "database/Track.hpp" #include "database/TrackList.hpp" -namespace Similarity { -namespace ClusterSearcher { +namespace Recommendation { + +std::unique_ptr createClustersClassifier(Database::Session& session) +{ + return std::make_unique(session); +} + +ClusterClassifier::ClusterClassifier(Database::Session& session) +{ + classify(session); +} + +void +ClusterClassifier::classify(Database::Session& session) +{ + auto transaction {session.createSharedTransaction()}; + + { + std::vector trackIds {Database::Track::getAllIdsWithClusters(session)}; + _classifiedTracks = std::unordered_set(std::cbegin(trackIds), std::cend(trackIds)); + } + + { + std::vector releaseIds {Database::Release::getAllIdsWithClusters(session)}; + _classifiedReleases = std::unordered_set(std::cbegin(releaseIds), std::cend(releaseIds)); + } + + { + std::vector artistIds {Database::Artist::getAllIdsWithClusters(session)}; + _classifiedArtists = std::unordered_set(std::cbegin(artistIds), std::cend(artistIds)); + } +} + +bool +ClusterClassifier::isTrackClassified(Database::IdType trackId) const +{ + return _classifiedTracks.find(trackId) != std::cend(_classifiedTracks); +} + +bool +ClusterClassifier::isReleaseClassified(Database::IdType releaseId) const +{ + return _classifiedReleases.find(releaseId) != std::cend(_classifiedReleases); +} + +bool +ClusterClassifier::isArtistClassified(Database::IdType artistId) const +{ + return _classifiedArtists.find(artistId) != std::cend(_classifiedArtists); +} + std::vector -getSimilarTracks(Database::Session& dbSession, const std::set& trackIds, std::size_t maxCount) +ClusterClassifier::getSimilarTracks(Database::Session& dbSession, const std::unordered_set& trackIds, std::size_t maxCount) const { auto transaction {dbSession.createSharedTransaction()}; @@ -43,7 +92,7 @@ getSimilarTracks(Database::Session& dbSession, const std::set& } std::vector -getSimilarTracksFromTrackList(Database::Session& session, Database::IdType tracklistId, std::size_t maxCount) +ClusterClassifier::getSimilarTracksFromTrackList(Database::Session& session, Database::IdType tracklistId, std::size_t maxCount) const { std::vector res; @@ -62,7 +111,7 @@ getSimilarTracksFromTrackList(Database::Session& session, Database::IdType track } std::vector -getSimilarReleases(Database::Session& dbSession, Database::IdType releaseId, std::size_t maxCount) +ClusterClassifier::getSimilarReleases(Database::Session& dbSession, Database::IdType releaseId, std::size_t maxCount) const { std::vector res; @@ -80,7 +129,7 @@ getSimilarReleases(Database::Session& dbSession, Database::IdType releaseId, std } std::vector -getSimilarArtists(Database::Session& dbSession, Database::IdType artistId, std::size_t maxCount) +ClusterClassifier::getSimilarArtists(Database::Session& dbSession, Database::IdType artistId, std::size_t maxCount) const { std::vector res; @@ -97,5 +146,4 @@ getSimilarArtists(Database::Session& dbSession, Database::IdType artistId, std:: return res; } -} // namespace ClusterSearcher -} // namespace Similarity +} // namespace Recommendation diff --git a/src/libs/recommendation/impl/clusters/ClustersClassifier.hpp b/src/libs/recommendation/impl/clusters/ClustersClassifier.hpp new file mode 100644 index 00000000..2f44205e --- /dev/null +++ b/src/libs/recommendation/impl/clusters/ClustersClassifier.hpp @@ -0,0 +1,56 @@ +/* + * Copyright (C) 2018 Emeric Poupon + * + * This file is part of LMS. + * + * LMS is free software: you can redistribute it and/or modify + * it under the terms of the GNU General Public License as published by + * the Free Software Foundation, either version 3 of the License, or + * (at your option) any later version. + * + * LMS is distributed in the hope that it will be useful, + * but WITHOUT ANY WARRANTY; without even the implied warranty of + * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the + * GNU General Public License for more details. + * + * You should have received a copy of the GNU General Public License + * along with LMS. If not, see . + */ + +#pragma once + +#include "recommendation/IClassifier.hpp" + + +namespace Recommendation +{ + + class ClusterClassifier : public IClassifier + { + public: + ClusterClassifier(Database::Session& session); + ClusterClassifier(const ClusterClassifier&) = delete; + ClusterClassifier(ClusterClassifier&&) = delete; + ClusterClassifier& operator=(const ClusterClassifier&) = delete; + ClusterClassifier& operator=(ClusterClassifier&&) = delete; + + private: + + bool isTrackClassified(Database::IdType trackId) const override; + bool isReleaseClassified(Database::IdType releaseId) const override; + bool isArtistClassified(Database::IdType artistId) const override; + + std::vector getSimilarTracksFromTrackList(Database::Session& session, Database::IdType tracklistId, std::size_t maxCount) const override; + std::vector getSimilarTracks(Database::Session& session, const std::unordered_set& tracksId, std::size_t maxCount) const override; + std::vector getSimilarReleases(Database::Session& session, Database::IdType releaseId, std::size_t maxCount) const override; + std::vector getSimilarArtists(Database::Session& session, Database::IdType artistId, std::size_t maxCount) const override; + + void classify(Database::Session& session); + + std::unordered_set _classifiedArtists; + std::unordered_set _classifiedReleases; + std::unordered_set _classifiedTracks; +}; + +} // namespace Recommendation + diff --git a/src/libs/recommendation/include/recommendation/ClustersClassifierCreator.hpp b/src/libs/recommendation/include/recommendation/ClustersClassifierCreator.hpp index 71e92ddc..3f4dfbe8 100644 --- a/src/libs/recommendation/include/recommendation/ClustersClassifierCreator.hpp +++ b/src/libs/recommendation/include/recommendation/ClustersClassifierCreator.hpp @@ -21,15 +21,10 @@ #include -namespace Database -{ - class Session; -} - namespace Recommendation { - class Classifier; + class IClassifier; - std::unique_ptr createClustersClassifier(); + std::unique_ptr createClustersClassifier(Database::Session& session); } diff --git a/src/libs/recommendation/include/recommendation/FeaturesClassifierCreator.hpp b/src/libs/recommendation/include/recommendation/FeaturesClassifierCreator.hpp index ce4e3a74..9a763020 100644 --- a/src/libs/recommendation/include/recommendation/FeaturesClassifierCreator.hpp +++ b/src/libs/recommendation/include/recommendation/FeaturesClassifierCreator.hpp @@ -21,15 +21,10 @@ #include -namespace Database -{ - class Session; -} - namespace Recommendation { - class Classifier; + class IClassifier; - std::unique_ptr createFeaturesClassifier(); + std::unique_ptr createFeaturesClassifier(); } diff --git a/src/libs/recommendation/include/recommendation/Classifier.hpp b/src/libs/recommendation/include/recommendation/IClassifier.hpp similarity index 85% rename from src/libs/recommendation/include/recommendation/Classifier.hpp rename to src/libs/recommendation/include/recommendation/IClassifier.hpp index 3c7c7e63..6619ccd2 100644 --- a/src/libs/recommendation/include/recommendation/Classifier.hpp +++ b/src/libs/recommendation/include/recommendation/IClassifier.hpp @@ -32,21 +32,19 @@ namespace Database namespace Recommendation { - class Classifier + class IClassifier { public: - virtual ~Classifier() = default; - - virtual void classify() = 0; + virtual ~IClassifier() = default; virtual bool isTrackClassified(Database::IdType trackId) const = 0; virtual bool isReleaseClassified(Database::IdType releaseId) const = 0; virtual bool isArtistClassified(Database::IdType artistId) const = 0; - virtual std::vector getSimilarTracksFromTrackList(Database::Session& session, Database::IdType tracklistId, std::size_t maxCount) = 0; + virtual std::vector getSimilarTracksFromTrackList(Database::Session& session, Database::IdType tracklistId, std::size_t maxCount) const = 0; virtual std::vector getSimilarTracks(Database::Session& session, const std::unordered_set& tracksId, std::size_t maxCount) const = 0; - virtual std::vector getSimilarReleases(Database::IdType releaseId, std::size_t maxCount) const = 0; - virtual std::vector getSimilarArtists(Database::IdType artistId, std::size_t maxCount) const = 0; + virtual std::vector getSimilarReleases(Database::Session& session, Database::IdType releaseId, std::size_t maxCount) const = 0; + virtual std::vector getSimilarArtists(Database::Session& session, Database::IdType artistId, std::size_t maxCount) const = 0; }; } // ns Recommendation diff --git a/src/libs/recommendation/include/recommendation/IEngine.hpp b/src/libs/recommendation/include/recommendation/IEngine.hpp index a28adfaf..91f10c50 100644 --- a/src/libs/recommendation/include/recommendation/IEngine.hpp +++ b/src/libs/recommendation/include/recommendation/IEngine.hpp @@ -23,7 +23,6 @@ #include #include "database/Types.hpp" -#include "Classifier.hpp" namespace Database { @@ -32,12 +31,13 @@ namespace Database namespace Recommendation { - class Classifier; class IEngine { public: virtual ~IEngine() = default; + virtual void reloadSettings(Database::Session& session) = 0; + // Closest results first virtual std::vector getSimilarTracksFromTrackList(Database::Session& session, Database::IdType tracklistId, std::size_t maxCount) = 0; virtual std::vector getSimilarTracks(Database::Session& session, const std::unordered_set& tracksId, std::size_t maxCount) = 0; @@ -45,7 +45,7 @@ namespace Recommendation virtual std::vector getSimilarArtists(Database::Session& session, Database::IdType artistId, std::size_t maxCount) = 0; }; - std::unique_ptr createEngine(); + std::unique_ptr createEngine(Database::Session& session); } // ns Recommendation diff --git a/src/lms/main.cpp b/src/lms/main.cpp index 7cf6af96..f2a71fa0 100644 --- a/src/lms/main.cpp +++ b/src/lms/main.cpp @@ -143,7 +143,10 @@ int main(int argc, char* argv[]) ServiceProvider::assign(Auth::createPasswordService(ServiceProvider::get()->getULong("login-throttler-max-entriees", 10000))); Scanner::IMediaScanner& mediaScanner {ServiceProvider::assign(Scanner::createMediaScanner(database))}; - ServiceProvider::assign(Recommendation::createEngine()); + { + Database::Session session {database}; + ServiceProvider::assign(Recommendation::createEngine(session)); + } CoverArt::IGrabber& coverArtGrabber {ServiceProvider::assign(CoverArt::createGrabber(argv[0]))}; coverArtGrabber.setDefaultCover(server.appRoot() + "/images/unknown-cover.jpg"); diff --git a/src/lms/ui/admin/DatabaseSettingsView.cpp b/src/lms/ui/admin/DatabaseSettingsView.cpp index 3c4423e1..323dfa7c 100644 --- a/src/lms/ui/admin/DatabaseSettingsView.cpp +++ b/src/lms/ui/admin/DatabaseSettingsView.cpp @@ -49,7 +49,7 @@ class DatabaseSettingsModel : public Wt::WFormModel static const Field MediaDirectoryField; static const Field UpdatePeriodField; static const Field UpdateStartTimeField; - static const Field SimilarityEngineTypeField; + static const Field RecommendationEngineTypeField; static const Field TagsField; DatabaseSettingsModel() @@ -60,7 +60,7 @@ class DatabaseSettingsModel : public Wt::WFormModel addField(MediaDirectoryField); addField(UpdatePeriodField); addField(UpdateStartTimeField); - addField(SimilarityEngineTypeField); + addField(RecommendationEngineTypeField); addField(TagsField); auto dirValidator {std::make_shared()}; @@ -69,7 +69,7 @@ class DatabaseSettingsModel : public Wt::WFormModel setValidator(UpdatePeriodField, createMandatoryValidator()); setValidator(UpdateStartTimeField, createMandatoryValidator()); - setValidator(SimilarityEngineTypeField, createMandatoryValidator()); + setValidator(RecommendationEngineTypeField, createMandatoryValidator()); setValidator(TagsField, createTagsValidator()); // populate the model with initial data @@ -78,7 +78,7 @@ class DatabaseSettingsModel : public Wt::WFormModel std::shared_ptr updatePeriodModel() { return _updatePeriodModel; } std::shared_ptr updateStartTimeModel() { return _updateStartTimeModel; } - std::shared_ptr similarityEngineTypeModel() { return _similarityEngineTypeModel; } + std::shared_ptr recommendationEngineTypeModel() { return _recommendationEngineTypeModel; } void loadData() { @@ -96,9 +96,9 @@ class DatabaseSettingsModel : public Wt::WFormModel if (startTimeRow) setValue(UpdateStartTimeField, _updateStartTimeModel->getString(*startTimeRow)); - auto similarityEngineTypeRow {_similarityEngineTypeModel->getRowFromValue(scanSettings->getSimilarityEngineType())}; - if (similarityEngineTypeRow) - setValue(SimilarityEngineTypeField, _similarityEngineTypeModel->getString(*similarityEngineTypeRow)); + auto recommendationEngineTypeRow {_recommendationEngineTypeModel->getRowFromValue(scanSettings->getRecommendationEngineType())}; + if (recommendationEngineTypeRow) + setValue(RecommendationEngineTypeField, _recommendationEngineTypeModel->getString(*recommendationEngineTypeRow)); auto clusterTypes {scanSettings->getClusterTypes()}; if (!clusterTypes.empty()) @@ -125,9 +125,9 @@ class DatabaseSettingsModel : public Wt::WFormModel if (startTimeRow) scanSettings.modify()->setUpdateStartTime(_updateStartTimeModel->getValue(*startTimeRow)); - auto similarityEngineTypeRow {_similarityEngineTypeModel->getRowFromString(valueText(SimilarityEngineTypeField))}; - if (similarityEngineTypeRow) - scanSettings.modify()->setSimilarityEngineType(_similarityEngineTypeModel->getValue(*similarityEngineTypeRow)); + auto recommendationEngineTypeRow {_recommendationEngineTypeModel->getRowFromString(valueText(RecommendationEngineTypeField))}; + if (recommendationEngineTypeRow) + scanSettings.modify()->setRecommendationEngineType(_recommendationEngineTypeModel->getValue(*recommendationEngineTypeRow)); auto clusterTypes {StringUtils::splitString(valueText(TagsField).toUTF8(), " ")}; scanSettings.modify()->setClusterTypes(LmsApp->getDbSession(), std::set(clusterTypes.begin(), clusterTypes.end())); @@ -156,22 +156,22 @@ class DatabaseSettingsModel : public Wt::WFormModel _updateStartTimeModel->add(time.toString(), time); } - _similarityEngineTypeModel = std::make_shared>(); - _similarityEngineTypeModel->add(Wt::WString::tr("Lms.Admin.Database.similarity-engine-type.clusters"), ScanSettings::SimilarityEngineType::Clusters); - _similarityEngineTypeModel->add(Wt::WString::tr("Lms.Admin.Database.similarity-engine-type.features"), ScanSettings::SimilarityEngineType::Features); + _recommendationEngineTypeModel = std::make_shared>(); + _recommendationEngineTypeModel->add(Wt::WString::tr("Lms.Admin.Database.recommendation-engine-type.clusters"), ScanSettings::RecommendationEngineType::Clusters); + _recommendationEngineTypeModel->add(Wt::WString::tr("Lms.Admin.Database.recommendation-engine-type.features"), ScanSettings::RecommendationEngineType::Features); } std::shared_ptr> _updatePeriodModel; std::shared_ptr> _updateStartTimeModel; - std::shared_ptr> _similarityEngineTypeModel; + std::shared_ptr> _recommendationEngineTypeModel; }; -const Wt::WFormModel::Field DatabaseSettingsModel::MediaDirectoryField = "media-directory"; -const Wt::WFormModel::Field DatabaseSettingsModel::UpdatePeriodField = "update-period"; -const Wt::WFormModel::Field DatabaseSettingsModel::UpdateStartTimeField = "update-start-time"; -const Wt::WFormModel::Field DatabaseSettingsModel::SimilarityEngineTypeField = "similarity-engine-type"; -const Wt::WFormModel::Field DatabaseSettingsModel::TagsField = "tags"; +const Wt::WFormModel::Field DatabaseSettingsModel::MediaDirectoryField = "media-directory"; +const Wt::WFormModel::Field DatabaseSettingsModel::UpdatePeriodField = "update-period"; +const Wt::WFormModel::Field DatabaseSettingsModel::UpdateStartTimeField = "update-start-time"; +const Wt::WFormModel::Field DatabaseSettingsModel::RecommendationEngineTypeField = "recommendation-engine-type"; +const Wt::WFormModel::Field DatabaseSettingsModel::TagsField = "tags"; DatabaseSettingsView::DatabaseSettingsView() { @@ -207,10 +207,10 @@ DatabaseSettingsView::refreshView() updateStartTime->setModel(model->updateStartTimeModel()); t->setFormWidget(DatabaseSettingsModel::UpdateStartTimeField, std::move(updateStartTime)); - // Similarity engine type - auto similarityEngineType {std::make_unique()}; - similarityEngineType->setModel(model->similarityEngineTypeModel()); - t->setFormWidget(DatabaseSettingsModel::SimilarityEngineTypeField, std::move(similarityEngineType)); + // recommendation engine type + auto recommendationEngineType {std::make_unique()}; + recommendationEngineType->setModel(model->recommendationEngineTypeModel()); + t->setFormWidget(DatabaseSettingsModel::RecommendationEngineTypeField, std::move(recommendationEngineType)); // Tags t->setFormWidget(DatabaseSettingsModel::TagsField, std::make_unique()); diff --git a/src/test/database/DatabaseTest.cpp b/src/test/database/DatabaseTest.cpp index 005439d1..5ad5b60b 100644 --- a/src/test/database/DatabaseTest.cpp +++ b/src/test/database/DatabaseTest.cpp @@ -453,12 +453,24 @@ testSingleTrackSingleCluster(Session& session) CHECK(track->getClusterIds().empty()); } + { + auto transaction {session.createSharedTransaction()}; + CHECK(Track::getAllIdsWithClusters(session).empty()); + } + { auto transaction {session.createUniqueTransaction()}; cluster1.get().modify()->addTrack(track.get()); } + { + auto transaction {session.createSharedTransaction()}; + auto tracks {Track::getAllIdsWithClusters(session)}; + CHECK(tracks.size() == 1); + CHECK(tracks.front() == track.getId()); + } + { auto transaction {session.createSharedTransaction()}; auto clusters {Cluster::getAllOrphans(session)}; @@ -611,6 +623,11 @@ testSingleTrackSingleReleaseSingleCluster(Session& session) ScopedClusterType clusterType {session, "MyClusterType"}; ScopedCluster cluster {session, clusterType .lockAndGet(), "MyCluster"}; + { + auto transaction {session.createSharedTransaction()}; + CHECK(Release::getAllIdsWithClusters(session).empty()); + } + { auto transaction {session.createUniqueTransaction()}; @@ -618,6 +635,13 @@ testSingleTrackSingleReleaseSingleCluster(Session& session) cluster.get().modify()->addTrack(track.get()); } + { + auto transaction {session.createSharedTransaction()}; + auto releases {Release::getAllIdsWithClusters(session)}; + CHECK(releases.size() == 1); + CHECK(releases.front() == release.getId()); + } + { auto transaction {session.createSharedTransaction()}; @@ -858,6 +882,11 @@ testSingleTrackSingleReleaseSingleArtistSingleCluster(Session& session) ScopedClusterType clusterType {session, "MyType"}; ScopedCluster cluster {session, clusterType.lockAndGet(), "MyCluster"}; + { + auto transaction {session.createSharedTransaction()}; + CHECK(Artist::getAllIdsWithClusters(session).empty()); + } + { auto transaction {session.createUniqueTransaction()}; @@ -875,6 +904,13 @@ testSingleTrackSingleReleaseSingleArtistSingleCluster(Session& session) CHECK(Release::getAllOrphans(session).empty()); } + { + auto transaction {session.createSharedTransaction()}; + auto artists {Artist::getAllIdsWithClusters(session)}; + CHECK(artists.size() == 1); + CHECK(artists.front() == artist.getId()); + } + { auto transaction {session.createSharedTransaction()}; diff --git a/src/tools/recommendation-features/LmsRecommendationFeatures.cpp b/src/tools/recommendation-features/LmsRecommendationFeatures.cpp index 6426dff5..a8c591be 100644 --- a/src/tools/recommendation-features/LmsRecommendationFeatures.cpp +++ b/src/tools/recommendation-features/LmsRecommendationFeatures.cpp @@ -32,6 +32,7 @@ #include "utils/Service.hpp" #include "utils/StreamLogger.hpp" #include "recommendation/IEngine.hpp" +#include "recommendation/IClassifier.hpp" #include "recommendation/FeaturesClassifierCreator.hpp" int main(int argc, char *argv[]) @@ -50,11 +51,9 @@ int main(int argc, char *argv[]) Database::Db db {ServiceProvider::get()->getPath("working-dir") / "lms.db"}; Database::Session session {db}; - auto classifier {Recommendation::createFeaturesClassifier()}; - - std::cout << "Classifying tracks..." << std::endl; // may be long... - classifier->classify(); + std::cout << "Classifying tracks..." << std::endl; + auto classifier {Recommendation::createFeaturesClassifier()}; std::cout << "Classifying tracks DONE" << std::endl; const std::vector trackIds = std::invoke([&]() @@ -106,7 +105,7 @@ int main(int argc, char *argv[]) }; std::cout << "Processing release '" << releaseToString(releaseId) << "'" << std::endl; - for (Database::IdType similarReleaseId : classifier->getSimilarReleases({releaseId}, 3)) + for (Database::IdType similarReleaseId : classifier->getSimilarReleases(session, {releaseId}, 3)) std::cout << "\t- Similar release '" << releaseToString(similarReleaseId) << "'" << std::endl; } @@ -128,7 +127,7 @@ int main(int argc, char *argv[]) }; std::cout << "Processing artist '" << artistToString(artistId) << "'" << std::endl; - for (Database::IdType similarArtistId : classifier->getSimilarArtists({artistId}, 3)) + for (Database::IdType similarArtistId : classifier->getSimilarArtists(session, {artistId}, 3)) std::cout << "\t- Similar artist '" << artistToString(similarArtistId) << "'" << std::endl; }