def similarity_search(query_embed, target_embeddings, content_list, k=5, threshold=0.05, temperature=0.5): """ Perform similarity search over embeddings and return top k results. """ # Calculate similarities similarities = torch.matmul(query_embed, target_embeddings.T) # Rescale similarities via softmax scores = torch.nn.functional.softmax(similarities/temperature, dim=1) # Get sorted indices and scores sorted_indices = scores.argsort(descending=True)[0] sorted_scores = scores[0][sorted_indices] # Filter by threshold and get top k filtered_indices = [ idx.item() for idx, score in zip(sorted_indices, sorted_scores) if score.item() >= threshold ][:k] # Get corresponding content items and scores top_results = [content_list[i] for i in filtered_indices] result_scores = [scores[0][i].item() for i in filtered_indices] return top_results, result_scores