Files
silo-server/internal/recommendations/diversity.go

216 lines
5.4 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package recommendations
import (
"math"
"sort"
)
const minGenreCapRetainedFraction = 0.5
// applyMMR re-ranks candidates using Maximal Marginal Relevance to balance
// relevance against diversity. It selects up to limit items from candidates,
// choosing each successive item to maximize:
//
// score = λ × normalizedRelevance - (1-λ) × maxSimilarityToSelected
//
// Candidates without an entry in embeddings are scored by relevance alone.
func applyMMR(candidates []ScoredItem, embeddings map[string][]float32, lambda float64, limit int) []ScoredItem {
if len(candidates) == 0 || limit <= 0 {
return nil
}
if limit > len(candidates) {
limit = len(candidates)
}
// Find max relevance score for normalization.
maxScore := candidates[0].Score
for _, c := range candidates[1:] {
if c.Score > maxScore {
maxScore = c.Score
}
}
if maxScore == 0 {
maxScore = 1 // avoid division by zero
}
// Track which candidates remain available and which have been selected.
remaining := make([]int, len(candidates))
for i := range remaining {
remaining[i] = i
}
selected := make([]ScoredItem, 0, limit)
selectedEmbeddings := make([][]float32, 0, limit)
// First pick: highest relevance score.
bestIdx := 0
for i, ri := range remaining {
if candidates[ri].Score > candidates[remaining[bestIdx]].Score {
bestIdx = i
}
}
first := remaining[bestIdx]
selected = append(selected, candidates[first])
if emb, ok := embeddings[candidates[first].MediaItemID]; ok {
selectedEmbeddings = append(selectedEmbeddings, emb)
}
remaining = append(remaining[:bestIdx], remaining[bestIdx+1:]...)
// Subsequent picks via MMR scoring.
for len(selected) < limit && len(remaining) > 0 {
bestMMRIdx := -1
bestMMRScore := math.Inf(-1)
for i, ri := range remaining {
normalizedRelevance := candidates[ri].Score / maxScore
candidateEmb := embeddings[candidates[ri].MediaItemID]
var maxSim float64
if candidateEmb != nil && len(selectedEmbeddings) > 0 {
for _, selEmb := range selectedEmbeddings {
sim := cosineSimilarity(candidateEmb, selEmb)
if sim > maxSim {
maxSim = sim
}
}
}
var mmrScore float64
if candidateEmb == nil {
// No embedding available — use relevance only.
mmrScore = normalizedRelevance
} else {
mmrScore = lambda*normalizedRelevance - (1-lambda)*maxSim
}
if mmrScore > bestMMRScore {
bestMMRScore = mmrScore
bestMMRIdx = i
}
}
pick := remaining[bestMMRIdx]
selected = append(selected, candidates[pick])
if emb, ok := embeddings[candidates[pick].MediaItemID]; ok {
selectedEmbeddings = append(selectedEmbeddings, emb)
}
remaining = append(remaining[:bestMMRIdx], remaining[bestMMRIdx+1:]...)
}
return selected
}
// applyGenreCap enforces that no single genre exceeds maxPct of the result set.
// When a genre is over-represented, the lowest-scored items from that genre are
// removed until the genre falls within the cap.
func applyGenreCap(items []ScoredItem, genres map[string][]string, maxPct float64) []ScoredItem {
if len(items) == 0 || maxPct <= 0 || maxPct >= 1 {
return items
}
minRetained := int(math.Ceil(float64(len(items)) * minGenreCapRetainedFraction))
if minRetained < 1 {
minRetained = 1
}
// Sort a copy by score descending so removals take lowest-scored first.
sorted := make([]ScoredItem, len(items))
copy(sorted, items)
sort.Slice(sorted, func(i, j int) bool {
if sorted[i].Score != sorted[j].Score {
return sorted[i].Score > sorted[j].Score
}
return sorted[i].MediaItemID < sorted[j].MediaItemID
})
// Iteratively remove until all genres are within cap. Re-check after each
// removal because the total count changes.
for {
// Count items per genre.
genreCounts := make(map[string]int)
for _, item := range sorted {
for _, genre := range genres[item.MediaItemID] {
if genre != "" {
genreCounts[genre]++
}
}
}
// Find a genre that exceeds the cap.
total := len(sorted)
maxAllowed := int(math.Floor(maxPct * float64(total)))
if maxAllowed < 1 {
maxAllowed = 1
}
overGenre := ""
for g, count := range genreCounts {
if count > maxAllowed {
if overGenre == "" ||
count > genreCounts[overGenre] ||
(count == genreCounts[overGenre] && g < overGenre) {
overGenre = g
}
}
}
if overGenre == "" {
break // all genres within cap
}
if len(sorted) <= minRetained {
break
}
// Remove the lowest-scored item of the over-represented genre.
// Items are sorted descending, so scan from the end.
removed := false
for i := len(sorted) - 1; i >= 0; i-- {
if itemHasGenre(genres[sorted[i].MediaItemID], overGenre) {
sorted = append(sorted[:i], sorted[i+1:]...)
removed = true
break
}
}
if !removed {
break
}
}
return sorted
}
func itemHasGenre(genres []string, target string) bool {
for _, genre := range genres {
if genre == target {
return true
}
}
return false
}
// cosineSimilarity computes the cosine similarity between two float32 vectors.
// Returns 0 if either vector is nil or empty.
func cosineSimilarity(a, b []float32) float64 {
if len(a) == 0 || len(b) == 0 {
return 0
}
n := len(a)
if len(b) < n {
n = len(b)
}
var dot, normA, normB float64
for i := 0; i < n; i++ {
ai := float64(a[i])
bi := float64(b[i])
dot += ai * bi
normA += ai * ai
normB += bi * bi
}
if normA == 0 || normB == 0 {
return 0
}
return dot / (math.Sqrt(normA) * math.Sqrt(normB))
}