diff --git a/examples/puzzletron/configs/families/qwen3_5/qwen3p5_0p8b/advanced.yaml b/examples/puzzletron/configs/families/qwen3_5/qwen3p5_0p8b/advanced.yaml new file mode 100644 index 00000000000..e951db2ad80 --- /dev/null +++ b/examples/puzzletron/configs/families/qwen3_5/qwen3p5_0p8b/advanced.yaml @@ -0,0 +1,54 @@ +# @package _global_ + +defaults: + - /families/qwen3_5/qwen3p5_0p8b/model@_global_ + - _self_ + +# Opt-in experimental search overlay. The axis structure is adapted from the +# existing Qwen 3.5 9B config, while the exact targets are derived from the +# pinned 0.8B geometry in model.yaml. These targets were not selected from a +# fully run 0.8B campaign and have not completed full runtime validation. +pruning: + intermediate_size_list: [3072, 2560, 2048, 1792, 1536] + attn_heads_list: + - [2, 1] + - [4, 1] + - [4, 2] + - [8, 2] + +search_space: + axes: + hidden_width: + enabled: true + teacher_value: 1024 + values: [768] + kv_groups: + enabled: true + teacher_value: 2 + values: [1] + q_heads_per_group: + enabled: true + teacher_value: 4 + values: [2] + ffn_intermediate: + enabled: true + teacher_value: 3584 + values: [3072, 2560, 2048, 1792, 1536] + gdn_key_groups: + enabled: true + teacher_value: 16 + values: [12, 8] + gdn_value_heads_per_group: + enabled: false + teacher_value: 1 + values: [] + # The proposed 128 -> 96 target remains blocked on runtime-equivalence + # evidence. Keep it documented but non-executable until that gate passes. + gdn_key_head_dim: + enabled: false + teacher_value: 128 + values: [] + gdn_value_head_dim: + enabled: true + teacher_value: 128 + values: [96] diff --git a/examples/puzzletron/configs/families/qwen3_5/qwen3p5_0p8b/model.yaml b/examples/puzzletron/configs/families/qwen3_5/qwen3p5_0p8b/model.yaml new file mode 100644 index 00000000000..574ce521227 --- /dev/null +++ b/examples/puzzletron/configs/families/qwen3_5/qwen3p5_0p8b/model.yaml @@ -0,0 +1,46 @@ +# @package _global_ + +# Runnable public Hugging Face repository ID; this is not a placeholder. +input_hf_model_path: Qwen/Qwen3.5-0.8B + +model_info: + # Immutable public model identity used by the focused MIP smoke. + # Geometry is reproducible from the config at this exact revision. + hf_repo: Qwen/Qwen3.5-0.8B + hf_revision: 2fc06364715b967f1860aea9cf38778875588b17 + model_type: qwen3_5 + architectures: [Qwen3_5ForConditionalGeneration] + num_hidden_layers: 24 + hidden_size: 1024 + intermediate_size: 3584 + num_attention_heads: 8 + num_key_value_heads: 2 + head_dim: 256 + vocab_size: 248320 + tie_word_embeddings: true + max_position_embeddings: 262144 + mtp_num_hidden_layers: 1 + layer_counts: + linear_attention: 18 + full_attention: 6 + mamba: + linear_key_head_dim: 128 + linear_num_key_heads: 16 + linear_num_value_heads: 16 + linear_value_head_dim: 128 + linear_conv_kernel_dim: 4 + +model: + revision: ${model_info.hf_revision} + +# Keep the default search aligned with the tracked Qwen 3.5 0.8B runtime +# campaign. Candidate construction adds the teacher value for the enabled axis. +pruning: + intermediate_size_list: [3072, 2048] + +search_space: + axes: + ffn_intermediate: + enabled: true + teacher_value: 3584 + values: [3072, 2048] diff --git a/modelopt/torch/puzzletron/stages/convert.py b/modelopt/torch/puzzletron/stages/convert.py index 9fe0a207ee0..97d15e14a7b 100644 --- a/modelopt/torch/puzzletron/stages/convert.py +++ b/modelopt/torch/puzzletron/stages/convert.py @@ -27,6 +27,11 @@ __all__ = ["convert_stage"] +_CONVERSION_SOURCE_METADATA = "puzzletron_conversion_source.json" +_CONVERSION_TRANSACTION_SUFFIX = ".puzzletron-convert-tmp" +_CONVERSION_BACKUP_SUFFIX = ".puzzletron-convert-backup" + + def _register_automodel_config_aliases() -> None: """Backward-compatible stage wrapper around the shared config registry.""" @@ -110,6 +115,139 @@ def _is_complete_checkpoint(path: Path, *, trust_remote_code: bool) -> bool: return True +def _conversion_source_matches( + path: Path, + *, + source: str, + revision: str | None, +) -> bool: + """Return whether a converted checkpoint belongs to the configured source revision.""" + metadata_path = path / _CONVERSION_SOURCE_METADATA + if not metadata_path.is_file(): + # Preserve resume compatibility for legacy unpinned checkpoints. Pinned + # runs must convert once to establish an auditable source revision. + return revision is None + try: + metadata = json.loads(metadata_path.read_text(encoding="utf-8")) + except (OSError, json.JSONDecodeError): + return False + if not isinstance(metadata, dict): + return False + return metadata.get("source") == source and metadata.get("revision") == revision + + +def _write_conversion_source_metadata( + path: Path, + *, + source: str, + revision: str | None, + source_identity: str, +) -> None: + metadata_path = path / _CONVERSION_SOURCE_METADATA + tmp_path = metadata_path.with_suffix(metadata_path.suffix + ".tmp") + tmp_path.write_text( + json.dumps( + { + "source": source, + "revision": revision, + "source_identity": source_identity, + "version": 1, + }, + indent=2, + sort_keys=True, + ) + + "\n", + encoding="utf-8", + ) + tmp_path.replace(metadata_path) + + +def _raise_on_master_failure(error: BaseException | None, *, action: str) -> None: + failure = f"{type(error).__name__}: {error}" if error is not None else None + failure = dist.broadcast(failure, src=0) + if failure is None: + return + if error is not None: + raise error + raise RuntimeError(f"rank 0 failed during {action}: {failure}") + + +def _conversion_sibling(path: Path, suffix: str) -> Path: + return path.with_name(f".{path.name}{suffix}") + + +def _remove_conversion_directory(path: Path) -> None: + if path.is_symlink() or not path.is_dir(): + raise RuntimeError(f"conversion transaction path is not a directory: {path}") + shutil.rmtree(path) + + +def _recover_conversion_transaction(path: Path) -> None: + """Restore the last published checkpoint after an interrupted directory swap.""" + transaction_dir = _conversion_sibling(path, _CONVERSION_TRANSACTION_SUFFIX) + backup_dir = _conversion_sibling(path, _CONVERSION_BACKUP_SUFFIX) + if backup_dir.exists(): + if path.exists(): + if transaction_dir.exists(): + raise RuntimeError(f"ambiguous conversion transaction state for checkpoint: {path}") + _remove_conversion_directory(backup_dir) + else: + backup_dir.replace(path) + if transaction_dir.exists(): + _remove_conversion_directory(transaction_dir) + + +def _prepare_conversion_transaction(path: Path) -> Path: + path.parent.mkdir(parents=True, exist_ok=True) + _recover_conversion_transaction(path) + transaction_dir = _conversion_sibling(path, _CONVERSION_TRANSACTION_SUFFIX) + transaction_dir.mkdir() + return transaction_dir + + +def _publish_conversion_transaction(path: Path, transaction_dir: Path) -> None: + """Publish a validated checkpoint while retaining the prior version until the swap.""" + backup_dir = _conversion_sibling(path, _CONVERSION_BACKUP_SUFFIX) + if backup_dir.exists(): + raise RuntimeError(f"conversion backup already exists: {backup_dir}") + moved_existing = False + try: + if path.exists(): + path.replace(backup_dir) + moved_existing = True + transaction_dir.replace(path) + except Exception: + if moved_existing and not path.exists() and backup_dir.exists(): + backup_dir.replace(path) + raise + if backup_dir.exists(): + _remove_conversion_directory(backup_dir) + + +def _validate_converted_checkpoint( + path: Path, + *, + source: str, + revision: str | None, + trust_remote_code: bool, + descriptor_override: str | None, +) -> None: + if not _is_complete_checkpoint(path, trust_remote_code=trust_remote_code): + raise RuntimeError(f"conversion did not produce a complete checkpoint: {path}") + if not _conversion_source_matches(path, source=source, revision=revision): + raise RuntimeError(f"conversion source metadata does not match checkpoint: {path}") + converted_config = AutoConfig.from_pretrained(path, trust_remote_code=trust_remote_code) + converted_resolution = resolve_descriptor_from_pretrained( + str(path), + trust_remote_code=trust_remote_code, + descriptor_override=descriptor_override, + ) + if not _descriptor_checkpoint_layout_complete( + path, converted_resolution.descriptor, converted_config + ): + raise RuntimeError(f"conversion produced an incompatible checkpoint layout: {path}") + + def _checkpoint_weight_map(path: Path) -> tuple[dict[str, str], str | None]: index_path = path / "model.safetensors.index.json" if index_path.is_file(): @@ -224,7 +362,7 @@ def _ensure_untied_word_embeddings( return True -def _resolve_source_path(source: str) -> Path: +def _resolve_source_path(source: str, *, revision: str | None = None) -> Path: source_path = Path(source) if source_path.exists(): return source_path @@ -234,7 +372,7 @@ def _resolve_source_path(source: str) -> Path: model_id = "/".join(source.rstrip("/").split("/")[-2:]) else: model_id = source - return Path(snapshot_download(repo_id=model_id)) + return Path(snapshot_download(repo_id=model_id, revision=revision)) def convert_stage(config: dict[str, Any], manifest: StageManifest): @@ -246,6 +384,7 @@ def convert_stage(config: dict[str, Any], manifest: StageManifest): raise ValueError("model.source is required for the convert stage") trust_remote_code = bool(model_cfg.get("trust_remote_code", False)) + revision = model_cfg.get("revision") convert_cfg = config.get("convert") or {} untie_word_embeddings = bool(convert_cfg.get("untie_word_embeddings", False)) teacher_dir = _teacher_dir(config) @@ -255,7 +394,21 @@ def convert_stage(config: dict[str, Any], manifest: StageManifest): skipped = False with _distributed_if_needed(): + recovery_error = None + if dist.is_master(): + try: + _recover_conversion_transaction(teacher_dir) + except BaseException as error: + recovery_error = error + _raise_on_master_failure(recovery_error, action="conversion transaction recovery") + dist.barrier() teacher_complete = _is_complete_checkpoint(teacher_dir, trust_remote_code=trust_remote_code) + if teacher_complete: + teacher_complete = _conversion_source_matches( + teacher_dir, + source=str(source), + revision=revision, + ) if teacher_complete: teacher_config = AutoConfig.from_pretrained( teacher_dir, trust_remote_code=trust_remote_code @@ -269,7 +422,9 @@ def convert_stage(config: dict[str, Any], manifest: StageManifest): teacher_dir, resolution.descriptor, teacher_config ) if teacher_complete and untie_word_embeddings: - teacher_config = AutoConfig.from_pretrained(teacher_dir, trust_remote_code=trust_remote_code) + teacher_config = AutoConfig.from_pretrained( + teacher_dir, trust_remote_code=trust_remote_code + ) teacher_tied = bool(_get(teacher_config, "tie_word_embeddings", False)) if teacher_tied: teacher_complete = False @@ -277,45 +432,70 @@ def convert_stage(config: dict[str, Any], manifest: StageManifest): if teacher_complete: skipped = True else: + source_error = None if dist.is_master(): - source_path = _resolve_source_path(str(source)) + try: + source_path = _resolve_source_path(str(source), revision=revision) + except BaseException as error: + source_path = None + source_error = error else: source_path = None + _raise_on_master_failure(source_error, action="source checkpoint resolution") source_path = Path(dist.broadcast(str(source_path), src=0)) - source_config = AutoConfig.from_pretrained(source_path, trust_remote_code=trust_remote_code) + source_config = AutoConfig.from_pretrained( + source_path, trust_remote_code=trust_remote_code + ) source_identity = model_identity(source_config).value + conversion_error = None if dist.is_master(): - if _is_anymodel_config(source_config): - teacher_dir.parent.mkdir(parents=True, exist_ok=True) - if source_path.resolve() != teacher_dir.resolve(): - shutil.copytree(source_path, teacher_dir, dirs_exist_ok=True) - already_anymodel = True - else: - resolution = resolve_descriptor_from_pretrained( - str(source_path), - trust_remote_code=trust_remote_code, - descriptor_override=model_cfg.get("descriptor_override"), - ) - converter = ConverterFactory.get(resolution.name) - converter.convert( - descriptor=resolution.descriptor, - input_dir=source_path, - output_dir=teacher_dir, - ) - descriptor_payload = resolution.to_dict() - already_anymodel = False - if untie_word_embeddings: - if "resolution" not in locals(): - resolution = resolve_descriptor_from_pretrained( - str(teacher_dir), + try: + transaction_dir = _prepare_conversion_transaction(teacher_dir) + if _is_anymodel_config(source_config): + shutil.copytree(source_path, transaction_dir, dirs_exist_ok=True) + already_anymodel = True + else: + source_resolution = resolve_descriptor_from_pretrained( + str(source_path), + trust_remote_code=trust_remote_code, + descriptor_override=model_cfg.get("descriptor_override"), + ) + converter = ConverterFactory.get(source_resolution.name) + converter.convert( + descriptor=source_resolution.descriptor, + input_dir=source_path, + output_dir=transaction_dir, + ) + descriptor_payload = source_resolution.to_dict() + already_anymodel = False + if untie_word_embeddings: + target_resolution = resolve_descriptor_from_pretrained( + str(transaction_dir), trust_remote_code=trust_remote_code, descriptor_override=model_cfg.get("descriptor_override"), ) - _ensure_untied_word_embeddings( - teacher_dir, - descriptor=resolution.descriptor, + _ensure_untied_word_embeddings( + transaction_dir, + descriptor=target_resolution.descriptor, + trust_remote_code=trust_remote_code, + ) + _write_conversion_source_metadata( + transaction_dir, + source=str(source), + revision=revision, + source_identity=source_identity, + ) + _validate_converted_checkpoint( + transaction_dir, + source=str(source), + revision=revision, trust_remote_code=trust_remote_code, + descriptor_override=model_cfg.get("descriptor_override"), ) + _publish_conversion_transaction(teacher_dir, transaction_dir) + except BaseException as error: + conversion_error = error + _raise_on_master_failure(conversion_error, action="teacher conversion") dist.barrier() teacher_config = AutoConfig.from_pretrained(teacher_dir, trust_remote_code=trust_remote_code) diff --git a/tests/unit/torch/puzzletron/test_convert_anymodel.py b/tests/unit/torch/puzzletron/test_convert_anymodel.py index 69bad3a4854..4affcd8f3ad 100644 --- a/tests/unit/torch/puzzletron/test_convert_anymodel.py +++ b/tests/unit/torch/puzzletron/test_convert_anymodel.py @@ -14,6 +14,7 @@ # limitations under the License. import json +from contextlib import nullcontext from types import SimpleNamespace import pytest @@ -29,9 +30,12 @@ from transformers import AutoModelForCausalLM import modelopt.torch.puzzletron as mtpz +import modelopt.torch.puzzletron.stages.convert as convert_stage_module from modelopt.torch.puzzletron.stages.convert import ( + _conversion_source_matches, _descriptor_checkpoint_layout_complete, _is_complete_checkpoint, + _write_conversion_source_metadata, ) from modelopt.torch.puzzletron.tools.checkpoint_utils_hf import load_model_config @@ -45,6 +49,480 @@ def _weight_map(checkpoint_dir): return dict.fromkeys(handle.keys(), weights_path.name) +def _patch_single_rank_convert(monkeypatch): + monkeypatch.setattr(convert_stage_module, "_register_automodel_config_aliases", lambda: None) + monkeypatch.setattr(convert_stage_module, "_distributed_if_needed", nullcontext) + + +def test_nonmaster_raises_broadcast_conversion_failure(monkeypatch): + monkeypatch.setattr( + convert_stage_module.dist, + "broadcast", + lambda value, src: "ConversionError: failed conversion", + ) + + with pytest.raises(RuntimeError, match="rank 0 failed during teacher conversion"): + convert_stage_module._raise_on_master_failure(None, action="teacher conversion") + + +def test_master_broadcasts_failure_before_reraising(monkeypatch): + events = [] + + class ConversionError(RuntimeError): + pass + + error = ConversionError("failed conversion") + + def broadcast(value, src): + events.append(("broadcast", value)) + return value + + monkeypatch.setattr(convert_stage_module.dist, "broadcast", broadcast) + + with pytest.raises(ConversionError) as raised: + convert_stage_module._raise_on_master_failure(error, action="teacher conversion") + events.append(("raised", str(raised.value))) + + assert raised.value is error + assert events == [ + ("broadcast", "ConversionError: failed conversion"), + ("raised", "failed conversion"), + ] + + +@pytest.mark.parametrize("revision", ["pinned-sha", None], ids=["pinned", "default"]) +def test_convert_stage_pins_optional_hugging_face_revision(tmp_path, monkeypatch, revision): + calls = [] + + class SourceResolvedError(RuntimeError): + pass + + def snapshot_download(*, repo_id, revision): + calls.append({"repo_id": repo_id, "revision": revision}) + raise SourceResolvedError + + _patch_single_rank_convert(monkeypatch) + monkeypatch.setattr( + convert_stage_module, "_is_complete_checkpoint", lambda *args, **kwargs: False + ) + monkeypatch.setattr("huggingface_hub.snapshot_download", snapshot_download) + model = {"source": "Qwen/Qwen3.5-0.8B"} + if revision is not None: + model["revision"] = revision + + with pytest.raises(SourceResolvedError): + convert_stage_module.convert_stage( + {"model": model, "convert": {"teacher_dir": str(tmp_path / "teacher")}}, + manifest=object(), + ) + + assert calls == [{"repo_id": "Qwen/Qwen3.5-0.8B", "revision": revision}] + + +def test_conversion_resume_requires_matching_source_revision(tmp_path): + teacher_dir = tmp_path / "teacher" + teacher_dir.mkdir() + + assert _conversion_source_matches(teacher_dir, source="Qwen/Qwen3.5-0.8B", revision=None) + assert not _conversion_source_matches( + teacher_dir, source="Qwen/Qwen3.5-0.8B", revision="new-sha" + ) + + _write_conversion_source_metadata( + teacher_dir, + source="Qwen/Qwen3.5-0.8B", + revision="old-sha", + source_identity="source_model_abc", + ) + assert _conversion_source_matches(teacher_dir, source="Qwen/Qwen3.5-0.8B", revision="old-sha") + assert not _conversion_source_matches( + teacher_dir, source="Qwen/Qwen3.5-0.8B", revision="new-sha" + ) + assert not _conversion_source_matches( + teacher_dir, source="Qwen/Another-Model", revision="old-sha" + ) + + metadata_path = teacher_dir / convert_stage_module._CONVERSION_SOURCE_METADATA + metadata_path.write_text("not-json") + assert not _conversion_source_matches( + teacher_dir, source="Qwen/Qwen3.5-0.8B", revision="old-sha" + ) + metadata_path.write_text("[]") + assert not _conversion_source_matches( + teacher_dir, source="Qwen/Qwen3.5-0.8B", revision="old-sha" + ) + + +def test_convert_stage_reconverts_teacher_from_different_revision(tmp_path, monkeypatch): + teacher_dir = tmp_path / "teacher" + teacher_dir.mkdir() + (teacher_dir / "old-shard.bin").write_text("old") + _write_conversion_source_metadata( + teacher_dir, + source="Qwen/Qwen3.5-0.8B", + revision="old-sha", + source_identity="source_model_abc", + ) + metadata_path = teacher_dir / convert_stage_module._CONVERSION_SOURCE_METADATA + metadata_before = metadata_path.read_bytes() + resolved = [] + events = [] + + class SourceResolvedError(RuntimeError): + pass + + def resolve_source(source, *, revision): + resolved.append({"source": source, "revision": revision}) + raise SourceResolvedError("source failed") + + def broadcast(value, src): + events.append(("broadcast", value)) + return value + + _patch_single_rank_convert(monkeypatch) + monkeypatch.setattr(convert_stage_module.dist, "broadcast", broadcast) + monkeypatch.setattr( + convert_stage_module.dist, "barrier", lambda: events.append(("barrier", None)) + ) + monkeypatch.setattr( + convert_stage_module, "_is_complete_checkpoint", lambda *args, **kwargs: True + ) + monkeypatch.setattr(convert_stage_module, "_resolve_source_path", resolve_source) + + with pytest.raises(SourceResolvedError): + convert_stage_module.convert_stage( + { + "model": {"source": "Qwen/Qwen3.5-0.8B", "revision": "new-sha"}, + "convert": {"teacher_dir": str(teacher_dir)}, + }, + manifest=object(), + ) + + assert resolved == [{"source": "Qwen/Qwen3.5-0.8B", "revision": "new-sha"}] + assert events == [ + ("broadcast", None), + ("barrier", None), + ("broadcast", "SourceResolvedError: source failed"), + ] + assert metadata_path.read_bytes() == metadata_before + assert (teacher_dir / "old-shard.bin").read_text() == "old" + assert not convert_stage_module._conversion_sibling( + teacher_dir, convert_stage_module._CONVERSION_TRANSACTION_SUFFIX + ).exists() + assert not convert_stage_module._conversion_sibling( + teacher_dir, convert_stage_module._CONVERSION_BACKUP_SUFFIX + ).exists() + + +def test_convert_stage_persists_and_reuses_matching_source_revision(tmp_path, monkeypatch): + source_dir = tmp_path / "source" + source_dir.mkdir() + (source_dir / "tokenizer.json").write_text("{}") + teacher_dir = tmp_path / "teacher" + source_config = SimpleNamespace(architectures=["AnyModel"]) + resolution = SimpleNamespace(descriptor=object()) + resolve_calls = [] + completions = [] + + _patch_single_rank_convert(monkeypatch) + monkeypatch.setattr( + convert_stage_module, + "_is_complete_checkpoint", + lambda path, **kwargs: (path / convert_stage_module._CONVERSION_SOURCE_METADATA).is_file(), + ) + monkeypatch.setattr( + convert_stage_module.AutoConfig, + "from_pretrained", + lambda *args, **kwargs: source_config, + ) + monkeypatch.setattr( + convert_stage_module, + "resolve_descriptor_from_pretrained", + lambda *args, **kwargs: resolution, + ) + monkeypatch.setattr( + convert_stage_module, "_descriptor_checkpoint_layout_complete", lambda *args: True + ) + monkeypatch.setattr( + convert_stage_module, + "model_identity", + lambda config: SimpleNamespace(value="source_model_abc"), + ) + + def resolve_source(source, *, revision): + resolve_calls.append({"source": source, "revision": revision}) + return source_dir + + def complete_stage(config, manifest, *, outputs, status, message=None): + completions.append({"outputs": outputs, "status": status, "message": message}) + return completions[-1] + + monkeypatch.setattr(convert_stage_module, "_resolve_source_path", resolve_source) + monkeypatch.setattr(convert_stage_module, "complete_stage", complete_stage) + config = { + "model": {"source": "Qwen/Qwen3.5-0.8B", "revision": "pinned-sha"}, + "convert": {"teacher_dir": str(teacher_dir)}, + } + + first = convert_stage_module.convert_stage(config, manifest=object()) + transaction_dir = convert_stage_module._conversion_sibling( + teacher_dir, convert_stage_module._CONVERSION_TRANSACTION_SUFFIX + ) + backup_dir = convert_stage_module._conversion_sibling( + teacher_dir, convert_stage_module._CONVERSION_BACKUP_SUFFIX + ) + assert not transaction_dir.exists() + assert not backup_dir.exists() + + second = convert_stage_module.convert_stage(config, manifest=object()) + + metadata = json.loads( + (teacher_dir / convert_stage_module._CONVERSION_SOURCE_METADATA).read_text() + ) + assert metadata == { + "revision": "pinned-sha", + "source": "Qwen/Qwen3.5-0.8B", + "source_identity": "source_model_abc", + "version": 1, + } + assert resolve_calls == [{"source": "Qwen/Qwen3.5-0.8B", "revision": "pinned-sha"}] + assert first["status"] == "success" + assert first["outputs"]["skipped"] is False + assert second["status"] == "skipped" + assert second["outputs"]["skipped"] is True + + +def test_convert_stage_reuses_legacy_unpinned_teacher(tmp_path, monkeypatch): + teacher_dir = tmp_path / "teacher" + teacher_dir.mkdir() + teacher_config = SimpleNamespace(architectures=["AnyModel"]) + + _patch_single_rank_convert(monkeypatch) + monkeypatch.setattr( + convert_stage_module, "_is_complete_checkpoint", lambda *args, **kwargs: True + ) + monkeypatch.setattr( + convert_stage_module.AutoConfig, + "from_pretrained", + lambda *args, **kwargs: teacher_config, + ) + monkeypatch.setattr( + convert_stage_module, + "resolve_descriptor_from_pretrained", + lambda *args, **kwargs: SimpleNamespace(descriptor=object()), + ) + monkeypatch.setattr( + convert_stage_module, "_descriptor_checkpoint_layout_complete", lambda *args: True + ) + monkeypatch.setattr( + convert_stage_module, + "model_identity", + lambda config: SimpleNamespace(value="teacher_model_abc"), + ) + monkeypatch.setattr( + convert_stage_module, + "_resolve_source_path", + lambda *args, **kwargs: pytest.fail("legacy unpinned teacher should be reused"), + ) + monkeypatch.setattr( + convert_stage_module, + "complete_stage", + lambda config, manifest, *, outputs, status, message=None: { + "outputs": outputs, + "status": status, + }, + ) + + result = convert_stage_module.convert_stage( + { + "model": {"source": "Qwen/Qwen3.5-0.8B"}, + "convert": {"teacher_dir": str(teacher_dir)}, + }, + manifest=object(), + ) + + assert result["status"] == "skipped" + assert result["outputs"]["skipped"] is True + + +def test_failed_reconversion_preserves_teacher_and_retries_cleanly(tmp_path, monkeypatch): + source_dir = tmp_path / "source" + source_dir.mkdir() + teacher_dir = tmp_path / "teacher" + teacher_dir.mkdir() + (teacher_dir / "old-shard.bin").write_text("old") + _write_conversion_source_metadata( + teacher_dir, + source="Qwen/Qwen3.5-0.8B", + revision="old-sha", + source_identity="source_model_old", + ) + source_config = SimpleNamespace(architectures=["Qwen3ForCausalLM"]) + + class ConversionError(RuntimeError): + pass + + class Resolution: + name = "qwen3" + descriptor = object() + + @staticmethod + def to_dict(): + return {"name": "qwen3"} + + class RetryConverter: + attempts = 0 + + def convert(self, *, output_dir, **kwargs): + self.attempts += 1 + (output_dir / "new-shard.bin").write_text("new") + if self.attempts == 1: + raise ConversionError + + converter = RetryConverter() + barriers = [] + + _patch_single_rank_convert(monkeypatch) + monkeypatch.setattr(convert_stage_module.dist, "barrier", lambda: barriers.append(None)) + monkeypatch.setattr( + convert_stage_module, + "_is_complete_checkpoint", + lambda path, **kwargs: path == teacher_dir + or ( + (path / "new-shard.bin").is_file() + and (path / convert_stage_module._CONVERSION_SOURCE_METADATA).is_file() + ), + ) + monkeypatch.setattr( + convert_stage_module, "_resolve_source_path", lambda *args, **kwargs: source_dir + ) + monkeypatch.setattr( + convert_stage_module.AutoConfig, + "from_pretrained", + lambda *args, **kwargs: source_config, + ) + monkeypatch.setattr( + convert_stage_module, + "model_identity", + lambda config: SimpleNamespace(value="source_model_new"), + ) + monkeypatch.setattr( + convert_stage_module, + "resolve_descriptor_from_pretrained", + lambda *args, **kwargs: Resolution(), + ) + monkeypatch.setattr( + convert_stage_module, "_descriptor_checkpoint_layout_complete", lambda *args: True + ) + monkeypatch.setattr( + convert_stage_module, + "complete_stage", + lambda config, manifest, *, outputs, status, message=None: { + "outputs": outputs, + "status": status, + }, + ) + monkeypatch.setattr(convert_stage_module.ConverterFactory, "get", lambda name: converter) + + config = { + "model": {"source": "Qwen/Qwen3.5-0.8B", "revision": "new-sha"}, + "convert": {"teacher_dir": str(teacher_dir)}, + } + + with pytest.raises(ConversionError): + convert_stage_module.convert_stage(config, manifest=object()) + + assert len(barriers) == 1 + + old_metadata = json.loads( + (teacher_dir / convert_stage_module._CONVERSION_SOURCE_METADATA).read_text() + ) + assert old_metadata["revision"] == "old-sha" + assert (teacher_dir / "old-shard.bin").read_text() == "old" + transaction_dir = convert_stage_module._conversion_sibling( + teacher_dir, convert_stage_module._CONVERSION_TRANSACTION_SUFFIX + ) + assert (transaction_dir / "new-shard.bin").read_text() == "new" + + result = convert_stage_module.convert_stage(config, manifest=object()) + + assert result["status"] == "success" + assert len(barriers) == 3 + assert converter.attempts == 2 + assert not (teacher_dir / "old-shard.bin").exists() + assert (teacher_dir / "new-shard.bin").read_text() == "new" + new_metadata = json.loads( + (teacher_dir / convert_stage_module._CONVERSION_SOURCE_METADATA).read_text() + ) + assert new_metadata["revision"] == "new-sha" + assert not transaction_dir.exists() + + +def test_conversion_transaction_recovers_interrupted_swap(tmp_path, monkeypatch): + teacher_dir = tmp_path / "teacher" + teacher_dir.mkdir() + (teacher_dir / "old-shard.bin").write_text("old") + transaction_dir = convert_stage_module._conversion_sibling( + teacher_dir, convert_stage_module._CONVERSION_TRANSACTION_SUFFIX + ) + transaction_dir.mkdir() + (transaction_dir / "new-shard.bin").write_text("new") + backup_dir = convert_stage_module._conversion_sibling( + teacher_dir, convert_stage_module._CONVERSION_BACKUP_SUFFIX + ) + teacher_dir.replace(backup_dir) + + convert_stage_module._recover_conversion_transaction(teacher_dir) + + assert (teacher_dir / "old-shard.bin").read_text() == "old" + assert not transaction_dir.exists() + assert not backup_dir.exists() + + teacher_dir.replace(backup_dir) + teacher_dir.mkdir() + (teacher_dir / "published-shard.bin").write_text("published") + + convert_stage_module._recover_conversion_transaction(teacher_dir) + + assert (teacher_dir / "published-shard.bin").read_text() == "published" + assert not backup_dir.exists() + + backup_dir.mkdir() + (backup_dir / "backup-shard.bin").write_text("backup") + transaction_dir.mkdir() + (transaction_dir / "transaction-shard.bin").write_text("transaction") + events = [] + + _patch_single_rank_convert(monkeypatch) + monkeypatch.setattr( + convert_stage_module.dist, + "broadcast", + lambda value, src: events.append(("broadcast", value)) or value, + ) + monkeypatch.setattr( + convert_stage_module.dist, "barrier", lambda: events.append(("barrier", None)) + ) + + with pytest.raises(RuntimeError, match="ambiguous conversion transaction state"): + convert_stage_module.convert_stage( + { + "model": {"source": "Qwen/Qwen3.5-0.8B"}, + "convert": {"teacher_dir": str(teacher_dir)}, + }, + manifest=object(), + ) + + assert events == [ + ( + "broadcast", + f"RuntimeError: ambiguous conversion transaction state for checkpoint: {teacher_dir}", + ) + ] + assert (teacher_dir / "published-shard.bin").read_text() == "published" + assert (backup_dir / "backup-shard.bin").read_text() == "backup" + assert (transaction_dir / "transaction-shard.bin").read_text() == "transaction" + + def test_convert_anymodel(tmp_path): input_dir = create_tiny_qwen3_dir(tmp_path, with_tokenizer=True) output_dir = tmp_path / "qwen3-0.6b-anymodel" @@ -85,9 +563,7 @@ def generic_decoder_contract(config): ) ) - assert not _descriptor_checkpoint_layout_complete( - checkpoint, Descriptor, SimpleNamespace() - ) + assert not _descriptor_checkpoint_layout_complete(checkpoint, Descriptor, SimpleNamespace()) def test_conversion_resume_skips_generic_layout_check_when_descriptor_has_no_contract(tmp_path): diff --git a/tests/unit/torch/puzzletron/test_portable_configs.py b/tests/unit/torch/puzzletron/test_portable_configs.py index 3389abf0455..6fc6552dba4 100644 --- a/tests/unit/torch/puzzletron/test_portable_configs.py +++ b/tests/unit/torch/puzzletron/test_portable_configs.py @@ -28,6 +28,7 @@ NEMOTRON3_NANO_30B_MODEL_CONFIG = ( "examples/puzzletron/configs/families/nemotron3/nano_30b_a3b_bf16/model.yaml" ) +QWEN3P5_0P8B_MODEL_CONFIG = "examples/puzzletron/configs/families/qwen3_5/qwen3p5_0p8b/model.yaml" QWEN3P5_9B_MODEL_CONFIG = "examples/puzzletron/configs/families/qwen3_5/qwen3p5_9b/model.yaml" QWEN3P6_35B_A3B_MODEL_CONFIG = ( "examples/puzzletron/configs/families/qwen3_5/qwen3p6_35b_a3b/model.yaml" @@ -113,6 +114,7 @@ def test_setup_defaults_example_is_portable() -> None: def test_model_examples_use_public_hugging_face_identities() -> None: paths = ( NEMOTRON3_NANO_30B_MODEL_CONFIG, + QWEN3P5_0P8B_MODEL_CONFIG, QWEN3P5_9B_MODEL_CONFIG, QWEN3P6_35B_A3B_MODEL_CONFIG, ) diff --git a/tests/unit/torch/puzzletron/test_qwen3p5_0p8b_example.py b/tests/unit/torch/puzzletron/test_qwen3p5_0p8b_example.py new file mode 100644 index 00000000000..b991954b421 --- /dev/null +++ b/tests/unit/torch/puzzletron/test_qwen3p5_0p8b_example.py @@ -0,0 +1,135 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""CPU contracts for the Qwen 3.5 0.8B model example.""" + +from pathlib import Path + +import yaml +from hydra import compose, initialize_config_dir +from omegaconf import OmegaConf + +REPOSITORY_ROOT = Path(__file__).resolve().parents[4] +CONFIG_ROOT = REPOSITORY_ROOT / "examples/puzzletron/configs" +MODEL_PATH = ( + REPOSITORY_ROOT / "examples/puzzletron/configs/families/qwen3_5/qwen3p5_0p8b/model.yaml" +) +ADVANCED_PATH = ( + REPOSITORY_ROOT / "examples/puzzletron/configs/families/qwen3_5/qwen3p5_0p8b/advanced.yaml" +) + + +def test_qwen3p5_0p8b_model_identity_and_geometry_are_pinned() -> None: + model = yaml.safe_load(MODEL_PATH.read_text()) + + assert model["input_hf_model_path"] == "Qwen/Qwen3.5-0.8B" + assert model["model_info"] == { + "hf_repo": model["input_hf_model_path"], + "hf_revision": "2fc06364715b967f1860aea9cf38778875588b17", + "model_type": "qwen3_5", + "architectures": ["Qwen3_5ForConditionalGeneration"], + "num_hidden_layers": 24, + "hidden_size": 1024, + "intermediate_size": 3584, + "num_attention_heads": 8, + "num_key_value_heads": 2, + "head_dim": 256, + "vocab_size": 248320, + "tie_word_embeddings": True, + "max_position_embeddings": 262144, + "mtp_num_hidden_layers": 1, + "layer_counts": {"linear_attention": 18, "full_attention": 6}, + "mamba": { + "linear_key_head_dim": 128, + "linear_num_key_heads": 16, + "linear_num_value_heads": 16, + "linear_value_head_dim": 128, + "linear_conv_kernel_dim": 4, + }, + } + + +def test_qwen3p5_0p8b_default_search_matches_tracked_runtime_campaign() -> None: + model = yaml.safe_load(MODEL_PATH.read_text()) + + assert model["pruning"] == {"intermediate_size_list": [3072, 2048]} + assert model["search_space"]["axes"] == { + "ffn_intermediate": { + "enabled": True, + "teacher_value": 3584, + "values": [3072, 2048], + } + } + + +def test_qwen3p5_0p8b_advanced_search_keeps_broad_domains_explicit() -> None: + advanced = yaml.safe_load(ADVANCED_PATH.read_text()) + axes = advanced["search_space"]["axes"] + + expected_enabled_domains = { + "hidden_width": (1024, [768]), + "kv_groups": (2, [1]), + "q_heads_per_group": (4, [2]), + "ffn_intermediate": (3584, [3072, 2560, 2048, 1792, 1536]), + "gdn_key_groups": (16, [12, 8]), + "gdn_value_head_dim": (128, [96]), + } + enabled_domains = { + axis_id: (axis["teacher_value"], axis["values"]) + for axis_id, axis in axes.items() + if axis["enabled"] + } + + assert advanced["pruning"] == { + "intermediate_size_list": [3072, 2560, 2048, 1792, 1536], + "attn_heads_list": [[2, 1], [4, 1], [4, 2], [8, 2]], + } + assert enabled_domains == expected_enabled_domains + assert { + axis_id: axes[axis_id] for axis_id in ("gdn_value_heads_per_group", "gdn_key_head_dim") + } == { + "gdn_value_heads_per_group": { + "enabled": False, + "teacher_value": 1, + "values": [], + }, + "gdn_key_head_dim": { + "enabled": False, + "teacher_value": 128, + "values": [], + }, + } + assert set(axes) == { + *expected_enabled_domains, + "gdn_value_heads_per_group", + "gdn_key_head_dim", + } + + +def test_qwen3p5_0p8b_advanced_search_composes_the_pinned_model() -> None: + with initialize_config_dir(version_base=None, config_dir=str(CONFIG_ROOT)): + config = compose(config_name="families/qwen3_5/qwen3p5_0p8b/advanced") + config = OmegaConf.to_container(config, resolve=True) + + assert config["input_hf_model_path"] == "Qwen/Qwen3.5-0.8B" + assert config["model_info"]["hf_revision"] == "2fc06364715b967f1860aea9cf38778875588b17" + assert config["model"]["revision"] == config["model_info"]["hf_revision"] + assert config["search_space"]["axes"]["ffn_intermediate"]["values"] == [ + 3072, + 2560, + 2048, + 1792, + 1536, + ]