Skip to content

Commit 40a45d2

Browse files
Merge pull request #4515 from AI-Hypercomputer:fix/decoder-sampling
PiperOrigin-RevId: 950892221
2 parents dfd8d29 + 156380f commit 40a45d2

1 file changed

Lines changed: 6 additions & 4 deletions

File tree

tests/integration/decode_tests.py

Lines changed: 6 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -181,16 +181,18 @@ def test_decode_greedy_sampling(self):
181181
def test_decode_weighted_sampling(self):
182182
config = DecodeTests.CONFIGS["decode_sampling"] + DecodeTests.SAMPLING_STRATEGY_CONFIG["weighted"]
183183
captured_out = run_decoding(config)
184-
expected_output = "Input `I love to` -> ` travel and I love to write"
185-
assert expected_output in captured_out
184+
expected_output_pre_v7x = "Input `I love to` -> ` travel and I love to write"
185+
expected_output_v7x = "Input `I love to` -> ` travel. I love to explore new places,"
186+
assert (expected_output_pre_v7x in captured_out) or (expected_output_v7x in captured_out)
186187

187188
@pytest.mark.tpu_only
188189
@pytest.mark.scheduled_only
189190
def test_decode_nucleus_sampling(self):
190191
config = DecodeTests.CONFIGS["decode_sampling"] + DecodeTests.SAMPLING_STRATEGY_CONFIG["nucleus"]
191192
captured_out = run_decoding(config)
192-
expected_output = "Input `I love to` -> ` travel and I love to write"
193-
assert expected_output in captured_out
193+
expected_output_pre_v7x = "Input `I love to` -> ` travel and I love to write"
194+
expected_output_v7x = "Input `I love to` -> ` travel. I love to explore new places,"
195+
assert (expected_output_pre_v7x in captured_out) or (expected_output_v7x in captured_out)
194196

195197
@pytest.mark.tpu_only
196198
@pytest.mark.scheduled_only

0 commit comments

Comments
 (0)