Skip to content

Commit 4650ef4

Browse files
committed
fix(conversion): preserve source generation config
Signed-off-by: yaoyu-33 <yaoyu.094@gmail.com>
1 parent 4b980e2 commit 4650ef4

2 files changed

Lines changed: 59 additions & 1 deletion

File tree

  • src/megatron/bridge/models/hf_pretrained
  • tests/unit_tests/models/hf_pretrained

src/megatron/bridge/models/hf_pretrained/base.py

Lines changed: 22 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -211,7 +211,28 @@ def save_artifacts(
211211
for name in self.OPTIONAL_ARTIFACTS:
212212
artifact = getattr(self, name, None)
213213
if artifact is not None and hasattr(artifact, "save_pretrained"):
214-
artifact.save_pretrained(save_path)
214+
try:
215+
artifact.save_pretrained(save_path)
216+
except ValueError:
217+
if name != "generation_config":
218+
raise
219+
220+
source_path = original_source_path or getattr(self, "model_name_or_path", None)
221+
if source_path is None:
222+
raise
223+
224+
copied_files = self._copy_custom_modeling_files(
225+
source_path=source_path,
226+
target_path=save_path,
227+
file_patterns=["generation_config.json"],
228+
)
229+
if "generation_config.json" not in copied_files:
230+
raise
231+
232+
logger.warning(
233+
"Generation config validation failed during export; preserving the source "
234+
"generation_config.json unchanged"
235+
)
215236

216237
# Download/copy additional files if specified
217238
if additional_files:

tests/unit_tests/models/hf_pretrained/test_base.py

Lines changed: 37 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,7 @@
2020
from pathlib import Path
2121
from unittest.mock import Mock, patch
2222

23+
import pytest
2324
from transformers.configuration_utils import PretrainedConfig
2425

2526

@@ -274,6 +275,42 @@ def test_save_artifacts_without_model_name_or_path():
274275
print("✅ test_save_artifacts_without_model_name_or_path passed")
275276

276277

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+
277314
def test_copy_handles_permission_errors():
278315
"""Test that copy failures are handled gracefully."""
279316
with tempfile.TemporaryDirectory() as tmp_dir:

0 commit comments

Comments
 (0)