/*
* 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