Simplified the similarity engine loading (since there is no more the features-based engine)

This commit is contained in:
emeric
2023-11-10 20:59:25 +01:00
parent 8ea532fa26
commit 636b70450b
10 changed files with 504 additions and 711 deletions
@@ -33,22 +33,17 @@
namespace Recommendation
{
namespace
{
Database::ScanSettings::SimilarityEngineType getSimilarityEngineType(Database::Session& session)
{
auto transaction{ session.createSharedTransaction() };
static
std::string_view
engineTypeToString(EngineType engineType)
{
switch (engineType)
{
case EngineType::Clusters: return "clusters";
case EngineType::Features: return "features";
return Database::ScanSettings::get(session)->getSimilarityEngineType();
}
}
throw LmsException {"Internal error"};
}
std::unique_ptr<IRecommendationService>
createRecommendationService(Database::Db& db)
std::unique_ptr<IRecommendationService> createRecommendationService(Database::Db& db)
{
return std::make_unique<RecommendationService>(db);
}
@@ -56,217 +51,73 @@ namespace Recommendation
RecommendationService::RecommendationService(Database::Db& db)
: _db{ db }
{
load();
}
TrackContainer
RecommendationService::findSimilarTracks(Database::TrackListId trackListId, std::size_t maxCount) const
TrackContainer RecommendationService::findSimilarTracks(Database::TrackListId trackListId, std::size_t maxCount) const
{
TrackContainer res;
std::shared_lock lock {_enginesMutex};
for (const auto& engineType : _enginePriorities)
{
auto itEngine {_engines.find(engineType)};
if (itEngine == std::cend(_engines))
continue;
res = itEngine->second->findSimilarTracksFromTrackList(trackListId, maxCount);
if (!res.empty())
break;
}
if (!_engine)
return res;
return _engine->findSimilarTracksFromTrackList(trackListId, maxCount);
}
TrackContainer
RecommendationService::findSimilarTracks(const std::vector<Database::TrackId>& trackIds, std::size_t maxCount) const
TrackContainer RecommendationService::findSimilarTracks(const std::vector<Database::TrackId>& trackIds, std::size_t maxCount) const
{
TrackContainer res;
std::shared_lock lock {_enginesMutex};
for (EngineType engineType : _enginePriorities)
{
auto itEngine {_engines.find(engineType)};
if (itEngine == std::cend(_engines))
continue;
LMS_LOG(RECOMMENDATION, DEBUG) << "Trying engine '" << engineTypeToString(engineType) << "' to get similar tracks";
const IEngine& engine {*itEngine->second};
res = engine.findSimilarTracks(trackIds, maxCount);
if (!res.empty())
{
LMS_LOG(RECOMMENDATION, DEBUG) << "Got " << res.size() << " similar tracks using engine '" << engineTypeToString(engineType) << "'";
break;
}
}
if (!_engine)
return res;
return _engine->findSimilarTracks(trackIds, maxCount);
}
ReleaseContainer
RecommendationService::getSimilarReleases(Database::ReleaseId releaseId, std::size_t maxCount) const
ReleaseContainer RecommendationService::getSimilarReleases(Database::ReleaseId releaseId, std::size_t maxCount) const
{
ReleaseContainer res;
std::shared_lock lock {_enginesMutex};
for (EngineType engineType : _enginePriorities)
{
auto itEngine {_engines.find(engineType)};
if (itEngine == std::cend(_engines))
continue;
LMS_LOG(RECOMMENDATION, DEBUG) << "Trying engine '" << engineTypeToString(engineType) << "' to get similar releases";
const IEngine& engine {*itEngine->second};
res = engine.getSimilarReleases(releaseId, maxCount);
if (!res.empty())
{
LMS_LOG(RECOMMENDATION, DEBUG) << "Got " << res.size() << " similar releases using engine '" << engineTypeToString(engineType) << "'";
break;
}
LMS_LOG(RECOMMENDATION, DEBUG) << "No result using engine '" << engineTypeToString(engineType) << "'";
}
if (!_engine)
return res;
return _engine->getSimilarReleases(releaseId, maxCount);;
}
ArtistContainer
RecommendationService::getSimilarArtists(Database::ArtistId artistId, EnumSet<Database::TrackArtistLinkType> linkTypes, std::size_t maxCount) const
ArtistContainer RecommendationService::getSimilarArtists(Database::ArtistId artistId, EnumSet<Database::TrackArtistLinkType> linkTypes, std::size_t maxCount) const
{
ArtistContainer res;
std::shared_lock lock {_enginesMutex};
for (EngineType engineType : _enginePriorities)
{
auto itEngine {_engines.find(engineType)};
if (itEngine == std::cend(_engines))
continue;
LMS_LOG(RECOMMENDATION, DEBUG) << "Trying engine '" << engineTypeToString(engineType) << "' to get similar artists";
const IEngine& engine {*itEngine->second};
res = engine.getSimilarArtists(artistId, linkTypes, maxCount);
if (!res.empty())
{
LMS_LOG(RECOMMENDATION, DEBUG) << "Got " << res.size() << " similar artists using engine '" << engineTypeToString(engineType) << "'";
if (!_engine)
return res;
}
}
return _engine->getSimilarArtists(artistId, linkTypes, maxCount);
return res;
}
static
Database::ScanSettings::SimilarityEngineType
getSimilarityEngineType(Database::Session& session)
{
auto transaction {session.createSharedTransaction()};
return Database::ScanSettings::get(session)->getSimilarityEngineType();
}
void
RecommendationService::load(bool forceReload, const ProgressCallback& progressCallback)
void RecommendationService::load()
{
using namespace Database;
LMS_LOG(RECOMMENDATION, INFO) << "Reloading recommendation engines...";
EngineContainer enginesToLoad;
{
std::unique_lock controlLock {_controlMutex};
{
std::unique_lock lock {_enginesMutex};
_engines.clear();
}
switch (getSimilarityEngineType(_db.getTLSSession()))
{
case ScanSettings::SimilarityEngineType::Clusters:
_enginePriorities = {EngineType::Clusters};
enginesToLoad.try_emplace(EngineType::Clusters, createClustersEngine(_db));
if (_engineType != EngineType::Clusters)
{
_engineType = EngineType::Clusters;
_engine = createClustersEngine(_db);
}
break;
case ScanSettings::SimilarityEngineType::Features:
_enginePriorities = {EngineType::Features, EngineType::Clusters};
// not same order since clusters is faster to load
enginesToLoad.try_emplace(EngineType::Clusters, createClustersEngine(_db));
enginesToLoad.try_emplace(EngineType::Features, createFeaturesEngine(_db));
break;
case ScanSettings::SimilarityEngineType::None:
_enginePriorities.clear();
_engineType.reset();
_engine.reset();
break;
}
assert(_pendingEngines.empty());
for (auto& [engineType, engine] : enginesToLoad)
_pendingEngines.push_back(engine.get());
if (_engine)
_engine->load(false);
}
for (auto& [engineType, engine] : enginesToLoad)
loadPendingEngine(engineType, std::move(engine), forceReload, progressCallback);
_pendingEnginesCondvar.notify_all();
LMS_LOG(RECOMMENDATION, INFO) << "Recommendation engines loaded!";
}
void
RecommendationService::loadPendingEngine(EngineType engineType, std::unique_ptr<IEngine> engine, bool forceReload, const ProgressCallback& progressCallback)
{
if (!_loadCancelled)
{
LMS_LOG(RECOMMENDATION, INFO) << "Initializing engine '" << engineTypeToString(engineType) << "'...";
auto progress {[&](const Progress& progress)
{
progressCallback(progress);
}};
engine->load(forceReload, progressCallback ? progress : ProgressCallback {});
LMS_LOG(RECOMMENDATION, INFO) << "Initializing engine '" << engineTypeToString(engineType) << "': " << (_loadCancelled ? "aborted" : "complete");
}
{
std::scoped_lock lock {_controlMutex};
_pendingEngines.erase(std::find(std::begin(_pendingEngines), std::end(_pendingEngines), engine.get()));
}
if (!_loadCancelled)
{
std::unique_lock lock {_enginesMutex};
_engines.emplace(engineType, std::move(engine));
}
}
void
RecommendationService::cancelLoad()
{
LMS_LOG(RECOMMENDATION, DEBUG) << "Cancelling loading...";
std::unique_lock controlLock {_controlMutex};
assert(!_loadCancelled);
_loadCancelled = true;
LMS_LOG(RECOMMENDATION, DEBUG) << "Still " << _pendingEngines.size() << " pending engines!";
for (IEngine* engine : _pendingEngines)
{
engine->requestCancelLoad();
}
_pendingEnginesCondvar.wait(controlLock, [this] {return _pendingEngines.empty();});
_loadCancelled = false;
LMS_LOG(RECOMMENDATION, DEBUG) << "Cancelling loading DONE";
}
} // ns Similarity
@@ -19,11 +19,7 @@
#pragma once
#include <condition_variable>
#include <mutex>
#include <shared_mutex>
#include <unordered_map>
#include <vector>
#include <optional>
#include "services/recommendation/IRecommendationService.hpp"
#include "IEngine.hpp"
@@ -48,13 +44,10 @@ namespace Recommendation
~RecommendationService() = default;
RecommendationService(const RecommendationService&) = delete;
RecommendationService(RecommendationService&&) = delete;
RecommendationService& operator=(const RecommendationService&) = delete;
RecommendationService& operator=(RecommendationService&&) = delete;
private:
void load(bool forceReload, const ProgressCallback& progressCallback) override;
void cancelLoad() override;
void load() override;
TrackContainer findSimilarTracks(Database::TrackListId tracklistId, std::size_t maxCount) const override;
TrackContainer findSimilarTracks(const std::vector<Database::TrackId>& tracksId, std::size_t maxCount) const override;
@@ -66,19 +59,8 @@ namespace Recommendation
void loadPendingEngine(EngineType engineType, std::unique_ptr<IEngine> engine, bool forceReload, const ProgressCallback& progressCallback);
Database::Db& _db;
std::mutex _controlMutex;
bool _loadCancelled {};
using EngineContainer = std::unordered_map<EngineType, std::unique_ptr<IEngine>>;
EngineContainer _engines;
mutable std::shared_mutex _enginesMutex;
std::vector<IEngine*> _pendingEngines;
std::shared_mutex _pendingEnginesMutex;
std::condition_variable _pendingEnginesCondvar;
std::vector<EngineType> _enginePriorities; // ordered by priority
std::optional<EngineType> _engineType;
std::unique_ptr<IEngine> _engine;
};
} // ns Recommendation
@@ -20,6 +20,7 @@
#pragma once
#include <memory>
#include <vector>
#include "utils/EnumSet.hpp"
#include "services/database/TrackListId.hpp"
#include "services/database/Types.hpp"
@@ -37,8 +38,7 @@ namespace Recommendation
public:
virtual ~IRecommendationService() = default;
virtual void load(bool forceReload, const ProgressCallback& progressCallback = {}) = 0;
virtual void cancelLoad() = 0; // wait for cancel done
virtual void load() = 0;
virtual TrackContainer findSimilarTracks(Database::TrackListId tracklistId, std::size_t maxCount) const = 0;
virtual TrackContainer findSimilarTracks(const std::vector<Database::TrackId>& tracksId, std::size_t maxCount) const = 0;
@@ -25,7 +25,6 @@
#include "services/database/Cluster.hpp"
#include "services/database/TrackFeatures.hpp"
#include "services/database/ScanSettings.hpp"
#include "services/recommendation/IRecommendationService.hpp"
#include "utils/Exception.hpp"
#include "utils/IConfig.hpp"
#include "utils/Logger.hpp"
@@ -38,12 +37,13 @@
#include "ScanStepScanFiles.hpp"
#include "ScanStepComputeClusterStats.hpp"
namespace Scanner
{
using namespace Database;
namespace {
Wt::WDate
getNextMonday(Wt::WDate current)
namespace
{
Wt::WDate getNextMonday(Wt::WDate current)
{
do
{
@@ -53,8 +53,7 @@ getNextMonday(Wt::WDate current)
return current;
}
Wt::WDate
getNextFirstOfMonth(Wt::WDate current)
Wt::WDate getNextFirstOfMonth(Wt::WDate current)
{
do
{
@@ -63,20 +62,15 @@ getNextFirstOfMonth(Wt::WDate current)
return current;
}
} // namespace
namespace Scanner {
std::unique_ptr<IScannerService>
createScannerService(Db& db, Recommendation::IRecommendationService& recommendationService)
std::unique_ptr<IScannerService> createScannerService(Db& db)
{
return std::make_unique<ScannerService>(db, recommendationService);
return std::make_unique<ScannerService>(db);
}
ScannerService::ScannerService(Db& db, Recommendation::IRecommendationService& recommendationService)
: _recommendationService {recommendationService}
, _db {db}
ScannerService::ScannerService(Db& db)
: _db{ db }
, _dbSession{ db }
{
_ioService.setThreadCount(1);
@@ -93,8 +87,7 @@ ScannerService::~ScannerService()
LMS_LOG(DBUPDATER, INFO) << "Service stopped!";
}
void
ScannerService::start()
void ScannerService::start()
{
std::scoped_lock lock{ _controlMutex };
@@ -103,30 +96,22 @@ ScannerService::start()
if (_abortScan)
return;
_recommendationService.load(false,
[](const Recommendation::Progress& progress)
{
LMS_LOG(DBUPDATER, DEBUG) << "Reloading recommendation : " << progress.processedElems << "/" << progress.totalElems;
});
scheduleNextScan();
});
_ioService.start();
}
void
ScannerService::stop()
void ScannerService::stop()
{
std::scoped_lock lock{ _controlMutex };
_abortScan = true;
_scheduleTimer.cancel();
_recommendationService.cancelLoad();
_ioService.stop();
}
void
ScannerService::abortScan()
void ScannerService::abortScan()
{
LMS_LOG(DBUPDATER, DEBUG) << "Aborting scan...";
std::scoped_lock lock{ _controlMutex };
@@ -135,7 +120,6 @@ ScannerService::abortScan()
_abortScan = true;
_scheduleTimer.cancel();
_recommendationService.cancelLoad();
_ioService.stop();
LMS_LOG(DBUPDATER, DEBUG) << "Scan abort done!";
@@ -143,8 +127,7 @@ ScannerService::abortScan()
_ioService.start();
}
void
ScannerService::requestImmediateScan(bool force)
void ScannerService::requestImmediateScan(bool force)
{
abortScan();
_ioService.post([=]()
@@ -156,8 +139,7 @@ ScannerService::requestImmediateScan(bool force)
});
}
void
ScannerService::requestReload()
void ScannerService::requestReload()
{
abortScan();
_ioService.post([=]()
@@ -169,8 +151,7 @@ ScannerService::requestReload()
});
}
ScannerService::Status
ScannerService::getStatus() const
ScannerService::Status ScannerService::getStatus() const
{
Status res;
@@ -184,8 +165,7 @@ ScannerService::getStatus() const
return res;
}
void
ScannerService::scheduleNextScan()
void ScannerService::scheduleNextScan()
{
LMS_LOG(DBUPDATER, DEBUG) << "Scheduling next scan";
@@ -238,8 +218,7 @@ ScannerService::scheduleNextScan()
_events.scanScheduled.emit(_nextScheduledScan);
}
void
ScannerService::scheduleScan(bool force, const Wt::WDateTime& dateTime)
void ScannerService::scheduleScan(bool force, const Wt::WDateTime& dateTime)
{
auto cb{ [=](boost::system::error_code ec)
{
@@ -267,8 +246,7 @@ ScannerService::scheduleScan(bool force, const Wt::WDateTime& dateTime)
}
}
void
ScannerService::scan(bool forceScan)
void ScannerService::scan(bool forceScan)
{
_events.scanStarted.emit();
@@ -328,8 +306,7 @@ ScannerService::scan(bool forceScan)
}
}
void
ScannerService::refreshScanSettings()
void ScannerService::refreshScanSettings()
{
ScannerSettings newSettings{ readSettings() };
if (_settings == newSettings)
@@ -362,8 +339,7 @@ ScannerService::refreshScanSettings()
_scanSteps.push_back(std::make_unique<ScanStepCheckDuplicatedDbFiles>(params));
}
ScannerSettings
ScannerService::readSettings()
ScannerSettings ScannerService::readSettings()
{
ScannerSettings newSettings;
@@ -383,7 +359,6 @@ ScannerService::readSettings()
std::transform(std::cbegin(fileExtensions), std::end(fileExtensions), std::back_inserter(newSettings.supportedExtensions),
[](const std::filesystem::path& extension) { return std::filesystem::path{ StringUtils::stringToLower(extension.string()) }; });
}
newSettings.similarityServiceType = scanSettings->getSimilarityEngineType();
newSettings.mediaDirectory = scanSettings->getMediaDirectory();
const auto clusterTypes = scanSettings->getClusterTypes();
@@ -399,8 +374,7 @@ ScannerService::readSettings()
return newSettings;
}
void
ScannerService::notifyInProgress(const ScanStepStats& stepStats)
void ScannerService::notifyInProgress(const ScanStepStats& stepStats)
{
{
std::unique_lock lock{ _statusMutex };
@@ -412,8 +386,7 @@ ScannerService::notifyInProgress(const ScanStepStats& stepStats)
_lastScanInProgressEmit = now;
}
void
ScannerService::notifyInProgressIfNeeded(const ScanStepStats& stepStats)
void ScannerService::notifyInProgressIfNeeded(const ScanStepStats& stepStats)
{
std::chrono::system_clock::time_point now{ std::chrono::system_clock::now() };
@@ -38,23 +38,16 @@
#include "IScanStep.hpp"
#include "ScannerSettings.hpp"
namespace Recommendation
{
class IRecommendationService;
}
namespace Scanner
{
class ScannerService : public IScannerService
{
public:
ScannerService(Database::Db& db, Recommendation::IRecommendationService& recommendationService);
ScannerService(Database::Db& db);
~ScannerService();
ScannerService(const ScannerService&) = delete;
ScannerService(ScannerService&&) = delete;
ScannerService& operator=(const ScannerService&) = delete;
ScannerService& operator=(ScannerService&&) = delete;
void requestReload() override;
void requestImmediateScan(bool force) override;
@@ -80,13 +73,12 @@ namespace Scanner
// Helpers
void refreshScanSettings();
ScannerSettings readSettings();
void reloadRecommendationService();
void notifyInProgressIfNeeded(const ScanStepStats& stats);
void notifyInProgress(const ScanStepStats& stats);
void reloadSimilarityEngine(ScanStats& stats);
Recommendation::IRecommendationService& _recommendationService;
std::vector<std::unique_ptr<IScanStep>> _scanSteps;
std::mutex _controlMutex;
@@ -34,7 +34,6 @@ namespace Scanner
Wt::WTime startTime;
Database::ScanSettings::UpdatePeriod updatePeriod {Database::ScanSettings::UpdatePeriod::Never};
std::vector<std::filesystem::path> supportedExtensions;
Database::ScanSettings::SimilarityEngineType similarityServiceType;
std::filesystem::path mediaDirectory;
bool skipDuplicateMBID {};
std::set<std::string> clusterTypeNames;
@@ -45,7 +44,6 @@ namespace Scanner
&& startTime == rhs.startTime
&& updatePeriod == rhs.updatePeriod
&& supportedExtensions == rhs.supportedExtensions
&& similarityServiceType == rhs.similarityServiceType
&& mediaDirectory == rhs.mediaDirectory
&& skipDuplicateMBID == rhs.skipDuplicateMBID
&& clusterTypeNames == rhs.clusterTypeNames;
@@ -29,11 +29,6 @@ namespace Database
class Db;
}
namespace Recommendation
{
class IRecommendationService;
}
namespace Scanner
{
@@ -66,7 +61,7 @@ namespace Scanner
virtual Events& getEvents() = 0;
};
std::unique_ptr<IScannerService> createScannerService(Database::Db& db, Recommendation::IRecommendationService& recommendationEngine);
std::unique_ptr<IScannerService> createScannerService(Database::Db& db);
} // Scanner
+1 -1
View File
@@ -270,7 +270,7 @@ int main(int argc, char* argv[])
Service<Cover::ICoverService> coverService{ Cover::createCoverService(database, argv[0], server.appRoot() + "/images/unknown-cover.jpg") };
Service<Recommendation::IRecommendationService> recommendationService{ Recommendation::createRecommendationService(database) };
Service<Recommendation::IPlaylistGeneratorService> playlistGeneratorService{ Recommendation::createPlaylistGeneratorService(database, *recommendationService.get()) };
Service<Scanner::IScannerService> scannerService{ Scanner::createScannerService(database, *recommendationService) };
Service<Scanner::IScannerService> scannerService{ Scanner::createScannerService(database) };
scannerService->getEvents().scanComplete.connect([&]
{
@@ -29,6 +29,7 @@
#include "services/database/Cluster.hpp"
#include "services/database/ScanSettings.hpp"
#include "services/database/Session.hpp"
#include "services/recommendation/IRecommendationService.hpp"
#include "services/scanner/IScannerService.hpp"
#include "utils/Logger.hpp"
#include "utils/Service.hpp"
@@ -232,6 +233,7 @@ DatabaseSettingsView::refreshView()
{
model->saveData();
Service<Recommendation::IRecommendationService>::get()->load();
Service<Scanner::IScannerService>::get()->requestImmediateScan(false);
LmsApp->notifyMsg(Notification::Type::Info, Wt::WString::tr("Lms.Admin.Database.database"), Wt::WString::tr("Lms.Admin.Database.settings-saved"));
}
@@ -157,16 +157,16 @@ int main(int argc, char* argv[])
Db db{ config->getPath("working-dir") / "lms.db" };
Session session{ db };
std::cout << "Creating recommendation recommendationService..." << std::endl;
std::cout << "Creating recommendation service..." << std::endl;
const auto recommendationService{ Recommendation::createRecommendationService(db) };
std::cout << "Recommendation recommendationService created!" << std::endl;
std::cout << "Recommendation service created!" << std::endl;
std::cout << "Loading recommendation recommendationService..." << std::endl;
recommendationService->load(false);
std::cout << "Loading recommendation service..." << std::endl;
recommendationService->load();
unsigned maxSimilarityCount{ vm["max"].as<unsigned>() };
std::cout << "Recommendation recommendationService loaded!" << std::endl;
std::cout << "Recommendation service loaded!" << std::endl;
if (vm.count("tracks"))
dumpTracksRecommendation(db, *recommendationService, maxSimilarityCount);