diff --git a/src/google/adk/models/interactions_utils.py b/src/google/adk/models/interactions_utils.py index 97f4343bf2..bb890c40c5 100644 --- a/src/google/adk/models/interactions_utils.py +++ b/src/google/adk/models/interactions_utils.py @@ -135,11 +135,7 @@ # google-genai release does not declare on its request model. That model # discards keys it has no field for while serializing, so these never reach the # API and sending them is indistinguishable from never setting them. -_UNDECLARED_SAMPLING_PARAMS = ( - 'temperature', - 'top_p', - 'top_k', -) +_UNDECLARED_SAMPLING_PARAMS = ('top_k',) # Sampling knobs the interactions API itself rejects as unknown parameters. # No release of the client can carry these, so the caller has to stop setting @@ -1334,6 +1330,10 @@ def build_generation_config( A dictionary containing generation configuration parameters. """ generation_config: GenerationConfigParam = {} + if config.temperature is not None: + generation_config['temperature'] = config.temperature + if config.top_p is not None: + generation_config['top_p'] = config.top_p if config.max_output_tokens is not None: generation_config['max_output_tokens'] = config.max_output_tokens if config.stop_sequences: diff --git a/tests/unittests/models/test_interactions_utils.py b/tests/unittests/models/test_interactions_utils.py index adfefeed33..169265e0da 100644 --- a/tests/unittests/models/test_interactions_utils.py +++ b/tests/unittests/models/test_interactions_utils.py @@ -1346,6 +1346,8 @@ def test_all_parameters(self): ) result = interactions_utils.build_generation_config(config) assert result == { + 'temperature': 0.7, + 'top_p': 0.9, 'max_output_tokens': 100, 'stop_sequences': ['END'], 'seed': 7, @@ -1358,7 +1360,7 @@ def test_partial_parameters(self): max_output_tokens=50, ) result = interactions_utils.build_generation_config(config) - assert result == {'max_output_tokens': 50} + assert result == {'temperature': 0.5, 'max_output_tokens': 50} def test_empty_config(self): """Test building config with no parameters.""" @@ -1380,7 +1382,13 @@ def test_every_key_is_a_real_generation_config_field(self): ) result = interactions_utils.build_generation_config(config) supported = set(typing.get_type_hints(interactions.GenerationConfigParam)) - assert set(result) == {'max_output_tokens', 'stop_sequences', 'seed'} + assert set(result) == { + 'temperature', + 'top_p', + 'max_output_tokens', + 'stop_sequences', + 'seed', + } assert set(result) <= supported def test_dropped_parameters_are_the_ones_the_request_cannot_carry(self): @@ -1404,9 +1412,9 @@ def test_undeclared_parameters_point_at_the_client(self, caplog): r.getMessage() for r in caplog.records if r.levelno >= logging.WARNING ] assert len(warnings) == 1 - assert 'temperature' in warnings[0] - assert 'top_p' in warnings[0] assert 'top_k' in warnings[0] + assert 'temperature' not in warnings[0] + assert 'top_p' not in warnings[0] assert 'google-genai' in warnings[0] assert 'use_interactions_api' not in warnings[0] @@ -1431,7 +1439,7 @@ def test_unsupported_parameters_point_at_the_api(self, caplog): def test_the_two_causes_are_reported_separately(self, caplog): """Test that one cause is not folded into the other's remedy.""" - config = types.GenerateContentConfig(temperature=0.7, presence_penalty=0.5) + config = types.GenerateContentConfig(top_k=40, presence_penalty=0.5) with caplog.at_level( logging.WARNING, logger=interactions_utils.logger.name @@ -1443,12 +1451,12 @@ def test_the_two_causes_are_reported_separately(self, caplog): ] assert len(warnings) == 2 client, api = sorted(warnings, key=lambda w: 'use_interactions_api' in w) - assert 'temperature' in client and 'presence_penalty' not in client - assert 'presence_penalty' in api and 'temperature' not in api + assert 'top_k' in client and 'presence_penalty' not in client + assert 'presence_penalty' in api and 'top_k' not in api def test_dropped_parameters_are_logged_once(self, caplog): """Test that a parameter is reported once, not on every model turn.""" - config = types.GenerateContentConfig(temperature=0.7) + config = types.GenerateContentConfig(top_k=40) with caplog.at_level( logging.WARNING, logger=interactions_utils.logger.name