From 5ade3f1b70d4a5cf5821dcd2005dba90cc0a848e Mon Sep 17 00:00:00 2001 From: frayle-ons <194791647+frayle-ons@users.noreply.github.com> Date: Fri, 4 Sep 2026 17:05:11 +0100 Subject: [PATCH] rework gcp transform method to include kwargs in embed content call --- src/classifai/vectorisers/gcp.py | 17 +++++++++++++---- 1 file changed, 13 insertions(+), 4 deletions(-) diff --git a/src/classifai/vectorisers/gcp.py b/src/classifai/vectorisers/gcp.py index 7cbb093..9323353 100644 --- a/src/classifai/vectorisers/gcp.py +++ b/src/classifai/vectorisers/gcp.py @@ -103,12 +103,13 @@ def __init__( context={"vectoriser": "gcp", "cause": str(e), "cause_type": type(e).__name__}, ) from e - def transform(self, texts: str | list[str]) -> np.ndarray: + def transform(self, texts: str | list[str], **kwargs) -> np.ndarray: """Transforms input text(s) into embeddings using the GenAI API. Args: texts (str | list[str]): The input text(s) to embed. Can be a single string or a list of strings. + **kwargs: Additional parameters to pass to the `embed_content` method. Returns: numpy.ndarray: A 2D array of embeddings, where each row @@ -119,15 +120,23 @@ def transform(self, texts: str | list[str]) -> np.ndarray: `VectorisationError`: If the response format from the GenAI API is unexpected. """ + from google import genai # type: ignore + # If a single string is passed as arg to texts, convert to list if isinstance(texts, str): texts = [texts] + # Dynamically create the EmbedContentConfig, preserving the constructor's task_type + config = genai.types.EmbedContentConfig( + task_type=kwargs.pop( + "task_type", self.model_config.task_type + ), # Use the constructor's task_type if not overridden + **kwargs, # Pass any additional configuration options + ) + # The Vertex AI call to embed content try: - embeddings = self.vectoriser.models.embed_content( - model=self.model_name, contents=texts, config=self.model_config - ) + embeddings = self.vectoriser.models.embed_content(model=self.model_name, contents=texts, config=config) except Exception as e: raise ExternalServiceError( "GCP embedding request failed.",