Files
NexQuant/rdagent/oai/llm_utils.py
T

45 lines
1.5 KiB
Python
Raw Normal View History

2024-06-12 15:12:11 +08:00
from __future__ import annotations
from typing import Any, Type
2024-05-21 22:48:41 +08:00
import numpy as np
2024-06-14 12:59:44 +08:00
from rdagent.core.utils import import_class
from rdagent.oai.backend.base import APIBackend as BaseAPIBackend
from rdagent.oai.llm_conf import LLM_SETTINGS
from rdagent.utils import md5_hash # for compatible with previous import
2024-05-21 22:48:41 +08:00
2024-06-12 15:12:11 +08:00
def calculate_embedding_distance_between_str_list(
source_str_list: list[str],
target_str_list: list[str],
2024-06-12 15:12:11 +08:00
) -> list[list[float]]:
if not source_str_list or not target_str_list:
2024-05-21 22:48:41 +08:00
return [[]]
embeddings = APIBackend().create_embedding(source_str_list + target_str_list)
source_embeddings = embeddings[: len(source_str_list)]
target_embeddings = embeddings[len(source_str_list) :]
2024-05-21 22:48:41 +08:00
source_embeddings_np = np.array(source_embeddings)
target_embeddings_np = np.array(target_embeddings)
source_embeddings_np = source_embeddings_np / np.linalg.norm(source_embeddings_np, axis=1, keepdims=True)
target_embeddings_np = target_embeddings_np / np.linalg.norm(target_embeddings_np, axis=1, keepdims=True)
similarity_matrix = np.dot(source_embeddings_np, target_embeddings_np.T)
return similarity_matrix.tolist() # type: ignore[no-any-return]
def get_api_backend(*args: Any, **kwargs: Any) -> BaseAPIBackend: # TODO: import it from base.py
"""
get llm api backend based on settings dynamically.
"""
api_backend_cls: Type[BaseAPIBackend] = import_class(LLM_SETTINGS.backend)
return api_backend_cls(*args, **kwargs)
# Alias
APIBackend = get_api_backend