|
20 | 20 | from pathlib import Path |
21 | 21 | from unittest.mock import Mock, patch |
22 | 22 |
|
| 23 | +import pytest |
23 | 24 | from transformers.configuration_utils import PretrainedConfig |
24 | 25 |
|
25 | 26 |
|
@@ -274,6 +275,42 @@ def test_save_artifacts_without_model_name_or_path(): |
274 | 275 | print("✅ test_save_artifacts_without_model_name_or_path passed") |
275 | 276 |
|
276 | 277 |
|
| 278 | +def test_save_artifacts_preserves_source_generation_config_after_validation_failure(): |
| 279 | + """Test invalid loaded generation configs are copied unchanged from the source.""" |
| 280 | + with tempfile.TemporaryDirectory() as tmp_dir: |
| 281 | + tmp_path = Path(tmp_dir) |
| 282 | + source_dir = tmp_path / "source" |
| 283 | + source_dir.mkdir() |
| 284 | + generation_config = '{"do_sample": false, "top_p": 0.95}\n' |
| 285 | + (source_dir / "generation_config.json").write_text(generation_config) |
| 286 | + |
| 287 | + target_dir = tmp_path / "target" |
| 288 | + base = MockPreTrainedBase(model_name_or_path=str(source_dir)) |
| 289 | + base._config = Mock() |
| 290 | + base._generation_config = Mock() |
| 291 | + base._generation_config.save_pretrained.side_effect = ValueError("invalid generation config") |
| 292 | + |
| 293 | + base.save_artifacts(target_dir) |
| 294 | + |
| 295 | + assert (target_dir / "generation_config.json").read_text() == generation_config |
| 296 | + |
| 297 | + |
| 298 | +def test_save_artifacts_reraises_generation_config_error_without_source_file(): |
| 299 | + """Test generation config validation errors remain fatal without a source artifact.""" |
| 300 | + with tempfile.TemporaryDirectory() as tmp_dir: |
| 301 | + tmp_path = Path(tmp_dir) |
| 302 | + source_dir = tmp_path / "source" |
| 303 | + source_dir.mkdir() |
| 304 | + |
| 305 | + base = MockPreTrainedBase(model_name_or_path=str(source_dir)) |
| 306 | + base._config = Mock() |
| 307 | + base._generation_config = Mock() |
| 308 | + base._generation_config.save_pretrained.side_effect = ValueError("invalid generation config") |
| 309 | + |
| 310 | + with pytest.raises(ValueError, match="invalid generation config"): |
| 311 | + base.save_artifacts(tmp_path / "target") |
| 312 | + |
| 313 | + |
277 | 314 | def test_copy_handles_permission_errors(): |
278 | 315 | """Test that copy failures are handled gracefully.""" |
279 | 316 | with tempfile.TemporaryDirectory() as tmp_dir: |
|
0 commit comments