This commit is contained in:
Timothy Jaeryang Baek
2026-03-24 17:03:08 -05:00
parent e1cdd7e4fe
commit d738044f47
+15 -4
View File
@@ -721,7 +721,7 @@ def generate_ollama_batch_embeddings(
) -> Optional[list[list[float]]]:
try:
log.debug(f'generate_ollama_batch_embeddings:model {model} batch size: {len(texts)}')
json_data = {'input': texts, 'model': model}
json_data = {'input': texts, 'model': model, 'truncate': True}
if isinstance(RAG_EMBEDDING_PREFIX_FIELD_NAME, str) and isinstance(prefix, str):
json_data[RAG_EMBEDDING_PREFIX_FIELD_NAME] = prefix
@@ -737,7 +737,9 @@ def generate_ollama_batch_embeddings(
headers=headers,
json=json_data,
)
r.raise_for_status()
if r.status_code != 200:
error_detail = r.json().get('error', r.text)
raise Exception(f'Ollama embed error ({r.status_code}): {error_detail}')
data = r.json()
if 'embeddings' in data:
@@ -759,7 +761,7 @@ async def agenerate_ollama_batch_embeddings(
) -> Optional[list[list[float]]]:
try:
log.debug(f'agenerate_ollama_batch_embeddings:model {model} batch size: {len(texts)}')
form_data = {'input': texts, 'model': model}
form_data = {'input': texts, 'model': model, 'truncate': True}
if isinstance(RAG_EMBEDDING_PREFIX_FIELD_NAME, str) and isinstance(prefix, str):
form_data[RAG_EMBEDDING_PREFIX_FIELD_NAME] = prefix
@@ -779,7 +781,10 @@ async def agenerate_ollama_batch_embeddings(
json=form_data,
ssl=AIOHTTP_CLIENT_SESSION_SSL,
) as r:
r.raise_for_status()
if r.status != 200:
error_data = await r.json()
error_detail = error_data.get('error', str(error_data))
raise Exception(f'Ollama embed error ({r.status}): {error_detail}')
data = await r.json()
if 'embeddings' in data:
return data['embeddings']
@@ -901,11 +906,15 @@ async def generate_embeddings(
'user': user,
}
)
if embeddings is None:
return None
return embeddings[0] if isinstance(text, str) else embeddings
elif engine == 'openai':
embeddings = await agenerate_openai_batch_embeddings(
model, text if isinstance(text, list) else [text], url, key, prefix, user
)
if embeddings is None:
return None
return embeddings[0] if isinstance(text, str) else embeddings
elif engine == 'azure_openai':
azure_api_version = kwargs.get('azure_api_version', '')
@@ -918,6 +927,8 @@ async def generate_embeddings(
prefix,
user,
)
if embeddings is None:
return None
return embeddings[0] if isinstance(text, str) else embeddings