fix(avatar): batch knowledge embedding requests
This commit is contained in:
@@ -51,7 +51,14 @@ def embed(texts):
|
|||||||
if api_url:
|
if api_url:
|
||||||
api_key = os.getenv("EMBEDDING_API_KEY", "")
|
api_key = os.getenv("EMBEDDING_API_KEY", "")
|
||||||
model = os.getenv("EMBEDDING_MODEL", "text-embedding-3-small")
|
model = os.getenv("EMBEDDING_MODEL", "text-embedding-3-small")
|
||||||
payload = json.dumps({"input": texts, "model": model}).encode("utf-8")
|
try:
|
||||||
|
batch_size = max(1, int(os.getenv("EMBEDDING_BATCH_SIZE", "10")))
|
||||||
|
except ValueError:
|
||||||
|
batch_size = 10
|
||||||
|
embeddings = []
|
||||||
|
for start in range(0, len(texts), batch_size):
|
||||||
|
batch = texts[start:start + batch_size]
|
||||||
|
payload = json.dumps({"input": batch, "model": model}).encode("utf-8")
|
||||||
req = urllib.request.Request(
|
req = urllib.request.Request(
|
||||||
api_url,
|
api_url,
|
||||||
data=payload,
|
data=payload,
|
||||||
@@ -66,7 +73,10 @@ def embed(texts):
|
|||||||
items = data["data"]
|
items = data["data"]
|
||||||
if items and "index" in items[0]:
|
if items and "index" in items[0]:
|
||||||
items = sorted(items, key=lambda x: x["index"])
|
items = sorted(items, key=lambda x: x["index"])
|
||||||
return [item["embedding"] for item in items]
|
if len(items) != len(batch):
|
||||||
|
raise ValueError("embedding response count does not match request")
|
||||||
|
embeddings.extend(item["embedding"] for item in items)
|
||||||
|
return embeddings
|
||||||
return _hash_embedding(texts)
|
return _hash_embedding(texts)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1,10 +1,26 @@
|
|||||||
|
import json
|
||||||
import os
|
import os
|
||||||
import tempfile
|
import tempfile
|
||||||
import unittest
|
import unittest
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
import embeddings
|
import embeddings
|
||||||
|
|
||||||
|
|
||||||
|
class FakeResponse:
|
||||||
|
def __init__(self, payload):
|
||||||
|
self.payload = payload
|
||||||
|
|
||||||
|
def __enter__(self):
|
||||||
|
return self
|
||||||
|
|
||||||
|
def __exit__(self, *_):
|
||||||
|
return None
|
||||||
|
|
||||||
|
def read(self):
|
||||||
|
return json.dumps(self.payload).encode("utf-8")
|
||||||
|
|
||||||
|
|
||||||
class TextExtractionTests(unittest.TestCase):
|
class TextExtractionTests(unittest.TestCase):
|
||||||
def write_text(self, suffix, content):
|
def write_text(self, suffix, content):
|
||||||
handle = tempfile.NamedTemporaryFile(suffix=suffix, delete=False)
|
handle = tempfile.NamedTemporaryFile(suffix=suffix, delete=False)
|
||||||
@@ -28,5 +44,33 @@ class TextExtractionTests(unittest.TestCase):
|
|||||||
embeddings.extract_text(path, ".csv")
|
embeddings.extract_text(path, ".csv")
|
||||||
|
|
||||||
|
|
||||||
|
class RemoteEmbeddingTests(unittest.TestCase):
|
||||||
|
def test_large_input_is_split_into_provider_safe_batches(self):
|
||||||
|
texts = [f"chunk-{index}" for index in range(14)]
|
||||||
|
batch_sizes = []
|
||||||
|
|
||||||
|
def fake_urlopen(request, timeout):
|
||||||
|
self.assertEqual(timeout, 30)
|
||||||
|
payload = json.loads(request.data.decode("utf-8"))
|
||||||
|
batch_sizes.append(len(payload["input"]))
|
||||||
|
return FakeResponse({
|
||||||
|
"data": [
|
||||||
|
{"index": index, "embedding": [float(text.split("-")[1])]}
|
||||||
|
for index, text in enumerate(payload["input"])
|
||||||
|
]
|
||||||
|
})
|
||||||
|
|
||||||
|
with patch.dict(os.environ, {
|
||||||
|
"EMBEDDING_API_URL": "https://embedding.example/v1/embeddings",
|
||||||
|
"EMBEDDING_API_KEY": "test-key",
|
||||||
|
"EMBEDDING_MODEL": "text-embedding-v4",
|
||||||
|
"EMBEDDING_BATCH_SIZE": "10",
|
||||||
|
}), patch("embeddings.urllib.request.urlopen", side_effect=fake_urlopen):
|
||||||
|
result = embeddings.embed(texts)
|
||||||
|
|
||||||
|
self.assertEqual(batch_sizes, [10, 4])
|
||||||
|
self.assertEqual(result, [[float(index)] for index in range(14)])
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user