/* * 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 . */ #include "ClustersEngine.hpp" #include #include #include #include #include #include "database/IDb.hpp" #include "database/Session.hpp" #include "database/objects/Artist.hpp" #include "database/objects/Cluster.hpp" #include "database/objects/Release.hpp" #include "database/objects/Track.hpp" #include "database/objects/TrackList.hpp" #include "core/ILogger.hpp" #include "core/ITraceLogger.hpp" #include "track-selection-constraints/DuplicateTrackConstraint.hpp" #include "track-selection-constraints/SameArtistConstraint.hpp" #include "track-selection-constraints/SameRecordingMBIDConstraint.hpp" #include "track-selection-constraints/SameReleaseConstraint.hpp" #include "track-selection-constraints/TrackCandidateContext.hpp" #define LOG(sev, message) LMS_LOG(RECOMMENDATION, sev, "[clusters] " << message) namespace lms::recommendation { namespace { template std::vector> computeClusterOverlap( const std::unordered_map>& profileMap, const std::unordered_set& excludeIds, const std::unordered_set& queryClusters) { std::vector> results; for (const auto& [candidateId, candidateClusters] : profileMap) { if (excludeIds.contains(candidateId)) continue; std::size_t count{}; for (const db::ClusterId clusterId : candidateClusters) if (queryClusters.contains(clusterId)) ++count; if (count > 0) results.emplace_back(candidateId, count); } return results; } template ResultContainer findSimilarByClusterOverlap( const std::unordered_map>& profileMap, IdType queryId, const std::vector& queryClusters, std::size_t maxCount) { const std::unordered_set querySet{ queryClusters.cbegin(), queryClusters.cend() }; auto overlapCounts{ computeClusterOverlap(profileMap, { queryId }, querySet) }; const std::size_t resultCount{ std::min(maxCount, overlapCounts.size()) }; std::partial_sort(overlapCounts.begin(), std::next(overlapCounts.begin(), resultCount), overlapCounts.end(), [](const auto& a, const auto& b) { return a.second > b.second; }); ResultContainer res; res.reserve(resultCount); for (std::size_t i{}; i < resultCount; ++i) res.push_back({ .id = overlapCounts[i].first, .distanceToFirst = {}, .distanceToPrevious = {} }); return res; } } // namespace std::unique_ptr createClustersEngine(db::IDb& db) { return std::make_unique(db); } ClusterEngine::ClusterEngine(db::IDb& db) : _db{ db } { constexpr float sameReleaseWeight{ 0.5F }; constexpr float sameArtistWeight{ 0.5F }; _trackEvaluator.addHardConstraint(std::make_unique()); _trackEvaluator.addHardConstraint(std::make_unique(_trackMetadata)); _trackEvaluator.addSoftConstraint(std::make_unique(_trackMetadata), sameReleaseWeight); _trackEvaluator.addSoftConstraint(std::make_unique(_trackMetadata), sameArtistWeight); } ClusterEngine::~ClusterEngine() = default; void ClusterEngine::load() { LMS_SCOPED_TRACE_OVERVIEW("ClustersEngine", "Loading"); LOG(INFO, "loading..."); _trackMetadata.clear(); _trackClusters.clear(); _releaseClusters.clear(); _artistClusters.clear(); db::Session& session{ _db.getTLSSession() }; auto transaction{ session.createReadTransaction() }; buildTrackClusters(session); buildTrackMetadata(session); buildReleaseClusters(); buildArtistClusters(); LOG(INFO, "loaded " << _trackClusters.size() << " tracks, " << _releaseClusters.size() << " releases, " << _artistClusters.size() << " artists"); } void ClusterEngine::buildTrackMetadata(db::Session& session) { LOG(DEBUG, "building track metadata..."); // Ensure cluster tracks with no release/artist have an entry for (const auto& [trackId, _] : _trackClusters) _trackMetadata.try_emplace(trackId); db::Track::find(session, db::Track::FindParameters{}, [&](const db::Track::pointer& track) { const auto it{ _trackMetadata.find(track->getId()) }; if (it != _trackMetadata.cend()) { it->second.releaseId = track->getReleaseId(); it->second.recordingMBID = track->getRecordingMBID(); } }); db::Artist::find(session, db::Artist::FindParameters{}, [&](const db::Artist::pointer& artist) { const auto mbid{ artist->getMBID() }; // skip "Various Artists" to avoid false artist matches if (mbid && mbid->toString() == "89ad4ac3-39f7-470e-963a-56509c546377") return; std::unordered_set artistTrackIds; { db::Release::FindParameters params; params.setArtist(artist->getId()); for (const db::ReleaseId releaseId : db::Release::findIds(session, params).results) { db::Track::FindParameters trackParams; trackParams.setRelease(releaseId); for (const db::TrackId trackId : db::Track::findIds(session, trackParams).results) artistTrackIds.insert(trackId); } } for (const db::TrackId trackId : artistTrackIds) _trackMetadata[trackId].artistIds.push_back(artist->getId()); }); for (auto& [trackId, metadata] : _trackMetadata) std::sort(metadata.artistIds.begin(), metadata.artistIds.end()); } void ClusterEngine::buildTrackClusters(db::Session& session) { LOG(DEBUG, "building track clusters..."); db::Cluster::find(session, db::Cluster::FindParameters{}, [&](const db::Cluster::pointer& cluster) { const db::ClusterId clusterId{ cluster->getId() }; for (const db::TrackId trackId : cluster->getTracks().results) _trackClusters[trackId].push_back(clusterId); }); } void ClusterEngine::buildReleaseClusters() { LOG(DEBUG, "building release clusters..."); for (const auto& [trackId, clusters] : _trackClusters) { const auto metaIt{ _trackMetadata.find(trackId) }; if (metaIt == _trackMetadata.cend()) continue; if (const db::ReleaseId releaseId{ metaIt->second.releaseId }; releaseId.isValid()) for (const db::ClusterId clusterId : clusters) _releaseClusters[releaseId].push_back(clusterId); } for (auto& [_, clusters] : _releaseClusters) { std::sort(clusters.begin(), clusters.end()); clusters.erase(std::unique(clusters.begin(), clusters.end()), clusters.end()); } } void ClusterEngine::buildArtistClusters() { LOG(DEBUG, "building artist clusters..."); for (const auto& [trackId, clusters] : _trackClusters) { const auto metaIt{ _trackMetadata.find(trackId) }; if (metaIt == _trackMetadata.cend()) continue; for (const db::ArtistId artistId : metaIt->second.artistIds) for (const db::ClusterId clusterId : clusters) _artistClusters[artistId].push_back(clusterId); } for (auto& [_, clusters] : _artistClusters) { std::sort(clusters.begin(), clusters.end()); clusters.erase(std::unique(clusters.begin(), clusters.end()), clusters.end()); } } TrackResults ClusterEngine::findSimilarTracks(std::span trackIds, std::size_t maxCount) const { LMS_SCOPED_TRACE_DETAILED("ClustersEngine", "Find similar tracks"); if (maxCount == 0 || trackIds.empty()) return {}; std::unordered_set queryClusters; for (const db::TrackId trackId : trackIds) { const auto it{ _trackClusters.find(trackId) }; if (it != _trackClusters.cend()) for (const db::ClusterId clusterId : it->second) queryClusters.insert(clusterId); } if (queryClusters.empty()) return {}; const std::unordered_set excludeSet{ std::cbegin(trackIds), std::cend(trackIds) }; auto overlapCounts{ computeClusterOverlap(_trackClusters, excludeSet, queryClusters) }; static constexpr std::size_t oversamplingFactor{ 5 }; const std::size_t candidateCount{ std::min(maxCount * oversamplingFactor, overlapCounts.size()) }; std::partial_sort(overlapCounts.begin(), std::next(overlapCounts.begin(), candidateCount), overlapCounts.end(), [](const auto& a, const auto& b) { return a.second > b.second; }); overlapCounts.resize(candidateCount); std::vector candidates; candidates.reserve(candidateCount); for (const auto& [trackId, count] : overlapCounts) candidates.push_back(trackId); std::vector seeds{ std::cbegin(trackIds), std::cend(trackIds) }; return greedySelect(std::move(candidates), std::move(seeds), maxCount); } TrackResults ClusterEngine::greedySelect(std::vector candidates, std::vector selectedTracks, std::size_t maxCount) const { selectedTracks.reserve(selectedTracks.size() + maxCount); TrackResults res; res.reserve(maxCount); while (res.size() < maxCount && !candidates.empty()) { std::optional bestIdx; float bestScore{ std::numeric_limits::max() }; for (std::size_t i{}; i < candidates.size(); ++i) { const TrackCandidateContext context{ .candidateTrackId = candidates[i], .selectedTracks = selectedTracks, .seedTrackIds = {}, }; if (_trackEvaluator.rejects(context)) continue; const float score{ _trackEvaluator.score(context) }; if (score < bestScore) { bestScore = score; bestIdx = i; } } if (!bestIdx) break; res.push_back({ .id = candidates[*bestIdx], .distanceToFirst = {}, .distanceToPrevious = {} }); selectedTracks.push_back(candidates[*bestIdx]); candidates.erase(std::begin(candidates) + static_cast(*bestIdx)); } return res; } TrackResults ClusterEngine::findSimilarTracksFromTrackList(db::TrackListId tracklistId, std::size_t maxCount) const { LMS_SCOPED_TRACE_DETAILED("ClustersEngine", "Find similar tracks from tracklist"); if (maxCount == 0) return {}; std::vector trackIds; { db::Session& dbSession{ _db.getTLSSession() }; auto transaction{ dbSession.createReadTransaction() }; const db::TrackList::pointer trackList{ db::TrackList::find(dbSession, tracklistId) }; if (!trackList) return {}; trackIds = trackList->getTrackIds(); } if (trackIds.empty()) return {}; return findSimilarTracks(trackIds, maxCount); } ReleaseResults ClusterEngine::findSimilarReleases(db::ReleaseId releaseId, std::size_t maxCount) const { LMS_SCOPED_TRACE_DETAILED("ClustersEngine", "Find similar releases"); if (maxCount == 0) return {}; const auto queryIt{ _releaseClusters.find(releaseId) }; if (queryIt == _releaseClusters.cend() || queryIt->second.empty()) return {}; return findSimilarByClusterOverlap(_releaseClusters, releaseId, queryIt->second, maxCount); } ArtistResults ClusterEngine::findSimilarArtists(db::ArtistId artistId, core::EnumSet linkTypes, std::size_t maxCount) const { LMS_SCOPED_TRACE_DETAILED("ClustersEngine", "Find similar artists"); if (maxCount == 0 || !linkTypes.contains(db::TrackArtistLinkType::Artist)) return {}; const auto queryIt{ _artistClusters.find(artistId) }; if (queryIt == _artistClusters.cend() || queryIt->second.empty()) return {}; return findSimilarByClusterOverlap(_artistClusters, artistId, queryIt->second, maxCount); } TrackResults ClusterEngine::findTrackSimilarityPath(db::TrackId startTrackId, db::TrackId endTrackId, std::size_t maxCount) const { LMS_SCOPED_TRACE_DETAILED("ClustersEngine", "Find track similarity path"); if (maxCount == 0) return {}; if (startTrackId == endTrackId) return { RecommendationResult{ .id = startTrackId, .distanceToFirst = {}, .distanceToPrevious = {} } }; db::Session& dbSession{ _db.getTLSSession() }; auto transaction{ dbSession.createReadTransaction() }; const auto startTrack{ db::Track::find(dbSession, startTrackId) }; const auto endTrack{ db::Track::find(dbSession, endTrackId) }; if (!startTrack || !endTrack) return {}; TrackResults res; res.reserve(std::min(maxCount, 2)); res.push_back({ .id = startTrackId, .distanceToFirst = {}, .distanceToPrevious = {} }); if (maxCount > 1) res.push_back({ .id = endTrackId, .distanceToFirst = {}, .distanceToPrevious = {} }); return res; } } // namespace lms::recommendation #undef LOG