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

216 lines
5.4 KiB
Go
Raw Normal View History

2026-05-22 20:26:11 -04:00
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))
}