(
clear_transformers_cache,
tmp_dir,
model,
start_tokens,
max_length,
expected_tokens,
)
| 236 | ids=[args[0] for args in _TRANSFORMERS_GENERATION_TESTS], |
| 237 | ) |
| 238 | def test_transformers_generation( |
| 239 | clear_transformers_cache, |
| 240 | tmp_dir, |
| 241 | model, |
| 242 | start_tokens, |
| 243 | max_length, |
| 244 | expected_tokens, |
| 245 | ): |
| 246 | converter = ctranslate2.converters.TransformersConverter(model) |
| 247 | output_dir = str(tmp_dir.join("ctranslate2_model")) |
| 248 | output_dir = converter.convert(output_dir) |
| 249 | |
| 250 | generator = ctranslate2.Generator(output_dir) |
| 251 | results = generator.generate_batch([start_tokens.split()], max_length=max_length) |
| 252 | output_tokens = " ".join(results[0].sequences[0]) |
| 253 | assert output_tokens == expected_tokens |
| 254 | |
| 255 | # Test empty inputs. |
| 256 | assert generator.generate_batch([]) == [] |
| 257 | |
| 258 | with pytest.raises(ValueError, match="start token"): |
| 259 | generator.generate_batch([[]]) |
| 260 | |
| 261 | |
| 262 | @test_utils.only_on_linux |
nothing calls this directly
no test coverage detected