Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 5 additions & 5 deletions src/google/adk/models/interactions_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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:
Expand Down
24 changes: 16 additions & 8 deletions tests/unittests/models/test_interactions_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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."""
Expand All @@ -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):
Expand All @@ -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]

Expand All @@ -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
Expand All @@ -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
Expand Down
Loading