trulens-benchmark 1.0.1a1__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- trulens_benchmark-1.0.1a1/PKG-INFO +24 -0
- trulens_benchmark-1.0.1a1/README.md +1 -0
- trulens_benchmark-1.0.1a1/pyproject.toml +34 -0
- trulens_benchmark-1.0.1a1/trulens/benchmark/__init__.py +15 -0
- trulens_benchmark-1.0.1a1/trulens/benchmark/answer_relevance_benchmark_small.ipynb +247 -0
- trulens_benchmark-1.0.1a1/trulens/benchmark/benchmark_frameworks/__init__.py +0 -0
- trulens_benchmark-1.0.1a1/trulens/benchmark/benchmark_frameworks/eval_as_recommendation.py +124 -0
- trulens_benchmark-1.0.1a1/trulens/benchmark/comprehensiveness_benchmark.ipynb +318 -0
- trulens_benchmark-1.0.1a1/trulens/benchmark/context_relevance_benchmark.ipynb +342 -0
- trulens_benchmark-1.0.1a1/trulens/benchmark/context_relevance_benchmark_calibration.ipynb +307 -0
- trulens_benchmark-1.0.1a1/trulens/benchmark/context_relevance_benchmark_small.ipynb +235 -0
- trulens_benchmark-1.0.1a1/trulens/benchmark/groundedness_benchmark.ipynb +235 -0
- trulens_benchmark-1.0.1a1/trulens/benchmark/groundedness_benchmark_abstentions.ipynb +676 -0
- trulens_benchmark-1.0.1a1/trulens/benchmark/groundedness_benchmark_snowflake_arctic.ipynb +274 -0
- trulens_benchmark-1.0.1a1/trulens/benchmark/test_cases.py +291 -0
|
@@ -0,0 +1,24 @@
|
|
|
1
|
+
Metadata-Version: 2.1
|
|
2
|
+
Name: trulens-benchmark
|
|
3
|
+
Version: 1.0.1a1
|
|
4
|
+
Summary: Library to systematically track and evaluate LLM based applications.
|
|
5
|
+
Home-page: https://trulens.org/
|
|
6
|
+
License: MIT
|
|
7
|
+
Author: Snowflake Inc.
|
|
8
|
+
Author-email: ml-observability-wg-dl@snowflake.com
|
|
9
|
+
Requires-Python: >=3.9,<4.0
|
|
10
|
+
Classifier: Development Status :: 3 - Alpha
|
|
11
|
+
Classifier: License :: OSI Approved :: MIT License
|
|
12
|
+
Classifier: Operating System :: OS Independent
|
|
13
|
+
Classifier: Programming Language :: Python :: 3
|
|
14
|
+
Classifier: Programming Language :: Python :: 3.9
|
|
15
|
+
Classifier: Programming Language :: Python :: 3.10
|
|
16
|
+
Classifier: Programming Language :: Python :: 3.11
|
|
17
|
+
Classifier: Programming Language :: Python :: 3.12
|
|
18
|
+
Requires-Dist: trulens-core (>=1.0.0,<2.0.0)
|
|
19
|
+
Project-URL: Documentation, https://trulens.org/trulens/getting_started/
|
|
20
|
+
Project-URL: Repository, https://github.com/truera/trulens
|
|
21
|
+
Description-Content-Type: text/markdown
|
|
22
|
+
|
|
23
|
+
# trulens-benchmark
|
|
24
|
+
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
# trulens-benchmark
|
|
@@ -0,0 +1,34 @@
|
|
|
1
|
+
[build-system]
|
|
2
|
+
build-backend = "poetry.core.masonry.api"
|
|
3
|
+
requires = [
|
|
4
|
+
"poetry-core",
|
|
5
|
+
]
|
|
6
|
+
|
|
7
|
+
[tool.poetry]
|
|
8
|
+
name = "trulens-benchmark"
|
|
9
|
+
version = "1.0.1a1"
|
|
10
|
+
description = "Library to systematically track and evaluate LLM based applications."
|
|
11
|
+
authors = [
|
|
12
|
+
"Snowflake Inc. <ml-observability-wg-dl@snowflake.com>",
|
|
13
|
+
]
|
|
14
|
+
license = "MIT"
|
|
15
|
+
readme = "README.md"
|
|
16
|
+
packages = [
|
|
17
|
+
{ include = "trulens" },
|
|
18
|
+
]
|
|
19
|
+
homepage = "https://trulens.org/"
|
|
20
|
+
documentation = "https://trulens.org/trulens/getting_started/"
|
|
21
|
+
repository = "https://github.com/truera/trulens"
|
|
22
|
+
classifiers = [
|
|
23
|
+
"Programming Language :: Python :: 3",
|
|
24
|
+
"Operating System :: OS Independent",
|
|
25
|
+
"Development Status :: 3 - Alpha",
|
|
26
|
+
"License :: OSI Approved :: MIT License",
|
|
27
|
+
]
|
|
28
|
+
|
|
29
|
+
[tool.poetry.dependencies]
|
|
30
|
+
python = "^3.9"
|
|
31
|
+
trulens-core = { version = "^1.0.0", allow-prereleases = true }
|
|
32
|
+
|
|
33
|
+
[tool.poetry.group.dev.dependencies]
|
|
34
|
+
trulens-core = { path = "../core", develop = true }
|
|
@@ -0,0 +1,15 @@
|
|
|
1
|
+
"""
|
|
2
|
+
!!! note "Additional Dependency Required"
|
|
3
|
+
|
|
4
|
+
To use this module, you must have the `trulens-benchmark` package installed.
|
|
5
|
+
|
|
6
|
+
```bash
|
|
7
|
+
pip install trulens-benchmark
|
|
8
|
+
```
|
|
9
|
+
"""
|
|
10
|
+
|
|
11
|
+
from importlib.metadata import version
|
|
12
|
+
|
|
13
|
+
from trulens.core.utils.imports import safe_importlib_package_name
|
|
14
|
+
|
|
15
|
+
__version__ = version(safe_importlib_package_name(__package__ or __name__))
|
|
@@ -0,0 +1,247 @@
|
|
|
1
|
+
{
|
|
2
|
+
"cells": [
|
|
3
|
+
{
|
|
4
|
+
"attachments": {},
|
|
5
|
+
"cell_type": "markdown",
|
|
6
|
+
"metadata": {},
|
|
7
|
+
"source": [
|
|
8
|
+
"# 📓 Answer Relevance Feedback Evaluation\n",
|
|
9
|
+
"In many ways, feedbacks can be thought of as LLM apps themselves. Given text,\n",
|
|
10
|
+
"they return some result. Thinking in this way, we can use _TruLens_ to evaluate\n",
|
|
11
|
+
"and track our feedback quality. We can even do this for different models (e.g.\n",
|
|
12
|
+
"gpt-3.5 and gpt-4) or prompting schemes (such as chain-of-thought reasoning).\n",
|
|
13
|
+
"\n",
|
|
14
|
+
"This notebook follows an evaluation of a set of test cases. You are encouraged\n",
|
|
15
|
+
"to run this on your own and even expand the test cases to evaluate performance\n",
|
|
16
|
+
"on test cases applicable to your scenario or domain."
|
|
17
|
+
]
|
|
18
|
+
},
|
|
19
|
+
{
|
|
20
|
+
"cell_type": "code",
|
|
21
|
+
"execution_count": null,
|
|
22
|
+
"metadata": {},
|
|
23
|
+
"outputs": [],
|
|
24
|
+
"source": [
|
|
25
|
+
"# Import relevance feedback function\n",
|
|
26
|
+
"from test_cases import answer_relevance_golden_set\n",
|
|
27
|
+
"from trulens.core import Feedback\n",
|
|
28
|
+
"from trulens.core import Select\n",
|
|
29
|
+
"from trulens.core import Tru\n",
|
|
30
|
+
"from trulens.core import TruBasicApp\n",
|
|
31
|
+
"from trulens.feedback import GroundTruthAgreement\n",
|
|
32
|
+
"from trulens.providers.litellm import LiteLLM\n",
|
|
33
|
+
"from trulens.providers.openai import OpenAI\n",
|
|
34
|
+
"\n",
|
|
35
|
+
"Tru().reset_database()"
|
|
36
|
+
]
|
|
37
|
+
},
|
|
38
|
+
{
|
|
39
|
+
"cell_type": "code",
|
|
40
|
+
"execution_count": null,
|
|
41
|
+
"metadata": {},
|
|
42
|
+
"outputs": [],
|
|
43
|
+
"source": [
|
|
44
|
+
"import os\n",
|
|
45
|
+
"\n",
|
|
46
|
+
"os.environ[\"OPENAI_API_KEY\"] = \"...\"\n",
|
|
47
|
+
"os.environ[\"COHERE_API_KEY\"] = \"...\"\n",
|
|
48
|
+
"os.environ[\"HUGGINGFACE_API_KEY\"] = \"...\"\n",
|
|
49
|
+
"os.environ[\"ANTHROPIC_API_KEY\"] = \"...\"\n",
|
|
50
|
+
"os.environ[\"TOGETHERAI_API_KEY\"] = \"...\""
|
|
51
|
+
]
|
|
52
|
+
},
|
|
53
|
+
{
|
|
54
|
+
"cell_type": "code",
|
|
55
|
+
"execution_count": null,
|
|
56
|
+
"metadata": {},
|
|
57
|
+
"outputs": [],
|
|
58
|
+
"source": [
|
|
59
|
+
"# GPT 3.5\n",
|
|
60
|
+
"turbo = OpenAI(model_engine=\"gpt-3.5-turbo\")\n",
|
|
61
|
+
"\n",
|
|
62
|
+
"\n",
|
|
63
|
+
"def wrapped_relevance_turbo(input, output):\n",
|
|
64
|
+
" return turbo.relevance(input, output)\n",
|
|
65
|
+
"\n",
|
|
66
|
+
"\n",
|
|
67
|
+
"# GPT 4\n",
|
|
68
|
+
"gpt4 = OpenAI(model_engine=\"gpt-4\")\n",
|
|
69
|
+
"\n",
|
|
70
|
+
"\n",
|
|
71
|
+
"def wrapped_relevance_gpt4(input, output):\n",
|
|
72
|
+
" return gpt4.relevance(input, output)\n",
|
|
73
|
+
"\n",
|
|
74
|
+
"\n",
|
|
75
|
+
"# Cohere\n",
|
|
76
|
+
"command_nightly = LiteLLM(model_engine=\"cohere/command-nightly\")\n",
|
|
77
|
+
"\n",
|
|
78
|
+
"\n",
|
|
79
|
+
"def wrapped_relevance_command_nightly(input, output):\n",
|
|
80
|
+
" return command_nightly.relevance(input, output)\n",
|
|
81
|
+
"\n",
|
|
82
|
+
"\n",
|
|
83
|
+
"# Anthropic\n",
|
|
84
|
+
"claude_1 = LiteLLM(model_engine=\"claude-instant-1\")\n",
|
|
85
|
+
"\n",
|
|
86
|
+
"\n",
|
|
87
|
+
"def wrapped_relevance_claude1(input, output):\n",
|
|
88
|
+
" return claude_1.relevance(input, output)\n",
|
|
89
|
+
"\n",
|
|
90
|
+
"\n",
|
|
91
|
+
"claude_2 = LiteLLM(model_engine=\"claude-2\")\n",
|
|
92
|
+
"\n",
|
|
93
|
+
"\n",
|
|
94
|
+
"def wrapped_relevance_claude2(input, output):\n",
|
|
95
|
+
" return claude_2.relevance(input, output)\n",
|
|
96
|
+
"\n",
|
|
97
|
+
"\n",
|
|
98
|
+
"# Meta\n",
|
|
99
|
+
"llama_2_13b = LiteLLM(\n",
|
|
100
|
+
" model_engine=\"together_ai/togethercomputer/Llama-2-7B-32K-Instruct\"\n",
|
|
101
|
+
")\n",
|
|
102
|
+
"\n",
|
|
103
|
+
"\n",
|
|
104
|
+
"def wrapped_relevance_llama2(input, output):\n",
|
|
105
|
+
" return llama_2_13b.relevance(input, output)"
|
|
106
|
+
]
|
|
107
|
+
},
|
|
108
|
+
{
|
|
109
|
+
"attachments": {},
|
|
110
|
+
"cell_type": "markdown",
|
|
111
|
+
"metadata": {},
|
|
112
|
+
"source": [
|
|
113
|
+
"Here we'll set up our golden set as a set of prompts, responses and expected\n",
|
|
114
|
+
"scores stored in `test_cases.py`. Then, our numeric_difference method will look\n",
|
|
115
|
+
"up the expected score for each prompt/response pair by **exact match**. After\n",
|
|
116
|
+
"looking up the expected score, we will then take the L1 difference between the\n",
|
|
117
|
+
"actual score and expected score."
|
|
118
|
+
]
|
|
119
|
+
},
|
|
120
|
+
{
|
|
121
|
+
"cell_type": "code",
|
|
122
|
+
"execution_count": null,
|
|
123
|
+
"metadata": {},
|
|
124
|
+
"outputs": [],
|
|
125
|
+
"source": [
|
|
126
|
+
"# Create a Feedback object using the numeric_difference method of the\n",
|
|
127
|
+
"# ground_truth object\n",
|
|
128
|
+
"ground_truth = GroundTruthAgreement(\n",
|
|
129
|
+
" answer_relevance_golden_set, provider=OpenAI()\n",
|
|
130
|
+
")\n",
|
|
131
|
+
"\n",
|
|
132
|
+
"# Call the numeric_difference method with app and record and aggregate to get\n",
|
|
133
|
+
"# the mean absolute error\n",
|
|
134
|
+
"f_mae = (\n",
|
|
135
|
+
" Feedback(ground_truth.mae, name=\"Mean Absolute Error\")\n",
|
|
136
|
+
" .on(Select.Record.calls[0].args.args[0])\n",
|
|
137
|
+
" .on(Select.Record.calls[0].args.args[1])\n",
|
|
138
|
+
" .on_output()\n",
|
|
139
|
+
")"
|
|
140
|
+
]
|
|
141
|
+
},
|
|
142
|
+
{
|
|
143
|
+
"cell_type": "code",
|
|
144
|
+
"execution_count": null,
|
|
145
|
+
"metadata": {},
|
|
146
|
+
"outputs": [],
|
|
147
|
+
"source": [
|
|
148
|
+
"tru_wrapped_relevance_turbo = TruBasicApp(\n",
|
|
149
|
+
" wrapped_relevance_turbo,\n",
|
|
150
|
+
" app_id=\"answer relevance gpt-3.5-turbo\",\n",
|
|
151
|
+
" feedbacks=[f_mae],\n",
|
|
152
|
+
")\n",
|
|
153
|
+
"\n",
|
|
154
|
+
"tru_wrapped_relevance_gpt4 = TruBasicApp(\n",
|
|
155
|
+
" wrapped_relevance_gpt4, app_id=\"answer relevance gpt-4\", feedbacks=[f_mae]\n",
|
|
156
|
+
")\n",
|
|
157
|
+
"\n",
|
|
158
|
+
"tru_wrapped_relevance_commandnightly = TruBasicApp(\n",
|
|
159
|
+
" wrapped_relevance_command_nightly,\n",
|
|
160
|
+
" app_id=\"answer relevance Command-Nightly\",\n",
|
|
161
|
+
" feedbacks=[f_mae],\n",
|
|
162
|
+
")\n",
|
|
163
|
+
"\n",
|
|
164
|
+
"tru_wrapped_relevance_claude1 = TruBasicApp(\n",
|
|
165
|
+
" wrapped_relevance_claude1,\n",
|
|
166
|
+
" app_id=\"answer relevance Claude 1\",\n",
|
|
167
|
+
" feedbacks=[f_mae],\n",
|
|
168
|
+
")\n",
|
|
169
|
+
"\n",
|
|
170
|
+
"tru_wrapped_relevance_claude2 = TruBasicApp(\n",
|
|
171
|
+
" wrapped_relevance_claude2,\n",
|
|
172
|
+
" app_id=\"answer relevance Claude 2\",\n",
|
|
173
|
+
" feedbacks=[f_mae],\n",
|
|
174
|
+
")\n",
|
|
175
|
+
"\n",
|
|
176
|
+
"tru_wrapped_relevance_llama2 = TruBasicApp(\n",
|
|
177
|
+
" wrapped_relevance_llama2,\n",
|
|
178
|
+
" app_id=\"answer relevance Llama-2-13b\",\n",
|
|
179
|
+
" feedbacks=[f_mae],\n",
|
|
180
|
+
")"
|
|
181
|
+
]
|
|
182
|
+
},
|
|
183
|
+
{
|
|
184
|
+
"cell_type": "code",
|
|
185
|
+
"execution_count": null,
|
|
186
|
+
"metadata": {},
|
|
187
|
+
"outputs": [],
|
|
188
|
+
"source": [
|
|
189
|
+
"for i in range(len(answer_relevance_golden_set)):\n",
|
|
190
|
+
" prompt = answer_relevance_golden_set[i][\"query\"]\n",
|
|
191
|
+
" response = answer_relevance_golden_set[i][\"response\"]\n",
|
|
192
|
+
"\n",
|
|
193
|
+
" with tru_wrapped_relevance_turbo as recording:\n",
|
|
194
|
+
" tru_wrapped_relevance_turbo.app(prompt, response)\n",
|
|
195
|
+
"\n",
|
|
196
|
+
" with tru_wrapped_relevance_gpt4 as recording:\n",
|
|
197
|
+
" tru_wrapped_relevance_gpt4.app(prompt, response)\n",
|
|
198
|
+
"\n",
|
|
199
|
+
" with tru_wrapped_relevance_commandnightly as recording:\n",
|
|
200
|
+
" tru_wrapped_relevance_commandnightly.app(prompt, response)\n",
|
|
201
|
+
"\n",
|
|
202
|
+
" with tru_wrapped_relevance_claude1 as recording:\n",
|
|
203
|
+
" tru_wrapped_relevance_claude1.app(prompt, response)\n",
|
|
204
|
+
"\n",
|
|
205
|
+
" with tru_wrapped_relevance_claude2 as recording:\n",
|
|
206
|
+
" tru_wrapped_relevance_claude2.app(prompt, response)\n",
|
|
207
|
+
"\n",
|
|
208
|
+
" with tru_wrapped_relevance_llama2 as recording:\n",
|
|
209
|
+
" tru_wrapped_relevance_llama2.app(prompt, response)"
|
|
210
|
+
]
|
|
211
|
+
},
|
|
212
|
+
{
|
|
213
|
+
"cell_type": "code",
|
|
214
|
+
"execution_count": null,
|
|
215
|
+
"metadata": {},
|
|
216
|
+
"outputs": [],
|
|
217
|
+
"source": [
|
|
218
|
+
"Tru().get_leaderboard(app_ids=[]).sort_values(by=\"Mean Absolute Error\")"
|
|
219
|
+
]
|
|
220
|
+
}
|
|
221
|
+
],
|
|
222
|
+
"metadata": {
|
|
223
|
+
"kernelspec": {
|
|
224
|
+
"display_name": "Python 3.11.4 ('agents')",
|
|
225
|
+
"language": "python",
|
|
226
|
+
"name": "python3"
|
|
227
|
+
},
|
|
228
|
+
"language_info": {
|
|
229
|
+
"codemirror_mode": {
|
|
230
|
+
"name": "ipython",
|
|
231
|
+
"version": 3
|
|
232
|
+
},
|
|
233
|
+
"file_extension": ".py",
|
|
234
|
+
"mimetype": "text/x-python",
|
|
235
|
+
"name": "python",
|
|
236
|
+
"nbconvert_exporter": "python",
|
|
237
|
+
"pygments_lexer": "ipython3"
|
|
238
|
+
},
|
|
239
|
+
"vscode": {
|
|
240
|
+
"interpreter": {
|
|
241
|
+
"hash": "7d153714b979d5e6d08dd8ec90712dd93bff2c9b6c1f0c118169738af3430cd4"
|
|
242
|
+
}
|
|
243
|
+
}
|
|
244
|
+
},
|
|
245
|
+
"nbformat": 4,
|
|
246
|
+
"nbformat_minor": 2
|
|
247
|
+
}
|
|
File without changes
|
|
@@ -0,0 +1,124 @@
|
|
|
1
|
+
import logging
|
|
2
|
+
import time
|
|
3
|
+
|
|
4
|
+
import numpy as np
|
|
5
|
+
from sklearn.metrics import ndcg_score
|
|
6
|
+
|
|
7
|
+
log = logging.getLogger(__name__)
|
|
8
|
+
"""score passages with feedback function, retrying if feedback function fails.
|
|
9
|
+
Args: df: dataframe with columns 'query_id', 'query', 'passage', 'is_selected'
|
|
10
|
+
feedback_func: function that takes query and passage as input and returns a score
|
|
11
|
+
backoff_time: time to wait between retries
|
|
12
|
+
n: number of samples to estimate conditional probabilities of feedback_func's scores
|
|
13
|
+
"""
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
def score_passages(
|
|
17
|
+
df,
|
|
18
|
+
feedback_func_name,
|
|
19
|
+
feedback_func,
|
|
20
|
+
backoff_time=0.5,
|
|
21
|
+
n=5,
|
|
22
|
+
temperature=0.0,
|
|
23
|
+
):
|
|
24
|
+
grouped = df.groupby("query_id")
|
|
25
|
+
scores = []
|
|
26
|
+
true_relevance = []
|
|
27
|
+
|
|
28
|
+
for name, group in grouped:
|
|
29
|
+
query_scores = []
|
|
30
|
+
query_relevance = []
|
|
31
|
+
for _, row in group.iterrows():
|
|
32
|
+
sampled_score = None
|
|
33
|
+
if feedback_func_name == "TruEra" or n == 1:
|
|
34
|
+
sampled_score = feedback_func(
|
|
35
|
+
row["query"], row["passage"], temperature
|
|
36
|
+
) # hard-coded for now, we don't need to sample for TruEra BERT-based model
|
|
37
|
+
time.sleep(backoff_time)
|
|
38
|
+
else:
|
|
39
|
+
sampled_scores = []
|
|
40
|
+
for _ in range(n):
|
|
41
|
+
sampled_scores.append(
|
|
42
|
+
feedback_func(row["query"], row["passage"], temperature)
|
|
43
|
+
)
|
|
44
|
+
time.sleep(backoff_time)
|
|
45
|
+
sampled_score = sum(sampled_scores) / len(sampled_scores)
|
|
46
|
+
query_scores.append(sampled_score)
|
|
47
|
+
query_relevance.append(row["is_selected"])
|
|
48
|
+
# print(f"Feedback avg score for query {name} is {sampled_score}, is_selected is {row['is_selected']}")
|
|
49
|
+
|
|
50
|
+
print(
|
|
51
|
+
f"Feedback function {name} scored {len(query_scores)} out of {len(group)} passages."
|
|
52
|
+
)
|
|
53
|
+
scores.append(query_scores)
|
|
54
|
+
true_relevance.append(query_relevance)
|
|
55
|
+
|
|
56
|
+
return scores, true_relevance
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
def compute_ndcg(scores, true_relevance):
|
|
60
|
+
ndcg_values = [
|
|
61
|
+
ndcg_score([true], [pred]) for true, pred in zip(true_relevance, scores)
|
|
62
|
+
]
|
|
63
|
+
return np.mean(ndcg_values)
|
|
64
|
+
|
|
65
|
+
|
|
66
|
+
def compute_ece(scores, true_relevance, n_bins=10):
|
|
67
|
+
ece = 0
|
|
68
|
+
for bin in np.linspace(0, 1, n_bins):
|
|
69
|
+
bin_scores = []
|
|
70
|
+
bin_truth = []
|
|
71
|
+
for score_list, truth_list in zip(scores, true_relevance):
|
|
72
|
+
for score, truth in zip(score_list, truth_list):
|
|
73
|
+
if bin <= score < bin + 1 / n_bins:
|
|
74
|
+
bin_scores.append(score)
|
|
75
|
+
bin_truth.append(truth)
|
|
76
|
+
|
|
77
|
+
if bin_scores:
|
|
78
|
+
bin_avg_confidence = np.mean(bin_scores)
|
|
79
|
+
bin_accuracy = np.mean(bin_truth)
|
|
80
|
+
ece += (
|
|
81
|
+
np.abs(bin_avg_confidence - bin_accuracy)
|
|
82
|
+
* len(bin_scores)
|
|
83
|
+
/ sum(map(len, scores))
|
|
84
|
+
)
|
|
85
|
+
|
|
86
|
+
return ece
|
|
87
|
+
|
|
88
|
+
|
|
89
|
+
def precision_at_k(scores, true_relevance, k):
|
|
90
|
+
sorted_scores = sorted(scores, reverse=True)
|
|
91
|
+
kth_score = sorted_scores[min(k - 1, len(scores) - 1)]
|
|
92
|
+
|
|
93
|
+
# Indices of items with scores >= kth highest score
|
|
94
|
+
top_k_indices = [i for i, score in enumerate(scores) if score >= kth_score]
|
|
95
|
+
|
|
96
|
+
# Calculate precision
|
|
97
|
+
true_positives = sum(np.take(true_relevance, top_k_indices))
|
|
98
|
+
return true_positives / len(top_k_indices) if top_k_indices else 0
|
|
99
|
+
|
|
100
|
+
|
|
101
|
+
def recall_at_k(scores, true_relevance, k):
|
|
102
|
+
"""
|
|
103
|
+
Calculate the recall at K.
|
|
104
|
+
|
|
105
|
+
Parameters:
|
|
106
|
+
true_relevance (list of int): List of binary values indicating relevance (1 for relevant, 0 for not).
|
|
107
|
+
scores (list of float): List of scores assigned by the model.
|
|
108
|
+
k (int): Number of top items to consider for calculating recall.
|
|
109
|
+
|
|
110
|
+
Returns:
|
|
111
|
+
float: Recall at K.
|
|
112
|
+
"""
|
|
113
|
+
sorted_scores = sorted(scores, reverse=True)
|
|
114
|
+
kth_score = sorted_scores[min(k - 1, len(scores) - 1)]
|
|
115
|
+
|
|
116
|
+
# Indices of items with scores >= kth highest score
|
|
117
|
+
top_k_indices = [i for i, score in enumerate(scores) if score >= kth_score]
|
|
118
|
+
|
|
119
|
+
# Calculate recall
|
|
120
|
+
relevant_indices = np.where(true_relevance)[0]
|
|
121
|
+
hits = sum(idx in top_k_indices for idx in relevant_indices)
|
|
122
|
+
total_relevant = sum(true_relevance)
|
|
123
|
+
|
|
124
|
+
return hits / total_relevant if total_relevant > 0 else 0
|