Restored recommendations based on acoustic similarities (using musicnn), fixes #301

This commit is contained in:
emeric
2026-06-02 08:32:43 +02:00
parent 1524106124
commit eb7f65878f
227 changed files with 10324 additions and 4673 deletions
@@ -19,6 +19,13 @@
#include "ClustersEngine.hpp"
#include <algorithm>
#include <limits>
#include <memory>
#include <optional>
#include <random>
#include <unordered_set>
#include "database/IDb.hpp"
#include "database/Session.hpp"
#include "database/objects/Artist.hpp"
@@ -27,85 +34,356 @@
#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/SameReleaseConstraint.hpp"
#include "track-selection-constraints/TrackCandidateContext.hpp"
#define LOG(sev, message) LMS_LOG(RECOMMENDATION, sev, "[clusters] " << message)
namespace lms::recommendation
{
using namespace db;
namespace
{
template<typename IdType>
std::vector<std::pair<IdType, std::size_t>> computeClusterOverlap(
const std::unordered_map<IdType, std::vector<db::ClusterId>>& profileMap,
const std::unordered_set<IdType>& excludeIds,
const std::unordered_set<db::ClusterId>& queryClusters)
{
std::vector<std::pair<IdType, std::size_t>> 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<typename IdType>
ResultContainer<IdType> findSimilarByClusterOverlap(
const std::unordered_map<IdType, std::vector<db::ClusterId>>& profileMap,
IdType queryId,
const std::vector<db::ClusterId>& queryClusters,
std::size_t maxCount)
{
const std::unordered_set<db::ClusterId> 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<IdType> res;
res.reserve(resultCount);
for (std::size_t i{}; i < resultCount; ++i)
res.push_back({ .id = overlapCounts[i].first, .distance = {} });
return res;
}
} // namespace
std::unique_ptr<IEngine> createClustersEngine(db::IDb& db)
{
return std::make_unique<ClusterEngine>(db);
}
TrackContainer ClusterEngine::findSimilarTracks(const std::vector<TrackId>& trackIds, std::size_t maxCount) const
ClusterEngine::ClusterEngine(db::IDb& db)
: _db{ db }
{
constexpr float sameReleaseWeight{ 0.5F };
constexpr float sameArtistWeight{ 0.5F };
_trackEvaluator.addHardConstraint(std::make_unique<DuplicateTrackConstraint>());
_trackEvaluator.addSoftConstraint(std::make_unique<SameReleaseConstraint>(_trackMetadata), sameReleaseWeight);
_trackEvaluator.addSoftConstraint(std::make_unique<SameArtistConstraint>(_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() };
buildTrackMetadata(session);
buildTrackClusters(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...");
db::Release::find(session, db::Release::FindParameters{}, [&](const db::Release::pointer& release) {
db::Track::FindParameters params;
params.setRelease(release->getId());
for (const db::TrackId trackId : db::Track::findIds(session, params).results)
_trackMetadata[trackId].releaseId = release->getId();
});
db::Artist::find(session, db::Artist::FindParameters{}, [&](const db::Artist::pointer& artist) {
std::unordered_set<db::TrackId> artistTrackIds;
{
db::Track::FindParameters params;
params.setArtist(artist->getId(), { db::TrackArtistLinkType::Artist });
for (const db::TrackId trackId : db::Track::findIds(session, params).results)
artistTrackIds.insert(trackId);
}
{
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<const db::TrackId> trackIds, std::size_t maxCount) const
{
LMS_SCOPED_TRACE_DETAILED("ClustersEngine", "Find similar tracks");
if (maxCount == 0 || trackIds.empty())
return {};
std::unordered_set<db::ClusterId> 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<db::TrackId> 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<db::TrackId> candidates;
candidates.reserve(candidateCount);
for (const auto& [trackId, count] : overlapCounts)
candidates.push_back(trackId);
std::vector<db::TrackId> seeds{ std::cbegin(trackIds), std::cend(trackIds) };
return greedySelect(std::move(candidates), std::move(seeds), maxCount);
}
TrackResults ClusterEngine::greedySelect(std::vector<db::TrackId> candidates, std::vector<db::TrackId> selectedTracks, std::size_t maxCount) const
{
selectedTracks.reserve(selectedTracks.size() + maxCount);
TrackResults res;
res.reserve(maxCount);
while (res.size() < maxCount && !candidates.empty())
{
std::optional<std::size_t> bestIdx;
float bestScore{ std::numeric_limits<float>::max() };
for (std::size_t i{}; i < candidates.size(); ++i)
{
const TrackCandidateContext context{
.candidateTrackId = candidates[i],
.selectedTracks = selectedTracks,
};
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], .distance = {} });
selectedTracks.push_back(candidates[*bestIdx]);
candidates.erase(std::begin(candidates) + static_cast<std::ptrdiff_t>(*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 {};
Session& dbSession{ _db.getTLSSession() };
auto transaction{ dbSession.createReadTransaction() };
auto similarTrackIds{ Track::findSimilarTrackIds(dbSession, trackIds, Range{ 0, maxCount }) };
return std::move(similarTrackIds.results);
}
TrackContainer ClusterEngine::findSimilarTracksFromTrackList(TrackListId tracklistId, std::size_t maxCount) const
{
TrackContainer res;
if (maxCount == 0)
return res;
std::vector<db::TrackId> trackIds;
{
Session& dbSession{ _db.getTLSSession() };
db::Session& dbSession{ _db.getTLSSession() };
auto transaction{ dbSession.createReadTransaction() };
const TrackList::pointer trackList{ TrackList::find(dbSession, tracklistId) };
const db::TrackList::pointer trackList{ db::TrackList::find(dbSession, tracklistId) };
if (!trackList)
return res;
return {};
const auto tracks{ trackList->getSimilarTracks(0, maxCount) };
res.reserve(tracks.size());
std::transform(std::cbegin(tracks), std::cend(tracks), std::back_inserter(res), [](const auto& track) { return track->getId(); });
trackIds = trackList->getTrackIds();
}
return res;
if (trackIds.empty())
return {};
return findSimilarTracks(trackIds, maxCount);
}
ReleaseContainer ClusterEngine::getSimilarReleases(ReleaseId releaseId, std::size_t maxCount) const
ReleaseResults ClusterEngine::findSimilarReleases(db::ReleaseId releaseId, std::size_t maxCount) const
{
ReleaseContainer res;
if (maxCount == 0)
return res;
LMS_SCOPED_TRACE_DETAILED("ClustersEngine", "Find similar releases");
{
Session& dbSession{ _db.getTLSSession() };
auto transaction{ dbSession.createReadTransaction() };
auto release{ Release::find(dbSession, releaseId) };
if (!release)
return res;
const auto releases{ release->getSimilarReleases(0, maxCount) };
res.reserve(releases.size());
std::transform(std::cbegin(releases), std::cend(releases), std::back_inserter(res), [](const auto& release) { return release->getId(); });
}
return res;
}
ArtistContainer ClusterEngine::getSimilarArtists(ArtistId artistId, core::EnumSet<TrackArtistLinkType> artistLinkTypes, std::size_t maxCount) const
{
if (maxCount == 0)
return {};
Session& dbSession{ _db.getTLSSession() };
const auto queryIt{ _releaseClusters.find(releaseId) };
if (queryIt == _releaseClusters.cend() || queryIt->second.empty())
return {};
return findSimilarByClusterOverlap<db::ReleaseId>(_releaseClusters, releaseId, queryIt->second, maxCount);
}
ArtistResults ClusterEngine::findSimilarArtists(db::ArtistId artistId, core::EnumSet<db::TrackArtistLinkType> 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<db::ArtistId>(_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<db::TrackId>{ .id = startTrackId, .distance = {} } };
db::Session& dbSession{ _db.getTLSSession() };
auto transaction{ dbSession.createReadTransaction() };
auto artist{ Artist::find(dbSession, artistId) };
if (!artist)
const auto startTrack{ db::Track::find(dbSession, startTrackId) };
const auto endTrack{ db::Track::find(dbSession, endTrackId) };
if (!startTrack || !endTrack)
return {};
auto similarArtistIds{ artist->findSimilarArtistIds(artistLinkTypes, Range{ 0, maxCount }) };
return std::move(similarArtistIds.results);
TrackResults res;
res.reserve(std::min<std::size_t>(maxCount, 2));
res.push_back({ .id = startTrackId, .distance = {} });
if (maxCount > 1)
res.push_back({ .id = endTrackId, .distance = {} });
return res;
}
} // namespace lms::recommendation
#undef LOG