2024-06-12 15:12:11 +08:00
|
|
|
from __future__ import annotations
|
|
|
|
|
|
2025-02-13 15:16:18 +08:00
|
|
|
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
|
|
|
|
2025-02-13 15:16:18 +08:00
|
|
|
from rdagent.core.utils import import_class
|
|
|
|
|
from rdagent.oai.backend.base import APIBackend as BaseAPIBackend
|
2024-10-14 17:34:09 +08:00
|
|
|
from rdagent.oai.llm_conf import LLM_SETTINGS
|
2025-02-13 15:16:18 +08:00
|
|
|
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(
|
2024-06-18 11:50:03 +08:00
|
|
|
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 [[]]
|
|
|
|
|
|
2024-10-08 19:34:19 +08:00
|
|
|
embeddings = APIBackend().create_embedding(source_str_list + target_str_list)
|
|
|
|
|
|
2024-06-18 11:50:03 +08:00
|
|
|
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)
|
|
|
|
|
|
2025-01-17 22:53:05 +08:00
|
|
|
return similarity_matrix.tolist() # type: ignore[no-any-return]
|
2025-02-13 15:16:18 +08:00
|
|
|
|
|
|
|
|
|
|
|
|
|
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
|