hyperforge-nucliadb-agentic 1.0.0.post64__py3-none-any.whl
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.
- hyperforge_nucliadb_agentic/__init__.py +3 -0
- hyperforge_nucliadb_agentic/agent.py +1642 -0
- hyperforge_nucliadb_agentic/ask/__init__.py +5 -0
- hyperforge_nucliadb_agentic/ask/audit.py +439 -0
- hyperforge_nucliadb_agentic/ask/exceptions.py +50 -0
- hyperforge_nucliadb_agentic/ask/lifespan.py +21 -0
- hyperforge_nucliadb_agentic/ask/model.py +1299 -0
- hyperforge_nucliadb_agentic/ask/predict.py +431 -0
- hyperforge_nucliadb_agentic/ask/predict_models.py +78 -0
- hyperforge_nucliadb_agentic/ask/search/__init__.py +0 -0
- hyperforge_nucliadb_agentic/ask/search/ask.py +1182 -0
- hyperforge_nucliadb_agentic/ask/search/graph_strategy.py +1138 -0
- hyperforge_nucliadb_agentic/ask/search/highlight.py +93 -0
- hyperforge_nucliadb_agentic/ask/search/hydrator.py +29 -0
- hyperforge_nucliadb_agentic/ask/search/metrics.py +112 -0
- hyperforge_nucliadb_agentic/ask/search/parsers/__init__.py +0 -0
- hyperforge_nucliadb_agentic/ask/search/parsers/ask.py +70 -0
- hyperforge_nucliadb_agentic/ask/search/parsers/fetcher.py +192 -0
- hyperforge_nucliadb_agentic/ask/search/parsers/find.py +729 -0
- hyperforge_nucliadb_agentic/ask/search/prompt.py +1298 -0
- hyperforge_nucliadb_agentic/ask/search/rank_fusion.py +157 -0
- hyperforge_nucliadb_agentic/ask/search/rerankers.py +161 -0
- hyperforge_nucliadb_agentic/ask/search/retrieval.py +750 -0
- hyperforge_nucliadb_agentic/ask/search/rpc.py +192 -0
- hyperforge_nucliadb_agentic/ask/settings.py +9 -0
- hyperforge_nucliadb_agentic/ask/utils/ids.py +188 -0
- hyperforge_nucliadb_agentic/ask/utils/proto.py +6 -0
- hyperforge_nucliadb_agentic/ask/utils/responses.py +6 -0
- hyperforge_nucliadb_agentic/ask/utils/text_blocks.py +51 -0
- hyperforge_nucliadb_agentic/config.py +47 -0
- hyperforge_nucliadb_agentic/internal_driver.py +87 -0
- hyperforge_nucliadb_agentic/py.typed +0 -0
- hyperforge_nucliadb_agentic-1.0.0.post64.dist-info/METADATA +24 -0
- hyperforge_nucliadb_agentic-1.0.0.post64.dist-info/RECORD +35 -0
- hyperforge_nucliadb_agentic-1.0.0.post64.dist-info/WHEEL +4 -0
|
@@ -0,0 +1,93 @@
|
|
|
1
|
+
import logging
|
|
2
|
+
import re
|
|
3
|
+
import string
|
|
4
|
+
|
|
5
|
+
logger = logging.getLogger(__name__)
|
|
6
|
+
PRE_WORD = string.punctuation + " "
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
def highlight_paragraph(
|
|
10
|
+
text: str, words: list[str] | None = None, ematches: list[str] | None = None
|
|
11
|
+
) -> str:
|
|
12
|
+
"""
|
|
13
|
+
Highlight `text` with <mark></mark> tags around the words in `words` and `ematches`.
|
|
14
|
+
|
|
15
|
+
Parameters:
|
|
16
|
+
- text: The text to highlight.
|
|
17
|
+
- words: A list of words to highlight.
|
|
18
|
+
- ematches: A list of exact matches to highlight.
|
|
19
|
+
|
|
20
|
+
Returns:
|
|
21
|
+
- The highlighted text.
|
|
22
|
+
"""
|
|
23
|
+
REGEX_TEMPLATE = r"(^|\s)({text})(\s|$)"
|
|
24
|
+
text_lower = text.lower()
|
|
25
|
+
|
|
26
|
+
marks = [0] * (len(text_lower) + 1)
|
|
27
|
+
ematches = ematches or []
|
|
28
|
+
for quote in ematches:
|
|
29
|
+
quote_regex = REGEX_TEMPLATE.format(text=re.escape(quote.lower()))
|
|
30
|
+
try:
|
|
31
|
+
for match in re.finditer(quote_regex, text_lower):
|
|
32
|
+
start, end = match.span(2)
|
|
33
|
+
marks[start] = 1
|
|
34
|
+
marks[end] = 2
|
|
35
|
+
except re.error:
|
|
36
|
+
logger.warning(
|
|
37
|
+
f"Regex errors while highlighting text. Regex: {quote_regex}"
|
|
38
|
+
)
|
|
39
|
+
continue
|
|
40
|
+
|
|
41
|
+
words = words or []
|
|
42
|
+
for word in words:
|
|
43
|
+
word_regex = REGEX_TEMPLATE.format(text=re.escape(word.lower()))
|
|
44
|
+
try:
|
|
45
|
+
for match in re.finditer(word_regex, text_lower):
|
|
46
|
+
start, end = match.span(2)
|
|
47
|
+
if marks[start] == 0 and marks[end] == 0:
|
|
48
|
+
marks[start] = 1
|
|
49
|
+
marks[end] = 2
|
|
50
|
+
except re.error:
|
|
51
|
+
logger.warning(f"Regex errors while highlighting text. Regex: {word_regex}")
|
|
52
|
+
continue
|
|
53
|
+
|
|
54
|
+
new_text = ""
|
|
55
|
+
actual = 0
|
|
56
|
+
mod = 0
|
|
57
|
+
skip = False
|
|
58
|
+
|
|
59
|
+
length = len(text)
|
|
60
|
+
|
|
61
|
+
for index, pos in enumerate(marks):
|
|
62
|
+
if skip:
|
|
63
|
+
skip = False
|
|
64
|
+
continue
|
|
65
|
+
if (index - mod) >= length:
|
|
66
|
+
char_pos = ""
|
|
67
|
+
else:
|
|
68
|
+
begining = True
|
|
69
|
+
if index > 0 and text[index - mod - 1] not in PRE_WORD:
|
|
70
|
+
begining = False
|
|
71
|
+
char_pos = text[index - mod]
|
|
72
|
+
if text[index - mod].lower() != text_lower[index]:
|
|
73
|
+
# May be incorrect positioning due to unicode lower
|
|
74
|
+
mod += 1
|
|
75
|
+
skip = True
|
|
76
|
+
if pos == 1 and actual == 0 and begining:
|
|
77
|
+
new_text += "<mark>"
|
|
78
|
+
new_text += char_pos
|
|
79
|
+
actual = 1
|
|
80
|
+
elif pos == 2 and actual == 1:
|
|
81
|
+
new_text += "</mark>"
|
|
82
|
+
new_text += char_pos
|
|
83
|
+
actual = 0
|
|
84
|
+
elif pos == 1 and actual > 0:
|
|
85
|
+
new_text += char_pos
|
|
86
|
+
actual += 1
|
|
87
|
+
elif pos == 2 and actual > 1:
|
|
88
|
+
new_text += char_pos
|
|
89
|
+
actual -= 1
|
|
90
|
+
else:
|
|
91
|
+
new_text += char_pos
|
|
92
|
+
|
|
93
|
+
return new_text
|
|
@@ -0,0 +1,29 @@
|
|
|
1
|
+
from nucliadb_models.common import FieldTypeName
|
|
2
|
+
from nucliadb_models.resource import ExtractedDataTypeName
|
|
3
|
+
from nucliadb_models.search import ResourceProperties
|
|
4
|
+
from pydantic import BaseModel
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class ResourceHydrationOptions(BaseModel):
|
|
8
|
+
"""
|
|
9
|
+
Options for hydrating resources.
|
|
10
|
+
"""
|
|
11
|
+
|
|
12
|
+
show: list[ResourceProperties] = []
|
|
13
|
+
extracted: list[ExtractedDataTypeName] = []
|
|
14
|
+
field_type_filter: list[FieldTypeName] = []
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
class TextBlockHydrationOptions(BaseModel):
|
|
18
|
+
"""
|
|
19
|
+
Options for hydrating text blocks (aka paragraphs).
|
|
20
|
+
"""
|
|
21
|
+
|
|
22
|
+
# whether to highlight the text block with `<mark>...</mark>` tags or not
|
|
23
|
+
highlight: bool = False
|
|
24
|
+
|
|
25
|
+
# list of exact matches to highlight
|
|
26
|
+
ematches: list[str] | None = None
|
|
27
|
+
|
|
28
|
+
# If true, only hydrate the text block if its text is not already populated
|
|
29
|
+
only_hydrate_empty: bool = False
|
|
@@ -0,0 +1,112 @@
|
|
|
1
|
+
import contextlib
|
|
2
|
+
import time
|
|
3
|
+
from typing import Any
|
|
4
|
+
|
|
5
|
+
from nucliadb_telemetry import metrics
|
|
6
|
+
|
|
7
|
+
buckets = [
|
|
8
|
+
0.005,
|
|
9
|
+
0.01,
|
|
10
|
+
0.025,
|
|
11
|
+
0.05,
|
|
12
|
+
0.075,
|
|
13
|
+
0.1,
|
|
14
|
+
0.25,
|
|
15
|
+
0.5,
|
|
16
|
+
0.75,
|
|
17
|
+
1.0,
|
|
18
|
+
2.5,
|
|
19
|
+
5.0,
|
|
20
|
+
7.5,
|
|
21
|
+
10.0,
|
|
22
|
+
30.0,
|
|
23
|
+
60.0,
|
|
24
|
+
metrics.INF,
|
|
25
|
+
]
|
|
26
|
+
|
|
27
|
+
generative_first_chunk_histogram = metrics.Histogram(
|
|
28
|
+
name="generative_reasoning_first_chunk",
|
|
29
|
+
buckets=buckets,
|
|
30
|
+
)
|
|
31
|
+
reasoning_first_chunk_histogram = metrics.Histogram(
|
|
32
|
+
name="generative_first_chunk",
|
|
33
|
+
buckets=buckets,
|
|
34
|
+
)
|
|
35
|
+
rag_histogram = metrics.Histogram(
|
|
36
|
+
name="rag",
|
|
37
|
+
labels={"step": ""},
|
|
38
|
+
buckets=buckets,
|
|
39
|
+
)
|
|
40
|
+
|
|
41
|
+
MetricsData = dict[str, int | float]
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
class Metrics:
|
|
45
|
+
def __init__(self: "Metrics", id: str):
|
|
46
|
+
self.id = id
|
|
47
|
+
self.child_spans: list[Metrics] = []
|
|
48
|
+
self._metrics: MetricsData = {}
|
|
49
|
+
|
|
50
|
+
@contextlib.contextmanager
|
|
51
|
+
def time(self, step: str):
|
|
52
|
+
start_time = time.monotonic()
|
|
53
|
+
try:
|
|
54
|
+
yield
|
|
55
|
+
finally:
|
|
56
|
+
elapsed = time.monotonic() - start_time
|
|
57
|
+
self._metrics[step] = elapsed
|
|
58
|
+
rag_histogram.observe(elapsed, labels={"step": step})
|
|
59
|
+
|
|
60
|
+
def child_span(self, id: str) -> "Metrics":
|
|
61
|
+
child_span = Metrics(id)
|
|
62
|
+
self.child_spans.append(child_span)
|
|
63
|
+
return child_span
|
|
64
|
+
|
|
65
|
+
def set(self, key: str, value: int | float):
|
|
66
|
+
self._metrics[key] = value
|
|
67
|
+
|
|
68
|
+
def get(self, key: str) -> int | float | None:
|
|
69
|
+
return self._metrics.get(key)
|
|
70
|
+
|
|
71
|
+
def to_dict(self) -> MetricsData:
|
|
72
|
+
return self._metrics
|
|
73
|
+
|
|
74
|
+
def dump(self) -> dict[str, Any]:
|
|
75
|
+
result = {}
|
|
76
|
+
for child in self.child_spans:
|
|
77
|
+
result.update(child.dump())
|
|
78
|
+
result[self.id] = self.to_dict()
|
|
79
|
+
return result
|
|
80
|
+
|
|
81
|
+
def __getitem__(self, key: str) -> int | float:
|
|
82
|
+
return self._metrics[key]
|
|
83
|
+
|
|
84
|
+
|
|
85
|
+
class AskMetrics(Metrics):
|
|
86
|
+
def __init__(self: "AskMetrics"):
|
|
87
|
+
super().__init__(id="ask")
|
|
88
|
+
self.global_start = time.monotonic()
|
|
89
|
+
self.first_chunk_yielded_at: float | None = None
|
|
90
|
+
self.first_reasoning_chunk_yielded_at: float | None = None
|
|
91
|
+
|
|
92
|
+
def record_first_chunk_yielded(self):
|
|
93
|
+
self.first_chunk_yielded_at = time.monotonic()
|
|
94
|
+
generative_first_chunk_histogram.observe(
|
|
95
|
+
self.first_chunk_yielded_at - self.global_start
|
|
96
|
+
)
|
|
97
|
+
|
|
98
|
+
def record_first_reasoning_chunk_yielded(self):
|
|
99
|
+
self.first_reasoning_chunk_yielded_at = time.monotonic()
|
|
100
|
+
reasoning_first_chunk_histogram.observe(
|
|
101
|
+
self.first_reasoning_chunk_yielded_at - self.global_start
|
|
102
|
+
)
|
|
103
|
+
|
|
104
|
+
def get_first_chunk_time(self) -> float | None:
|
|
105
|
+
if self.first_chunk_yielded_at is None:
|
|
106
|
+
return None
|
|
107
|
+
return self.first_chunk_yielded_at - self.global_start
|
|
108
|
+
|
|
109
|
+
def get_first_reasoning_chunk_time(self) -> float | None:
|
|
110
|
+
if self.first_reasoning_chunk_yielded_at is None:
|
|
111
|
+
return None
|
|
112
|
+
return self.first_reasoning_chunk_yielded_at - self.global_start
|
|
File without changes
|
|
@@ -0,0 +1,70 @@
|
|
|
1
|
+
from pydantic import BaseModel
|
|
2
|
+
from typing_extensions import assert_never
|
|
3
|
+
|
|
4
|
+
from hyperforge_nucliadb_agentic.ask.model import AskRequest, MaxTokens
|
|
5
|
+
from hyperforge_nucliadb_agentic.ask.search.parsers.fetcher import (
|
|
6
|
+
Fetcher,
|
|
7
|
+
)
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class Generation(BaseModel):
|
|
11
|
+
"""Request field related with response generation"""
|
|
12
|
+
|
|
13
|
+
use_visual_llm: bool
|
|
14
|
+
max_context_tokens: int
|
|
15
|
+
max_answer_tokens: int | None
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
class _AskParser:
|
|
19
|
+
def __init__(self, kbid: str, item: AskRequest, fetcher: Fetcher):
|
|
20
|
+
self.kbid = kbid
|
|
21
|
+
self.item = item
|
|
22
|
+
self.fetcher = fetcher
|
|
23
|
+
|
|
24
|
+
async def parse(self) -> Generation:
|
|
25
|
+
use_visual_llm = await self.fetcher.get_visual_llm_enabled()
|
|
26
|
+
|
|
27
|
+
if self.item.max_tokens is None:
|
|
28
|
+
max_tokens = None
|
|
29
|
+
elif isinstance(self.item.max_tokens, int):
|
|
30
|
+
max_tokens = MaxTokens(
|
|
31
|
+
context=None,
|
|
32
|
+
answer=self.item.max_tokens,
|
|
33
|
+
)
|
|
34
|
+
elif isinstance(self.item.max_tokens, MaxTokens):
|
|
35
|
+
max_tokens = self.item.max_tokens
|
|
36
|
+
else: # pragma: no cover
|
|
37
|
+
assert_never(self.item.max_tokens)
|
|
38
|
+
|
|
39
|
+
max_context_tokens = await self.fetcher.get_max_context_tokens(max_tokens)
|
|
40
|
+
max_answer_tokens = self.fetcher.get_max_answer_tokens(max_tokens)
|
|
41
|
+
|
|
42
|
+
return Generation(
|
|
43
|
+
use_visual_llm=use_visual_llm,
|
|
44
|
+
max_context_tokens=max_context_tokens,
|
|
45
|
+
max_answer_tokens=max_answer_tokens,
|
|
46
|
+
)
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
async def parse_ask(
|
|
50
|
+
kbid: str,
|
|
51
|
+
item: AskRequest,
|
|
52
|
+
*,
|
|
53
|
+
fetcher: Fetcher | None = None,
|
|
54
|
+
) -> Generation:
|
|
55
|
+
fetcher = fetcher or fetcher_for_ask(kbid, item)
|
|
56
|
+
parser = _AskParser(kbid, item, fetcher)
|
|
57
|
+
return await parser.parse()
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
def fetcher_for_ask(kbid: str, item: AskRequest) -> Fetcher:
|
|
61
|
+
return Fetcher(
|
|
62
|
+
kbid=kbid,
|
|
63
|
+
query=item.query,
|
|
64
|
+
user_vector=None,
|
|
65
|
+
vectorset=item.vectorset,
|
|
66
|
+
rephrase=item.rephrase,
|
|
67
|
+
rephrase_prompt=None,
|
|
68
|
+
generative_model=item.generative_model,
|
|
69
|
+
query_image=item.query_image,
|
|
70
|
+
)
|
|
@@ -0,0 +1,192 @@
|
|
|
1
|
+
from google.protobuf.json_format import ParseDict
|
|
2
|
+
from nucliadb_models.internal.predict import QueryInfo
|
|
3
|
+
from nucliadb_protos import knowledgebox_pb2, utils_pb2
|
|
4
|
+
from nucliadb_sdk import NucliaDBAsync
|
|
5
|
+
|
|
6
|
+
from hyperforge_nucliadb_agentic.ask import logger
|
|
7
|
+
from hyperforge_nucliadb_agentic.ask.exceptions import (
|
|
8
|
+
InvalidQueryError,
|
|
9
|
+
)
|
|
10
|
+
from hyperforge_nucliadb_agentic.ask.model import Image, MaxTokens
|
|
11
|
+
from hyperforge_nucliadb_agentic.ask.predict import (
|
|
12
|
+
SendToPredictError,
|
|
13
|
+
convert_relations,
|
|
14
|
+
get_predict,
|
|
15
|
+
)
|
|
16
|
+
from hyperforge_nucliadb_agentic.ask.predict_models import QueryModel
|
|
17
|
+
from hyperforge_nucliadb_agentic.ask.search import rpc
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
class Fetcher:
|
|
21
|
+
"""This class is an encapsulation of data gathering across different parts of
|
|
22
|
+
the system. Given the user query input, it aims to be as efficient as
|
|
23
|
+
possible removing redundant expensive calls to other parts of the system. An
|
|
24
|
+
instance of a fetcher caches it's results and it's thought to be used in the
|
|
25
|
+
context of a single request.
|
|
26
|
+
|
|
27
|
+
*DO NOT* use this as a global object!
|
|
28
|
+
|
|
29
|
+
"""
|
|
30
|
+
|
|
31
|
+
def __init__(
|
|
32
|
+
self,
|
|
33
|
+
kbid: str,
|
|
34
|
+
*,
|
|
35
|
+
query: str,
|
|
36
|
+
user_vector: list[float] | None,
|
|
37
|
+
vectorset: str | None,
|
|
38
|
+
rephrase: bool,
|
|
39
|
+
rephrase_prompt: str | None,
|
|
40
|
+
generative_model: str | None,
|
|
41
|
+
query_image: Image | None,
|
|
42
|
+
):
|
|
43
|
+
self.kbid = kbid
|
|
44
|
+
self.query = query
|
|
45
|
+
self.user_vector = user_vector
|
|
46
|
+
self.user_vectorset = vectorset
|
|
47
|
+
self.user_vectorset_validated = False
|
|
48
|
+
self.rephrase = rephrase
|
|
49
|
+
self.rephrase_prompt = rephrase_prompt
|
|
50
|
+
self.generative_model = generative_model
|
|
51
|
+
self.query_image = query_image
|
|
52
|
+
|
|
53
|
+
self._query_info: QueryInfo | None = None
|
|
54
|
+
self._vectorset: str | None = None
|
|
55
|
+
|
|
56
|
+
async def query_information(self) -> QueryInfo:
|
|
57
|
+
if self._query_info is None:
|
|
58
|
+
predict = get_predict()
|
|
59
|
+
item = QueryModel(
|
|
60
|
+
text=self.query,
|
|
61
|
+
semantic_models=[self.user_vectorset] if self.user_vectorset else None,
|
|
62
|
+
generative_model=self.generative_model,
|
|
63
|
+
rephrase=self.rephrase,
|
|
64
|
+
rephrase_prompt=self.rephrase_prompt,
|
|
65
|
+
query_image=self.query_image,
|
|
66
|
+
)
|
|
67
|
+
try:
|
|
68
|
+
self._query_info = await predict.query(self.kbid, item)
|
|
69
|
+
except TimeoutError as exc:
|
|
70
|
+
raise SendToPredictError(
|
|
71
|
+
"timeout while requesting Predict API /query"
|
|
72
|
+
) from exc
|
|
73
|
+
|
|
74
|
+
return self._query_info
|
|
75
|
+
|
|
76
|
+
# Retrieval
|
|
77
|
+
|
|
78
|
+
async def get_rephrased_query(self) -> str | None:
|
|
79
|
+
query_info = await self.query_information()
|
|
80
|
+
return query_info.rephrased_query
|
|
81
|
+
|
|
82
|
+
def get_cached_rephrased_query(self) -> str | None:
|
|
83
|
+
if self._query_info is None:
|
|
84
|
+
return None
|
|
85
|
+
return self._query_info.rephrased_query
|
|
86
|
+
|
|
87
|
+
async def get_detected_entities(self) -> list[utils_pb2.RelationNode]:
|
|
88
|
+
query_info = await self.query_information()
|
|
89
|
+
if query_info.entities is not None:
|
|
90
|
+
detected_entities = convert_relations(query_info.entities.model_dump())
|
|
91
|
+
else:
|
|
92
|
+
detected_entities = []
|
|
93
|
+
return detected_entities
|
|
94
|
+
|
|
95
|
+
async def get_semantic_min_score(self) -> float | None:
|
|
96
|
+
query_info = await self.query_information()
|
|
97
|
+
vectorset = await self.get_vectorset()
|
|
98
|
+
return query_info.semantic_thresholds.get(vectorset, None)
|
|
99
|
+
|
|
100
|
+
async def get_vectorset(self) -> str:
|
|
101
|
+
if self._vectorset is None:
|
|
102
|
+
if self.user_vectorset is not None:
|
|
103
|
+
self._vectorset = self.user_vectorset
|
|
104
|
+
else:
|
|
105
|
+
# when it's not provided, we get the default from Predict API
|
|
106
|
+
query_info = await self.query_information()
|
|
107
|
+
if query_info.sentence is None or len(query_info.sentence.vectors) == 0:
|
|
108
|
+
logger.error(
|
|
109
|
+
"Asking for a vectorset but /query didn't return one",
|
|
110
|
+
extra={"kbid": self.kbid},
|
|
111
|
+
)
|
|
112
|
+
raise SendToPredictError(
|
|
113
|
+
"Predict API didn't return a sentence vectorset"
|
|
114
|
+
)
|
|
115
|
+
# vectors field is enforced by the data model to have at least one key
|
|
116
|
+
for vectorset in query_info.sentence.vectors.keys():
|
|
117
|
+
self._vectorset = vectorset
|
|
118
|
+
break
|
|
119
|
+
assert self._vectorset is not None
|
|
120
|
+
return self._vectorset
|
|
121
|
+
|
|
122
|
+
async def get_query_vector(self) -> list[float]:
|
|
123
|
+
if self.user_vector is not None:
|
|
124
|
+
return self.user_vector
|
|
125
|
+
|
|
126
|
+
query_info = await self.query_information()
|
|
127
|
+
if query_info.sentence is None:
|
|
128
|
+
logger.error(
|
|
129
|
+
"Asking for a semantic query vector but /query didn't return a sentence",
|
|
130
|
+
extra={"kbid": self.kbid},
|
|
131
|
+
)
|
|
132
|
+
raise SendToPredictError(
|
|
133
|
+
"Predict API didn't return a sentence for semantic search"
|
|
134
|
+
)
|
|
135
|
+
|
|
136
|
+
vectorset = await self.get_vectorset()
|
|
137
|
+
if vectorset not in query_info.sentence.vectors:
|
|
138
|
+
logger.error(
|
|
139
|
+
"Predict is not responding with a valid query nucliadb vectorset",
|
|
140
|
+
extra={
|
|
141
|
+
"kbid": self.kbid,
|
|
142
|
+
"vectorset": vectorset,
|
|
143
|
+
"predict_vectorsets": ",".join(query_info.sentence.vectors.keys()),
|
|
144
|
+
},
|
|
145
|
+
)
|
|
146
|
+
raise SendToPredictError(
|
|
147
|
+
"Predict API didn't return the requested vectorset"
|
|
148
|
+
)
|
|
149
|
+
|
|
150
|
+
query_vector = query_info.sentence.vectors[vectorset]
|
|
151
|
+
return query_vector
|
|
152
|
+
|
|
153
|
+
async def get_classification_labels(
|
|
154
|
+
self, reader_sdk: NucliaDBAsync
|
|
155
|
+
) -> knowledgebox_pb2.Labels:
|
|
156
|
+
labelsets = await rpc.labelsets(reader_sdk=reader_sdk, kbid=self.kbid)
|
|
157
|
+
|
|
158
|
+
# TODO(decoupled-ask): remove this conversion and refactor code to use API models instead of protobuf
|
|
159
|
+
kb_labels = knowledgebox_pb2.Labels()
|
|
160
|
+
for labelset, labels in labelsets.labelsets.items():
|
|
161
|
+
ParseDict(labels.model_dump(), kb_labels.labelset[labelset])
|
|
162
|
+
|
|
163
|
+
return kb_labels
|
|
164
|
+
|
|
165
|
+
# Generative
|
|
166
|
+
|
|
167
|
+
async def get_visual_llm_enabled(self) -> bool:
|
|
168
|
+
query_info = await self.query_information()
|
|
169
|
+
if query_info is None:
|
|
170
|
+
raise SendToPredictError("Error while using predict's query endpoint")
|
|
171
|
+
|
|
172
|
+
return query_info.visual_llm
|
|
173
|
+
|
|
174
|
+
async def get_max_context_tokens(self, max_tokens: MaxTokens | None) -> int:
|
|
175
|
+
query_info = await self.query_information()
|
|
176
|
+
if query_info is None:
|
|
177
|
+
raise SendToPredictError("Error while using predict's query endpoint")
|
|
178
|
+
|
|
179
|
+
model_max = query_info.max_context
|
|
180
|
+
if max_tokens is not None and max_tokens.context is not None:
|
|
181
|
+
if max_tokens.context > model_max:
|
|
182
|
+
raise InvalidQueryError(
|
|
183
|
+
"max_tokens.context",
|
|
184
|
+
f"Max context tokens is higher than the model's limit of {model_max}",
|
|
185
|
+
)
|
|
186
|
+
return max_tokens.context
|
|
187
|
+
return model_max
|
|
188
|
+
|
|
189
|
+
def get_max_answer_tokens(self, max_tokens: MaxTokens | None) -> int | None:
|
|
190
|
+
if max_tokens is not None and max_tokens.answer is not None:
|
|
191
|
+
return max_tokens.answer
|
|
192
|
+
return None
|