From 831ab5510839720cbc9c3e3ebf8c522862d4c1b8 Mon Sep 17 00:00:00 2001 From: Cyber MacGeddon Date: Thu, 25 Sep 2025 19:57:03 +0100 Subject: [PATCH] Add temperature parameter to LlmService and roll out to VertexAI --- .../trustgraph/base/llm_service.py | 25 +++++++++++++-- .../model/text_completion/vertexai/llm.py | 32 +++++++++++-------- 2 files changed, 42 insertions(+), 15 deletions(-) diff --git a/trustgraph-base/trustgraph/base/llm_service.py b/trustgraph-base/trustgraph/base/llm_service.py index 9fb31a66..8eb6751c 100644 --- a/trustgraph-base/trustgraph/base/llm_service.py +++ b/trustgraph-base/trustgraph/base/llm_service.py @@ -5,7 +5,7 @@ LLM text completion base class import time import logging -from prometheus_client import Histogram +from prometheus_client import Histogram, Info from .. schema import TextCompletionRequest, TextCompletionResponse, Error from .. exceptions import TooManyRequests @@ -62,6 +62,12 @@ class LlmService(FlowProcessor): ) ) + self.register_specification( + ParameterSpec( + name = "temperature", + ) + ) + if not hasattr(__class__, "text_completion_metric"): __class__.text_completion_metric = Histogram( 'text_completion_duration', @@ -76,6 +82,13 @@ class LlmService(FlowProcessor): ] ) + if not hasattr(__class__, "text_completion_model_metric"): + __class__.text_completion_model_metric = Info( + 'text_completion_model', + 'Text completion model', + ["processor", "flow"] + ) + async def on_request(self, msg, consumer, flow): try: @@ -92,11 +105,19 @@ class LlmService(FlowProcessor): ).time(): model = flow("model") + temperature = flow("temperature") response = await self.generate_content( - request.system, request.prompt, model + request.system, request.prompt, model, temperature ) + await __class__.text_completion_model_metric.labels( + id = flow.id, flow = flow.name + ).info({ + "model": model, + "temperature": temperature, + }) + await flow("response").send( TextCompletionResponse( error=None, diff --git a/trustgraph-vertexai/trustgraph/model/text_completion/vertexai/llm.py b/trustgraph-vertexai/trustgraph/model/text_completion/vertexai/llm.py index 0ec9aca1..2fdfee62 100755 --- a/trustgraph-vertexai/trustgraph/model/text_completion/vertexai/llm.py +++ b/trustgraph-vertexai/trustgraph/model/text_completion/vertexai/llm.py @@ -152,29 +152,35 @@ class Processor(LlmService): return self.anthropic_client - def _get_gemini_model(self, model_name): + def _get_gemini_model(self, model_name, temperature=None): """Get or create a Gemini model instance""" if model_name not in self.model_clients: logger.info(f"Creating GenerativeModel instance for '{model_name}'") self.model_clients[model_name] = GenerativeModel(model_name) - # Create generation config for this model - self.generation_configs[model_name] = GenerationConfig( - temperature=self.temperature, - top_p=1.0, - top_k=10, - candidate_count=1, - max_output_tokens=self.max_output, - ) + # Use provided temperature or fall back to default + effective_temperature = temperature if temperature is not None else self.temperature - return self.model_clients[model_name], self.generation_configs[model_name] + # Create generation config with the effective temperature + generation_config = GenerationConfig( + temperature=effective_temperature, + top_p=1.0, + top_k=10, + candidate_count=1, + max_output_tokens=self.max_output, + ) - async def generate_content(self, system, prompt, model=None): + return self.model_clients[model_name], generation_config + + async def generate_content(self, system, prompt, model=None, temperature=None): # Use provided model or fall back to default model_name = model or self.default_model + # Use provided temperature or fall back to default + effective_temperature = temperature if temperature is not None else self.temperature logger.debug(f"Using model: {model_name}") + logger.debug(f"Using temperature: {effective_temperature}") try: if 'claude' in model_name.lower(): @@ -187,7 +193,7 @@ class Processor(LlmService): system=system, messages=[{"role": "user", "content": prompt}], max_tokens=self.api_params['max_output_tokens'], - temperature=self.api_params['temperature'], + temperature=effective_temperature, top_p=self.api_params['top_p'], top_k=self.api_params['top_k'], ) @@ -203,7 +209,7 @@ class Processor(LlmService): logger.debug(f"Sending request to Gemini model '{model_name}'...") full_prompt = system + "\n\n" + prompt - llm, generation_config = self._get_gemini_model(model_name) + llm, generation_config = self._get_gemini_model(model_name, effective_temperature) response = llm.generate_content( full_prompt, generation_config = generation_config,