diff --git a/examples/hf_ptq/example_utils.py b/examples/hf_ptq/example_utils.py index 1b0a230e66d..ed2f0961aa8 100755 --- a/examples/hf_ptq/example_utils.py +++ b/examples/hf_ptq/example_utils.py @@ -57,15 +57,7 @@ snapshot_download = None from modelopt.torch.utils import distributed as dist_utils -from modelopt.torch.utils.mlflow import ( - EXPERIMENT_JSON, - MlflowRunLogger, - checkpoint_run_tags, - resolved_recipe_texts, - track_run, -) -from modelopt.torch.utils.mlflow import add_mlflow_args as _add_mlflow_args -from modelopt.torch.utils.mlflow import resolve_mlflow_args as _resolve_mlflow_args +from modelopt.torch.utils.mlflow import EXPERIMENT_JSON, Tool, resolved_recipe_texts, tracked_run logger = logging.getLogger(__name__) @@ -1228,103 +1220,40 @@ def set_layerwise_export_dir(quant_cfg: dict, export_path: str) -> dict: return quant_cfg -def add_mlflow_args(parser: argparse.ArgumentParser) -> None: - """Add the MLflow tracking flags.""" - _add_mlflow_args( - parser, - "hf_ptq", - tracks=( - "Track this run on an MLflow server (e.g. https:///), " - "uploading the command, the resolved recipe, the run log and the quantization " - "summaries, and writing .experiment.json into --export_path so the checkpoint " - "names the run that produced it." - ), - variant_help="recipe name, or --qformat if no --recipe", - ) - - -def resolve_mlflow_args(args: argparse.Namespace, parser: argparse.ArgumentParser) -> None: - """Settle where tracking is configured from, and name the experiment.""" - _resolve_mlflow_args( - args, - parser, - tool="hf_ptq", - model=args.pyt_ckpt_path, - variant=Path(args.recipe).stem if args.recipe else args.qformat, - ) - - -_MLFLOW_NON_PARAM_ARGS = frozenset( - { - "checkpoint_exported", - "dist_state", - "mlflow", - "mlflow_experiment", - "mlflow_required", - "mlflow_run_name", - } +HF_PTQ = Tool( + name="hf_ptq", + tracks=( + "Track this run on an MLflow server (e.g. https:///), " + "uploading the command, the resolved recipe, the run log and the quantization " + "summaries, and writing .experiment.json into --export_path so the checkpoint " + "names the run that produced it." + ), + variant_help="recipe name, or --qformat if no --recipe", + variant=lambda args: Path(args.recipe).stem if args.recipe else args.qformat, + model=lambda args: args.pyt_ckpt_path, + checkpoint=lambda args: args.export_path, + texts=lambda args: resolved_recipe_texts(args.recipe), + # Missing entries are skipped: the MoE table only exists for MoE models, and neither + # file is written under --no-verbose. + outputs=lambda args: { + "summary/quant_summary.txt": Path(args.export_path) / ".quant_summary.txt", + "summary/moe.html": Path(args.export_path) / ".moe.html", + }, + # dist_state is an object rather than a setting, and checkpoint_exported is this + # script's own bookkeeping. + non_params=frozenset({"dist_state", "checkpoint_exported"}), ) -def _mlflow_run_inputs(args: argparse.Namespace) -> tuple[dict, dict]: - """Params and start-time artifacts describing this PTQ run.""" - params = {k: v for k, v in vars(args).items() if k not in _MLFLOW_NON_PARAM_ARGS} - # dist_state is an object, so record the one field worth searching on. - params["world_size"] = args.dist_state.world_size - return params, resolved_recipe_texts(args.recipe) - - -def _mlflow_logger(args: argparse.Namespace) -> MlflowRunLogger: - """Build this run's logger; inert unless --mlflow was given and this is the main rank.""" - return MlflowRunLogger( - args.mlflow, - args.mlflow_experiment, - run_name=args.mlflow_run_name, - enabled=bool(args.mlflow) and args.dist_state.is_main, - required=args.mlflow_required, - ) - - -def _mlflow_describe(args: argparse.Namespace) -> dict: - """Everything the run uploads, gathered once -- reading the recipe twice would print a - second "[load_recipe] loading:" line on every tracked run.""" - params, texts = _mlflow_run_inputs(args) - return { - "params": params, - "tags": _mlflow_run_tags(args), - "texts": texts, - "files": _mlflow_run_outputs(args), - } - - @contextmanager def mlflow_run(args: argparse.Namespace) -> Iterator[None]: """Track this invocation for the duration of the block; see - :func:`~modelopt.torch.utils.mlflow.track_run`.""" - with track_run( - _mlflow_logger(args), - args.export_path, + :func:`~modelopt.torch.utils.mlflow.tracked_run`.""" + with tracked_run( + args, + HF_PTQ, is_main=args.dist_state.is_main, exported=lambda: args.checkpoint_exported, - describe=lambda: _mlflow_describe(args), + world_size=args.dist_state.world_size, ): yield - - -def _mlflow_run_tags(args: argparse.Namespace) -> dict[str, str]: - """This run's shared join keys, from the arguments that name its input and output.""" - return checkpoint_run_tags(args.pyt_ckpt_path, args.export_path) - - -def _mlflow_run_outputs(args: argparse.Namespace) -> dict[str, Path]: - """Summaries written by post_quantize, keyed by artifact path. - - Uploaded without the leading dot, which is awkward to browse in the MLflow UI. Missing - entries are skipped: the MoE table only exists for MoE models, and neither file is - written under ``--no-verbose``. - """ - export_path = Path(args.export_path) - return { - "summary/quant_summary.txt": export_path / ".quant_summary.txt", - "summary/moe.html": export_path / ".moe.html", - } diff --git a/examples/hf_ptq/hf_ptq.py b/examples/hf_ptq/hf_ptq.py index 1266857d844..347f6dc86c7 100755 --- a/examples/hf_ptq/hf_ptq.py +++ b/examples/hf_ptq/hf_ptq.py @@ -28,8 +28,8 @@ from cast_mxfp4_to_nvfp4 import apply_to_model as apply_cast_mxfp4_to_nvfp4 from cast_mxfp4_to_nvfp4 import force_weight_quantizers_static from example_utils import ( + HF_PTQ, _resolve_model_path, - add_mlflow_args, build_quant_cfg, cleanup_distributed, copy_custom_model_files, @@ -45,7 +45,6 @@ needs_checkpoint_path_update, recipe_layerwise_blocks, resolve_checkpoint_dir, - resolve_mlflow_args, run_nemotron_vl_preview, save_processor_config, save_source_config, @@ -101,6 +100,7 @@ get_supported_datasets, ) from modelopt.torch.utils.memory_monitor import launch_memory_monitor +from modelopt.torch.utils.mlflow import add_mlflow_args, resolve_mlflow_args from modelopt.torch.utils.plugins.model_load_utils import parallel_load_and_prepare_fsdp2 from modelopt.torch.utils.speech_dataset_utils import get_speech_dataset_dataloader from modelopt.torch.utils.vlm_dataset_utils import get_vlm_dataset_dataloader @@ -1750,13 +1750,13 @@ def parse_args() -> argparse.Namespace: ), ) - add_mlflow_args(parser) + add_mlflow_args(parser, HF_PTQ) args = parser.parse_args() # Flipped by export_quantized once a checkpoint is actually on disk. The MLflow pointer # is gated on it rather than on --export_path existing, which proves nothing. args.checkpoint_exported = False - resolve_mlflow_args(args, parser) + resolve_mlflow_args(args, parser, HF_PTQ) if args.moe_calib_experts_ratio is not None and not (0.0 < args.moe_calib_experts_ratio <= 1.0): parser.error("--moe_calib_experts_ratio must be in the range (0.0, 1.0].") diff --git a/examples/megatron_bridge/mlflow_utils.py b/examples/megatron_bridge/mlflow_utils.py index 9782c764ca0..392e9e620b5 100644 --- a/examples/megatron_bridge/mlflow_utils.py +++ b/examples/megatron_bridge/mlflow_utils.py @@ -13,125 +13,35 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""MLflow tracking for ``quantize.py``, mirroring ``examples/hf_ptq``. +"""Shared MLflow wiring for the Megatron-Bridge examples, mirroring ``examples/hf_ptq``. -Every rank parses and validates the same flags, so a typo in the URI fails identically -everywhere instead of on one rank while the others wait in a collective. Only the master rank -opens a run, so the log capture and the uploads happen once. +Each script declares what it records as a :class:`~modelopt.torch.utils.mlflow.Tool` beside +its own flags; this module knows none of them, only how a run is opened and closed here. -Nothing here imports Megatron, so the tracking can be exercised without it. +Every rank parses the same flags, so a typo in the URI fails identically everywhere rather +than on one rank while the others wait in a collective. """ import argparse from collections.abc import Iterator from contextlib import contextmanager -from pathlib import Path import modelopt.torch.utils.distributed as dist -from modelopt.torch.utils.mlflow import ( - MlflowRunLogger, - checkpoint_run_tags, - resolved_recipe_texts, - track_run, -) -from modelopt.torch.utils.mlflow import add_mlflow_args as _add_mlflow_args -from modelopt.torch.utils.mlflow import resolve_mlflow_args as _resolve_mlflow_args +from modelopt.torch.utils.mlflow import Tool, tracked_run -TOOL_NAME = "megatron_bridge_quantize" - -# The tracking settings describe the destination rather than the quantization, and -# checkpoint_exported is this script's own bookkeeping. -_NON_PARAM_ARGS = frozenset( - { - "checkpoint_exported", - "mlflow", - "mlflow_experiment", - "mlflow_required", - "mlflow_run_name", - } -) - - -def add_mlflow_args(parser: argparse.ArgumentParser) -> None: - """Add the MLflow tracking flags.""" - _add_mlflow_args( - parser, - TOOL_NAME, - tracks=( - "Track this run on an MLflow server (e.g. https:///), " - "uploading the command, the resolved recipe, the run log and the quantizer " - "summary, and writing .experiment.json into --export_megatron_path so the " - "checkpoint names the run that produced it." - ), - variant_help="recipe name, or --quant_cfg if no --recipe", - ) - - -def resolve_mlflow_args(args: argparse.Namespace, parser: argparse.ArgumentParser) -> None: - """Settle where tracking is configured from, and name the experiment.""" - _resolve_mlflow_args( - args, - parser, - tool=TOOL_NAME, - model=args.hf_model_name_or_path, - # ``or "none"``: neither flag is required by the parser, and the run that reaches - # get_quant_config without one fails there rather than while being named. - variant=Path(args.recipe).stem if args.recipe else (args.quant_cfg or "none"), - ) - - -def _run_inputs(args: argparse.Namespace) -> tuple[dict, dict]: - """Params and start-time artifacts describing this PTQ run.""" - params = {k: v for k, v in vars(args).items() if k not in _NON_PARAM_ARGS} - # The parallelism flags say how the run was laid out but not how many GPUs it took: - # data parallelism is implicit in the launcher's world size. - params["world_size"] = dist.size() - return params, resolved_recipe_texts(args.recipe) - - -def _run_tags(args: argparse.Namespace) -> dict[str, str]: - """This run's shared join keys, from the arguments that name its input and output.""" - return checkpoint_run_tags(args.hf_model_name_or_path, args.export_megatron_path) - - -def _run_outputs(args: argparse.Namespace) -> dict[str, Path]: - """Summaries written beside the checkpoint, keyed by artifact path. - - Uploaded without the leading dot, which is awkward to browse in the MLflow UI. A missing - entry is skipped: the summary is written by the master rank only once quantization - has finished. - """ - return {"summary/quant_summary.txt": Path(args.export_megatron_path) / ".quant_summary.txt"} - - -def _describe(args: argparse.Namespace) -> dict: - """Everything the run uploads, gathered once -- reading the recipe twice would print a - second "[load_recipe] loading:" line on every tracked run.""" - params, texts = _run_inputs(args) - return { - "params": params, - "tags": _run_tags(args), - "texts": texts, - "files": _run_outputs(args), - } +# These scripts' own bookkeeping, on top of the tracking settings the library already keeps +# out of the params. Every Tool here passes it as ``non_params``. +NON_PARAMS = frozenset({"checkpoint_exported"}) @contextmanager -def mlflow_run(args: argparse.Namespace) -> Iterator[None]: - """Track this invocation for the duration of the block; see - :func:`~modelopt.torch.utils.mlflow.track_run`.""" - logger = MlflowRunLogger( - args.mlflow or "", - args.mlflow_experiment, - run_name=args.mlflow_run_name, - enabled=bool(args.mlflow) and dist.is_master(), - required=args.mlflow_required, - ) - with track_run( - logger, - args.export_megatron_path, +def mlflow_run(args: argparse.Namespace, tool: Tool) -> Iterator[None]: + """Track this invocation for the duration of the block.""" + with tracked_run( + args, + tool, is_main=dist.is_master(), exported=lambda: args.checkpoint_exported, - describe=lambda: _describe(args), + world_size=dist.size(), ): yield diff --git a/examples/megatron_bridge/quantize.py b/examples/megatron_bridge/quantize.py index 2b3e8e47fa1..a35f4d0c7f3 100644 --- a/examples/megatron_bridge/quantize.py +++ b/examples/megatron_bridge/quantize.py @@ -60,10 +60,11 @@ import argparse import copy import gc +from pathlib import Path import torch from megatron.bridge.models.hf_pretrained.utils import is_safe_repo -from mlflow_utils import add_mlflow_args, mlflow_run, resolve_mlflow_args +from mlflow_utils import NON_PARAMS, mlflow_run from transformers import AutoProcessor import modelopt.torch.quantization as mtq @@ -77,7 +78,13 @@ ) from modelopt.torch.utils import print_args, print_rank_0, warn_rank_0 from modelopt.torch.utils.dataset_utils import get_supported_datasets -from modelopt.torch.utils.mlflow import masked_args +from modelopt.torch.utils.mlflow import ( + Tool, + add_mlflow_args, + masked_args, + resolve_mlflow_args, + resolved_recipe_texts, +) from modelopt.torch.utils.plugins.mbridge import ( get_language_model, load_mbridge_model_from_hf, @@ -104,6 +111,27 @@ # single fixed --quant_cfg / --recipe to the whole model. +QUANTIZE = Tool( + name="megatron_bridge_quantize", + tracks=( + "Track this run on an MLflow server, uploading the command, the resolved recipe, the " + "run log and the quantizer summary, and writing .experiment.json into " + "--export_megatron_path." + ), + variant_help="recipe name, or --quant_cfg if no --recipe", + # ``or "none"``: neither flag is required, and a run without one fails in get_quant_config + # rather than while being named. + variant=lambda args: Path(args.recipe).stem if args.recipe else (args.quant_cfg or "none"), + model=lambda args: args.hf_model_name_or_path, + checkpoint=lambda args: args.export_megatron_path, + texts=lambda args: resolved_recipe_texts(args.recipe), + outputs=lambda args: { + "summary/quant_summary.txt": Path(args.export_megatron_path) / ".quant_summary.txt" + }, + non_params=NON_PARAMS, +) + + def get_args() -> argparse.Namespace: parser = argparse.ArgumentParser(formatter_class=argparse.ArgumentDefaultsHelpFormatter) parser.add_argument("--hf_model_name_or_path", type=str, required=True) @@ -224,10 +252,10 @@ def get_args() -> argparse.Namespace: help="Skip the post-quantization generation sanity check.", ) - add_mlflow_args(parser) + add_mlflow_args(parser, QUANTIZE) args = parser.parse_args() - resolve_mlflow_args(args, parser) + resolve_mlflow_args(args, parser, QUANTIZE) print_args(masked_args(args)) @@ -475,7 +503,7 @@ def forward_loop(_model=None): try: # Entered inside the try: opening the run is fatal by design, and the peers of a rank # that exits without dist.abort() stay blocked on the first collective. - with mlflow_run(args): + with mlflow_run(args, QUANTIZE): main(args) except BaseException: dist.abort() # peers may be stuck in a collective this rank will never reach diff --git a/examples/vllm_serve/vllm_mlflow_utils.py b/examples/vllm_serve/vllm_mlflow_utils.py index 382f0b8c70e..68f2c74d459 100644 --- a/examples/vllm_serve/vllm_mlflow_utils.py +++ b/examples/vllm_serve/vllm_mlflow_utils.py @@ -46,6 +46,7 @@ from modelopt.torch.utils.mlflow import ( TRACKING_URI_ENV, MlflowRunLogger, + Tool, command_text, default_experiment_name, resolve_tracking_uri, @@ -84,20 +85,27 @@ ) +# This example serves a checkpoint rather than producing one, so the Tool carries only what +# names the run: there is no output to point at and no chain to join. +TOOL = Tool( + name=TOOL_NAME, + tracks=( + "Track this server's calibration on an MLflow server " + "(e.g. https:///), uploading the command, the resolved " + "recipe, the quantization config actually applied, the worker log and the " + "quantizer summary. This is the quantization tracking server, which is " + "unrelated to any tracking server an evaluation harness exports its scores to." + ), + variant_help="recipe name, or the quantization config", + # Read from the environment, which is where this example's quantization settings live. + variant=lambda args: quant_variant(), + model=lambda args: args.model, +) + + def add_mlflow_args(parser: argparse.ArgumentParser) -> None: """Add the MLflow tracking flags to the launcher's parser.""" - _add_mlflow_args( - parser, - TOOL_NAME, - tracks=( - "Track this server's calibration on an MLflow server " - "(e.g. https:///), uploading the command, the resolved " - "recipe, the quantization config actually applied, the worker log and the " - "quantizer summary. This is the quantization tracking server, which is " - "unrelated to any tracking server an evaluation harness exports its scores to." - ), - variant_help="recipe name, or the quantization config", - ) + _add_mlflow_args(parser, TOOL) def resolve_mlflow_args(args: argparse.Namespace, parser: argparse.ArgumentParser) -> None: @@ -119,7 +127,7 @@ def resolve_mlflow_args(args: argparse.Namespace, parser: argparse.ArgumentParse os.environ[EXPERIMENT_ENV] = ( args.mlflow_experiment or os.environ.get(EXPERIMENT_ENV) - or default_experiment_name(TOOL_NAME, args.model, quant_variant()) + or default_experiment_name(TOOL.name, TOOL.model(args), TOOL.variant(args)) ) if args.mlflow_run_name: os.environ[RUN_NAME_ENV] = args.mlflow_run_name diff --git a/modelopt/torch/utils/mlflow.py b/modelopt/torch/utils/mlflow.py index 8a7f5bbd4de..eb264055e8b 100644 --- a/modelopt/torch/utils/mlflow.py +++ b/modelopt/torch/utils/mlflow.py @@ -37,6 +37,7 @@ import warnings from collections.abc import Callable, Iterator, Mapping from contextlib import contextmanager +from dataclasses import dataclass, field from datetime import datetime, timezone from pathlib import Path, PurePosixPath from typing import Any @@ -51,18 +52,21 @@ "EXPERIMENT_JSON", "TRACKING_URI_ENV", "MlflowRunLogger", + "Tool", "add_mlflow_args", - "checkpoint_run_tags", "command_text", "current_user", "default_experiment_name", + "default_run_name", + "describe_run", "drop_experiment_json", "mask_tracking_uri", "masked_args", "resolve_mlflow_args", "resolve_tracking_uri", "resolved_recipe_texts", - "track_run", + "run_tags", + "tracked_run", "validate_tracking_uri", ] @@ -89,6 +93,23 @@ TRACKING_URI_ENV = "MLFLOW_TRACKING_URI" +def _experiment_json( + tracking_uri: str, experiment_name: str, info: Any, run_name: str | None = None +) -> dict[str, str]: + """The provenance record's fields, read off the run the server returned.""" + uri = _redact(tracking_uri).rstrip("/") + experiment_id = str(info.experiment_id) + run_id = str(info.run_id) + return { + "tracking_uri": uri, + "experiment_name": experiment_name, + "experiment_id": experiment_id, + "run_id": run_id, + "run_name": getattr(info, "run_name", None) or run_name or "", + "run_url": f"{uri}/#/experiments/{experiment_id}/runs/{run_id}", + } + + def _stat_key(path: Path) -> tuple[int, int] | None: """Identity of a file's contents-in-time, or ``None`` when it does not exist.""" try: @@ -169,6 +190,15 @@ def default_experiment_name(tool: str, model: str, variant: str, user: str | Non return name[:_MAX_NAME_LEN] +def default_run_name() -> str: + """The UTC start time as ``YYYYmmdd-HHMMSS``, which is what the flags document. + + Used by :class:`MlflowRunLogger` and by a caller handing the name to something else that + opens the run, so both honour the documented default rather than MLflow's random one. + """ + return datetime.now(timezone.utc).strftime("%Y%m%d-%H%M%S") + + def current_user() -> str: """Return the current username, or ``"unknown"`` if the uid has no passwd entry.""" try: @@ -238,6 +268,31 @@ def command_text(argv: list[str] | None = None) -> str: return "\n".join(lines) + "\n" +def _ask(what: str, callback: Callable[[], Any], default: Any) -> Any: + """Read what a run reports about itself, never at the cost of the caller's exception.""" + try: + return callback() + except Exception as e: + print(f"[mlflow] WARNING: could not read this run's {what}: {e}") + return default + + +@contextmanager +def _closing_run(finish: Callable[[str], None]) -> Iterator[None]: + """Run the block, then *finish* the run with the status the block earned.""" + status = "FAILED" + try: + yield + status = "FINISHED" + except SystemExit as e: + # A script that ends by exiting -- Megatron-Bridge does, from inside its training + # loop -- finished if it exited cleanly. + status = "FINISHED" if e.code in (0, None) else "FAILED" + raise + finally: + finish(status) + + class MlflowRunLogger: """Record one script invocation as an MLflow run. @@ -303,11 +358,7 @@ def __init__( @property def run_url(self) -> str: """Link to this run in the MLflow UI, or ``""`` before the run is open.""" - if self._run is None: - return "" - info = self._run.info - uri = _redact(self.tracking_uri) - return f"{uri}/#/experiments/{info.experiment_id}/runs/{info.run_id}" + return self.run_info.get("run_url", "") @property def run_info(self) -> dict[str, str]: @@ -320,15 +371,9 @@ def run_info(self) -> dict[str, str]: """ if self._run is None: return {} - info = self._run.info - return { - "tracking_uri": _redact(self.tracking_uri), - "experiment_name": self.experiment_name, - "experiment_id": str(info.experiment_id), - "run_id": str(info.run_id), - "run_name": getattr(info, "run_name", None) or self.run_name or "", - "run_url": self.run_url, - } + return _experiment_json( + self.tracking_uri, self.experiment_name, self._run.info, self.run_name + ) def start( self, @@ -393,12 +438,8 @@ def track( ... quantize_and_export() """ self.start(params=params, tags=tags, texts=texts, files=files) - status = "FAILED" - try: + with _closing_run(lambda status: self.finish(status, files=files, metrics=metrics)): yield self - status = "FINISHED" - finally: - self.finish(status, files=files, metrics=metrics) def log_text(self, artifact_path: str, text: str) -> None: """Upload *text* as an artifact while the run is open, best-effort. @@ -516,7 +557,7 @@ def _open_run(self) -> None: mlflow.set_experiment(self.experiment_name) # Settled here rather than passed straight through, so run_info reports the name the # run actually carries. - self.run_name = self.run_name or datetime.now(timezone.utc).strftime("%Y%m%d-%H%M%S") + self.run_name = self.run_name or default_run_name() self._run = mlflow.start_run(run_name=self.run_name) print(f"[mlflow] experiment: {self.experiment_name}\n[mlflow] run: {self.run_url}") @@ -633,6 +674,38 @@ def _stop_capture(self) -> None: # The CLI surface below is shared by the example scripts that offer tracking, so a run is # configured the same way and named by the same convention whichever script opened it. + +@dataclass(frozen=True) +class Tool: + """What distinguishes one script's tracking from another's. + + A script declares one of these and the functions below do the rest. Each callable reads + the arguments naming this run's input, output and what it consumed -- *source* is what + the next run in a chain joins on, *variant* names what this run did. *settles_pointer* is + false when something other than :func:`tracked_run` writes the provenance pointer. + """ + + name: str + tracks: str + variant_help: str + variant: Callable[[argparse.Namespace], str] + model: Callable[[argparse.Namespace], str] + checkpoint: Callable[[argparse.Namespace], str | None] = field(default=lambda args: None) + source: Callable[[argparse.Namespace], str] | None = None + texts: Callable[[argparse.Namespace], dict[str, str]] = field(default=lambda args: {}) + outputs: Callable[[argparse.Namespace], dict[str, Path]] = field(default=lambda args: {}) + # Read on the way *out*, so it can report something the run computed -- a pruning score, + # say -- which the script stashes on its own namespace. + metrics: Callable[[argparse.Namespace], dict[str, float]] = field(default=lambda args: {}) + non_params: frozenset[str] = frozenset() + settles_pointer: bool = True + + +# The tracking settings describe the destination rather than the work, so they are never +# params; a script adds its own bookkeeping through ``Tool.non_params``. +_NEVER_PARAMS = frozenset({"mlflow", "mlflow_experiment", "mlflow_required", "mlflow_run_name"}) + + _ENV_HELP = ( f"MLflow's own ${TRACKING_URI_ENV} enables tracking without this flag, which overrides " "it. A URI taken from the environment is best-effort: if it is unusable the run warns and " @@ -645,18 +718,12 @@ def _stop_capture(self) -> None: ) -def add_mlflow_args( - parser: argparse.ArgumentParser, - tool: str, - tracks: str = _TRACKS_HELP, - variant_help: str = "recipe name, or the quantization format", -) -> None: +def add_mlflow_args(parser: argparse.ArgumentParser, tool: Tool) -> None: """Add ``--mlflow``, ``--mlflow_experiment`` and ``--mlflow_run_name`` to *parser*. - *tool* names the script in the default experiment ``//-`` (see - :func:`default_experiment_name`), *tracks* is the leading description of ``--mlflow`` -- - what this particular script uploads -- and *variant_help* says what the script derives the - variant from. Pair with :func:`resolve_mlflow_args`. + The help text comes from *tool*: its ``tracks`` describes what this script uploads and its + ``variant_help`` says what the experiment name's variant is derived from. Pair with + :func:`resolve_mlflow_args`. The multi-word flags are registered under both the underscored and the dashed spelling: vLLM's ``FlexibleArgumentParser`` rewrites every ``--foo_bar`` on the command line to @@ -664,12 +731,15 @@ def add_mlflow_args( reachable there at all, and a user moving between the example scripts should not have to remember which spelling each one took. """ - parser.add_argument("--mlflow", default=None, help=f"{tracks} {_ENV_HELP}") + parser.add_argument("--mlflow", default=None, help=f"{tool.tracks} {_ENV_HELP}") parser.add_argument( "--mlflow_experiment", "--mlflow-experiment", default=None, - help=f"MLflow experiment name. Default: $USER/{tool}/-<{variant_help}>.", + help=( + f"MLflow experiment name. Default: " + f"$USER/{tool.name}/-<{tool.variant_help}>." + ), ) parser.add_argument( "--mlflow_run_name", @@ -724,22 +794,18 @@ def resolve_tracking_uri( def resolve_mlflow_args( - args: argparse.Namespace, - parser: argparse.ArgumentParser, - tool: str, - model: str, - variant: str, + args: argparse.Namespace, parser: argparse.ArgumentParser, tool: Tool ) -> None: """Settle where tracking is configured from, and name the experiment, in place. Sets ``args.mlflow`` to the validated URI or ``None``, ``args.mlflow_required`` to whether - the flag asked for it, and defaults ``args.mlflow_experiment`` from *tool*, *model* and - *variant*. Pair with :func:`add_mlflow_args`. + the flag asked for it, and defaults ``args.mlflow_experiment`` from *tool*. Pair with + :func:`add_mlflow_args`. """ args.mlflow, args.mlflow_required = resolve_tracking_uri(args.mlflow, parser) if args.mlflow: args.mlflow_experiment = args.mlflow_experiment or default_experiment_name( - tool, model, variant + tool.name, tool.model(args), tool.variant(args) ) @@ -770,21 +836,6 @@ def masked_args(args: argparse.Namespace, attr: str = "mlflow") -> argparse.Name return argparse.Namespace(**{**vars(args), attr: mask_tracking_uri(getattr(args, attr, None))}) -def checkpoint_run_tags(source_model: str, checkpoint_dir: Path | str) -> dict[str, str]: - """Tags a quantization run and whatever is later done with the checkpoint it wrote. - - Shared so the two can be found together on one tracking server. ``checkpoint_path`` is - the checkpoint the run *writes*, because that is what an export or an evaluation is later - pointed at (NEL takes ``deployment.checkpoint_path``); the input is kept separately. It is - resolved because a relative path is useless as a join key. - """ - return { - "model": Path(source_model).name, - "checkpoint_path": str(Path(checkpoint_dir).resolve()), - "source_checkpoint_path": source_model, - } - - def resolved_recipe_texts(recipe: str | None) -> dict[str, str]: r"""``{artifact path: content}`` for *recipe*, or ``{}`` when the run used none. @@ -801,40 +852,103 @@ def resolved_recipe_texts(recipe: str | None) -> dict[str, str]: return {"recipe/resolved_recipe.yaml": yaml.safe_dump(resolved, sort_keys=False)} +def run_tags(args: argparse.Namespace, tool: Tool) -> dict[str, str]: + """This run's join keys, shared with whatever is later done with what it produced. + + ``checkpoint_path`` is the checkpoint the run *writes*, because that is what an export or + an evaluation is later pointed at (NEL takes ``deployment.checkpoint_path``), and + ``source_checkpoint_path`` is what it consumed, so a chain of runs joins on the pair: a + distillation's source is the checkpoint it continues from, not the model that was + quantized. Both are resolved, since a relative path is useless as a join key -- except a + source that names no directory, such as a Hub ``org/name`` id. + """ + source = tool.source(args) if tool.source else tool.model(args) + checkpoint = tool.checkpoint(args) + tags = { + "model": Path(tool.model(args)).name, + "source_checkpoint_path": ( + str(Path(source).resolve()) if os.path.exists(source) else str(source) + ), + } + # Omitted rather than empty when the run writes no checkpoint: a search for runs that + # produced one should not match it. + if checkpoint is not None: + tags["checkpoint_path"] = str(Path(checkpoint).resolve()) + return tags + + +def describe_run(args: argparse.Namespace, tool: Tool, world_size: int = 1) -> dict: + """The keyword arguments :meth:`MlflowRunLogger.track` takes, for this run. + + Every command-line argument becomes a searchable param, so a flag added later is tracked + without touching this. *world_size* is recorded separately because the parallelism flags + say how a run was laid out but not how many processes it took. + """ + skip = _NEVER_PARAMS | tool.non_params + params = {k: v for k, v in vars(args).items() if k not in skip} + params["world_size"] = world_size + return { + "params": params, + "tags": run_tags(args, tool), + "texts": tool.texts(args), + "files": tool.outputs(args), + } + + @contextmanager -def track_run( - logger: MlflowRunLogger, - checkpoint_dir: Path | str, +def tracked_run( + args: argparse.Namespace, + tool: Tool, is_main: bool, exported: Callable[[], bool], - describe: Callable[[], Mapping[str, Any]] | None = None, + world_size: int = 1, ) -> Iterator[MlflowRunLogger]: - """Track a checkpoint-producing run, keeping its provenance pointer honest either way. - - *logger* is inert unless tracking was configured *and* this is the rank that records it, - so the caller needs no branching. *checkpoint_dir* is where the run writes its checkpoint - and *is_main* gates writes every rank would otherwise race on. *exported* is read on the - way out, not on the way in: only a completed export may claim the checkpoint the pointer - sits next to, since the directory usually exists before the weights do. + """Track one invocation of *tool* for the duration of the block. - *describe* returns the keyword arguments for :meth:`MlflowRunLogger.track` (``params``, - ``tags``, ``texts``, ``files``) and is called only when the run is tracked, so an - untracked run does not pay for gathering them -- re-reading a recipe, say. + Inert unless ``--mlflow`` settled a URI and this is the rank that records it, so the + caller needs no branching. *is_main* is that rank, and also gates the writes every rank + would otherwise race on; *exported* is read on the way out, once the run knows whether it + wrote the checkpoint its pointer would claim. Example: - >>> with track_run(logger, args.export_path, is_main, lambda: args.exported, describe): + >>> with tracked_run(args, HF_PTQ, is_main, lambda: args.exported, world_size): ... quantize_and_export(args) """ - path = Path(checkpoint_dir) + logger = MlflowRunLogger( + args.mlflow or "", + args.mlflow_experiment, + run_name=args.mlflow_run_name, + enabled=bool(args.mlflow) and is_main, + required=args.mlflow_required, + ) + # None when the run writes no checkpoint at all -- a pruning run that only scores, say, + # or a script that points each of several checkpoints at the run itself -- so there is + # nothing to point at and nothing that could inherit a stale pointer. + path = None + if tool.settles_pointer and (checkpoint := tool.checkpoint(args)) is not None: + path = Path(checkpoint) if not logger.enabled: + # Gathering the inputs re-reads the recipe, so keep it off the untracked path. try: yield logger finally: - if exported() and is_main: + if path is not None and is_main and _ask("exported flag", exported, False): drop_experiment_json(path) return - with logger.track(**(describe() if describe is not None else {})): - try: - yield logger - finally: - logger.log_experiment_json(path if exported() else None) + + described = describe_run(args, tool, world_size) + logger.start(**described) + + def close(status: str) -> None: + # Only a completed export may claim the checkpoint the pointer sits next to: the + # directory usually exists before the weights do. + wrote_it = path is not None and _ask("exported flag", exported, False) + logger.log_experiment_json(path if wrote_it else None) + logger.finish( + status, + files=described["files"], + metrics=_ask("metrics", lambda: tool.metrics(args), {}), + ) + + with _closing_run(close): + yield logger diff --git a/tests/_test_utils/mlflow.py b/tests/_test_utils/mlflow.py new file mode 100644 index 00000000000..dc1852f39cd --- /dev/null +++ b/tests/_test_utils/mlflow.py @@ -0,0 +1,118 @@ +# 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. + +"""Shared stand-ins for the MLflow client, used by every suite that exercises tracking. + +One fake rather than one per suite: four copies drifted, and a recording method that was a +no-op in one of them made a test pass while asserting nothing. +""" + +import getpass +from pathlib import Path +from types import SimpleNamespace + +import pytest + + +class FakeMlflow: + """Stand-in for the ``mlflow`` module: no server, no dependency, records every call.""" + + def __init__(self): + self.tracking_uri = None + self.experiment = None + self.run_name = None + self.status = None + self.params = {} + self.tags = {} + self.texts = {} + self.metrics = {} + # {uploaded name: (artifact path, contents)} + self.artifacts = {} + # What the server says the run is called, which need not be what was requested. + self.server_run_name = None + # Runs the fluent API opened by itself, which is always a bug in the caller. + self.strays = 0 + self._run = None + + def set_tracking_uri(self, uri): + self.tracking_uri = uri + + def set_experiment(self, name): + self.experiment = name + + def start_run(self, run_name=None, tags=None, description=None): + self.run_name = run_name + # Replaced, not merged: a run starts with only the tags it was opened with, so an + # earlier run's cannot satisfy an assertion about this one. set_tags adds to these. + self.tags = dict(tags or {}) + self._run = SimpleNamespace( + info=SimpleNamespace( + experiment_id="7", run_id="deadbeef", run_name=self.server_run_name or run_name + ) + ) + return self._run + + def _get_or_start_run(self): + """What every fluent call does first: with nothing active, it opens a run of its own.""" + if self._run is None: + self.strays += 1 + self._run = SimpleNamespace( + info=SimpleNamespace(experiment_id="7", run_id="stray", run_name=None) + ) + + def log_params(self, params): + self._get_or_start_run() + self.params.update(params) + + def set_tags(self, tags): + self._get_or_start_run() + self.tags.update(tags) + + def log_text(self, text, artifact_file): + self._get_or_start_run() + self.texts[artifact_file] = text + + def log_artifact(self, local_path, artifact_path=None): + self._get_or_start_run() + self.artifacts[Path(local_path).name] = (artifact_path, Path(local_path).read_text()) + + def log_metrics(self, metrics): + self._get_or_start_run() + self.metrics.update(metrics) + + def end_run(self, status=None): + self.status = status + self._run = None + + +def pin_tracking_env(monkeypatch): + """Pin what the tracking reads from the environment; see :func:`clean_env`.""" + monkeypatch.setattr(getpass, "getuser", lambda: "tester") + for name in ("MLFLOW_TRACKING_URI", "MLFLOW_TRACKING_USERNAME", "MLFLOW_TRACKING_PASSWORD"): + # setenv before delenv: monkeypatch records nothing for a variable that was already + # absent, so a test whose code *sets* one would otherwise leave it behind for the + # rest of the session. + monkeypatch.setenv(name, "") + monkeypatch.delenv(name) + + +@pytest.fixture(autouse=True) +def clean_env(monkeypatch): + """Pin what the tracking reads from the environment. + + A developer shell that exports $MLFLOW_TRACKING_URI -- exactly the population this feature + is built for -- would otherwise flip the branch under test. Tests that want it set it. + """ + pin_tracking_env(monkeypatch) diff --git a/tests/conftest.py b/tests/conftest.py index 3109c3ecd24..af77546a244 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -15,12 +15,14 @@ import os import platform +import sys from pathlib import Path import pytest import torch import torch.distributed as dist from _test_utils.fs_utils import assert_unmodified_tree +from _test_utils.mlflow import FakeMlflow, pin_tracking_env from _test_utils.torch.distributed.utils import init_process import modelopt.torch.opt as mto @@ -179,3 +181,14 @@ def tiny_wan22_path(tmp_path_factory): pipeline_dir = create_tiny_wan22_pipeline_dir(tmp_path_factory.mktemp("tiny_wan22")) with assert_unmodified_tree(pipeline_dir) as path: yield str(path) + + +@pytest.fixture +def fake_mlflow(monkeypatch): + """Stand in for the ``mlflow`` module; see ``_test_utils.mlflow.FakeMlflow``.""" + # A suite that takes the fake without also importing clean_env would otherwise read the + # developer's own $MLFLOW_TRACKING_URI and flip the branch under test. + pin_tracking_env(monkeypatch) + fake = FakeMlflow() + monkeypatch.setitem(sys.modules, "mlflow", fake) + return fake diff --git a/tests/examples/hf_ptq/test_hf_ptq_args.py b/tests/examples/hf_ptq/test_hf_ptq_args.py index b216fc11eb1..d4811f57496 100644 --- a/tests/examples/hf_ptq/test_hf_ptq_args.py +++ b/tests/examples/hf_ptq/test_hf_ptq_args.py @@ -24,6 +24,7 @@ import pytest import torch import yaml +from _test_utils.mlflow import clean_env # noqa: F401 from _test_utils.torch.transformers_models import get_tiny_qwen3 from modelopt.recipe import load_recipe @@ -31,21 +32,12 @@ from modelopt.recipe.presets import QUANT_CFG_CHOICES, RecipeSupersededAction from modelopt.torch.quantization import tensor_quant from modelopt.torch.quantization.config import QuantizeConfig +from modelopt.torch.utils import mlflow as mlflow_lib +from modelopt.torch.utils.mlflow import describe_run, run_tags _EXAMPLES_DIR = Path(__file__).resolve().parents[3] / "examples" / "hf_ptq" -@pytest.fixture(autouse=True) -def clean_env(monkeypatch): - """Pin what the tracking reads from the environment. - - ``resolve_tracking_uri`` consults $MLFLOW_TRACKING_URI, so a developer shell or runner - that exports it -- exactly the population this feature is built for -- would otherwise - flip the tracked/untracked branch under test. Tests that want the variable set it. - """ - monkeypatch.delenv("MLFLOW_TRACKING_URI", raising=False) - - def _import_hf_ptq(monkeypatch): monkeypatch.syspath_prepend(str(_EXAMPLES_DIR)) return importlib.import_module("hf_ptq") @@ -489,7 +481,7 @@ def test_mlflow_provenance_is_not_logged_as_a_param(monkeypatch, example_utils): hf_ptq, args = _parse_hf_ptq_args(monkeypatch, "--pyt_ckpt_path", "/models/Qwen3-0.6B") args.dist_state = SimpleNamespace(is_main=True, world_size=1) - params, _ = example_utils._mlflow_run_inputs(args) + params = describe_run(args, example_utils.HF_PTQ, args.dist_state.world_size)["params"] assert "mlflow_required" not in params @@ -504,7 +496,8 @@ def test_mlflow_run_inputs_carry_the_resolved_recipe(monkeypatch, example_utils) ) args.dist_state = SimpleNamespace(is_main=True, world_size=1) - params, texts = example_utils._mlflow_run_inputs(args) + described = describe_run(args, example_utils.HF_PTQ, args.dist_state.world_size) + params, texts = described["params"], described["texts"] assert params["pyt_ckpt_path"] == "/models/Qwen3-0.6B" assert params["recipe"] == "general/ptq/nvfp4_default-kv_fp8_cast" @@ -518,7 +511,8 @@ def test_mlflow_run_inputs_omit_the_recipe_when_unused(monkeypatch, example_util hf_ptq, args = _parse_hf_ptq_args(monkeypatch, "--pyt_ckpt_path", "/models/Qwen3-0.6B") args.dist_state = SimpleNamespace(is_main=True, world_size=1) - params, texts = example_utils._mlflow_run_inputs(args) + described = describe_run(args, example_utils.HF_PTQ, args.dist_state.world_size) + params, texts = described["params"], described["texts"] assert texts == {} assert params["recipe"] is None @@ -529,7 +523,7 @@ def test_mlflow_run_outputs_name_the_summaries(monkeypatch, example_utils): monkeypatch, "--pyt_ckpt_path", "/models/Qwen3-0.6B", "--export_path", "/tmp/out" ) - files = example_utils._mlflow_run_outputs(args) + files = example_utils.HF_PTQ.outputs(args) assert files["summary/quant_summary.txt"] == Path("/tmp/out/.quant_summary.txt") assert files["summary/moe.html"] == Path("/tmp/out/.moe.html") @@ -547,12 +541,12 @@ def test_untracked_runs_do_not_gather_mlflow_inputs(monkeypatch, example_utils): ) args.dist_state = SimpleNamespace(is_main=True, world_size=1) calls = [] - monkeypatch.setattr(example_utils, "_mlflow_run_inputs", lambda a: calls.append(a) or ({}, {})) + # Patched on the library, which is where tracked_run resolves it. + monkeypatch.setattr(mlflow_lib, "describe_run", lambda a, t, w=1: calls.append(a) or {}) with example_utils.mlflow_run(args): pass - assert not example_utils._mlflow_logger(args).enabled assert calls == [] @@ -567,12 +561,12 @@ def test_non_main_ranks_do_not_open_a_run(monkeypatch, example_utils): ) args.dist_state = SimpleNamespace(is_main=False, world_size=8) calls = [] - monkeypatch.setattr(example_utils, "_mlflow_run_inputs", lambda a: calls.append(a) or ({}, {})) + # Patched on the library, which is where tracked_run resolves it. + monkeypatch.setattr(mlflow_lib, "describe_run", lambda a, t, w=1: calls.append(a) or {}) with example_utils.mlflow_run(args): pass - assert not example_utils._mlflow_logger(args).enabled assert calls == [] @@ -581,16 +575,20 @@ def test_mlflow_params_track_every_cli_argument(monkeypatch, example_utils): hf_ptq, args = _parse_hf_ptq_args(monkeypatch, "--pyt_ckpt_path", "/models/Qwen3-0.6B") args.dist_state = SimpleNamespace(is_main=True, world_size=4) - params, _ = example_utils._mlflow_run_inputs(args) + params = describe_run(args, example_utils.HF_PTQ, args.dist_state.world_size)["params"] - tracked = set(vars(args)) - example_utils._MLFLOW_NON_PARAM_ARGS + tracked = ( + set(vars(args)) + - example_utils.HF_PTQ.non_params + - {"mlflow", "mlflow_experiment", "mlflow_required", "mlflow_run_name"} + ) assert tracked <= set(params) # The tracking settings describe the destination, not the run, and dist_state is an object. assert not {"mlflow", "mlflow_experiment", "mlflow_run_name", "dist_state"} & set(params) assert params["world_size"] == 4 # A flag added to the parser later is picked up without editing _mlflow_run_inputs. args.some_future_flag = "future" - assert example_utils._mlflow_run_inputs(args)[0]["some_future_flag"] == "future" + assert describe_run(args, example_utils.HF_PTQ)["params"]["some_future_flag"] == "future" def test_mlflow_tags_identify_the_produced_checkpoint(monkeypatch, example_utils, tmp_path): @@ -601,7 +599,7 @@ def test_mlflow_tags_identify_the_produced_checkpoint(monkeypatch, example_utils monkeypatch, "--pyt_ckpt_path", "/models/Qwen3-0.6B", "--export_path", str(export) ) - assert example_utils._mlflow_run_tags(args) == { + assert run_tags(args, example_utils.HF_PTQ) == { "model": "Qwen3-0.6B", "checkpoint_path": str(export), "source_checkpoint_path": "/models/Qwen3-0.6B", @@ -614,52 +612,7 @@ def test_mlflow_checkpoint_tag_is_absolute(monkeypatch, example_utils): monkeypatch, "--pyt_ckpt_path", "/models/Qwen3-0.6B", "--export_path", "exported_model" ) - assert Path(example_utils._mlflow_run_tags(args)["checkpoint_path"]).is_absolute() - - -class FakeMlflow: - """Stand-in for the mlflow module, so these tests need no server and no dependency.""" - - def __init__(self): - self.status = None - self.run_name = None - self.texts = {} - self.artifacts = {} - - def set_tracking_uri(self, uri): - self.tracking_uri = uri - - def set_experiment(self, name): - self.experiment = name - - def start_run(self, run_name=None): - self.run_name = run_name - return SimpleNamespace(info=SimpleNamespace(experiment_id="7", run_id="deadbeef")) - - def log_params(self, params): - pass - - def set_tags(self, tags): - pass - - def log_text(self, text, artifact_file): - self.texts[artifact_file] = text - - def log_artifact(self, local_path, artifact_path=None): - self.artifacts[Path(local_path).name] = artifact_path - - def log_metrics(self, metrics): - pass - - def end_run(self, status=None): - self.status = status - - -@pytest.fixture -def fake_mlflow(monkeypatch): - fake = FakeMlflow() - monkeypatch.setitem(sys.modules, "mlflow", fake) - return fake + assert Path(run_tags(args, example_utils.HF_PTQ)["checkpoint_path"]).is_absolute() def _tracked_run(monkeypatch, export_path, *extra): diff --git a/tests/examples/megatron_bridge/test_mlflow_utils.py b/tests/examples/megatron_bridge/test_mlflow_utils.py index a3e103b1c24..ce3902dca27 100644 --- a/tests/examples/megatron_bridge/test_mlflow_utils.py +++ b/tests/examples/megatron_bridge/test_mlflow_utils.py @@ -22,155 +22,84 @@ """ import argparse -import getpass import importlib.util import json import sys from pathlib import Path -from types import SimpleNamespace import pytest import yaml +from _test_utils.mlflow import clean_env # noqa: F401 +from modelopt.torch.utils import mlflow as mlflow_lib from modelopt.torch.utils.mlflow import masked_args _EXAMPLE_DIR = Path(__file__).resolve().parents[3] / "examples" / "megatron_bridge" _SCRIPT = _EXAMPLE_DIR / "quantize.py" -_SPEC = importlib.util.spec_from_file_location( - "megatron_bridge_mlflow_utils", _EXAMPLE_DIR / "mlflow_utils.py" -) -assert _SPEC is not None and _SPEC.loader is not None -mlflow_utils = importlib.util.module_from_spec(_SPEC) -_SPEC.loader.exec_module(mlflow_utils) -URI = "https://mlflow.example.com" -RECIPE = "general/ptq/nvfp4_default-kv_fp8_cast" - - -class FakeMlflow: - """Stand-in for the mlflow module, so these tests need no server and no dependency.""" - - def __init__(self): - self.status = None - self.texts = {} - self.artifacts = {} - def set_tracking_uri(self, uri): - self.tracking_uri = uri +def _load(name: str): + """Import one file from ``examples/megatron_bridge`` as a module. - def set_experiment(self, name): - self.experiment = name - - def start_run(self, run_name=None): - self.run_name = run_name - return SimpleNamespace(info=SimpleNamespace(experiment_id="7", run_id="deadbeef")) - - def log_params(self, params): - pass - - def set_tags(self, tags): - pass - - def log_text(self, text, artifact_file): - self.texts[artifact_file] = text - - def log_artifact(self, local_path, artifact_path=None): - self.artifacts[Path(local_path).name] = artifact_path - - def log_metrics(self, metrics): - pass - - def end_run(self, status=None): - self.status = status - - -@pytest.fixture(autouse=True) -def clean_env(monkeypatch): - """Pin what the tracking reads from the environment. - - ``resolve_tracking_uri`` consults $MLFLOW_TRACKING_URI, so a developer shell or runner - that exports it -- exactly the population this feature is built for -- would otherwise - flip the tracked/untracked branch under test. Tests that want the variable set it. + The scripts need Megatron to import, which this lane has: each declares its own ``Tool`` + beside the flags that Tool reads, so the two cannot drift. """ - monkeypatch.setattr(getpass, "getuser", lambda: "tester") - monkeypatch.delenv("MLFLOW_TRACKING_URI", raising=False) + # The scripts import each other by bare name, the way they do when run directly. + if str(_EXAMPLE_DIR) not in sys.path: + sys.path.insert(0, str(_EXAMPLE_DIR)) + spec = importlib.util.spec_from_file_location( + f"megatron_bridge_{name}", _EXAMPLE_DIR / f"{name}.py" + ) + assert spec is not None and spec.loader is not None + module = importlib.util.module_from_spec(spec) + sys.modules[spec.name] = module + spec.loader.exec_module(module) + return module -@pytest.fixture -def fake_mlflow(monkeypatch): - fake = FakeMlflow() - monkeypatch.setitem(sys.modules, "mlflow", fake) - return fake +mlflow_utils = _load("mlflow_utils") +QUANTIZE = _load("quantize").QUANTIZE + +URI = "https://mlflow.example.com" +RECIPE = "general/ptq/nvfp4_default-kv_fp8_cast" -def _parse(monkeypatch, *argv): - """The parser side of ``quantize.py``, limited to what the tracking reads off args.""" +def _parse(*argv): + """Parse *argv* the way ``quantize.py`` would, limited to what the tracking reads.""" parser = argparse.ArgumentParser() parser.add_argument("--hf_model_name_or_path", default="/models/Qwen3-0.6B") parser.add_argument("--export_megatron_path", default="/tmp/out") - parser.add_argument("--recipe", default=None) - parser.add_argument("--quant_cfg", default=None) + parser.add_argument("--recipe") + parser.add_argument("--quant_cfg") parser.add_argument("--tp_size", type=int, default=1) - mlflow_utils.add_mlflow_args(parser) + mlflow_lib.add_mlflow_args(parser, QUANTIZE) args = parser.parse_args(list(argv)) - mlflow_utils.resolve_mlflow_args(args, parser) + mlflow_lib.resolve_mlflow_args(args, parser, QUANTIZE) args.checkpoint_exported = False return args -def _tracked(monkeypatch, export_path, *extra): - return _parse(monkeypatch, "--export_megatron_path", str(export_path), "--mlflow", URI, *extra) +def _tracked(export_path, *extra): + return _parse("--export_megatron_path", str(export_path), "--mlflow", URI, *extra) # --- flags ---------------------------------------------------------------------------- -def test_tracking_is_off_by_default(monkeypatch): - args = _parse(monkeypatch) +def test_tracking_is_off_by_default(): + args = _parse() assert args.mlflow is None assert args.mlflow_experiment is None assert not args.mlflow_required -def test_the_flag_defaults_the_experiment_name(monkeypatch): - args = _parse(monkeypatch, "--recipe", RECIPE, "--mlflow", f"{URI}/") - - assert args.mlflow == URI # trailing slash stripped - assert args.mlflow_experiment == ( - "tester/megatron_bridge_quantize/Qwen3-0.6B-nvfp4_default-kv_fp8_cast" - ) - assert args.mlflow_run_name is None - - -def test_the_experiment_falls_back_to_quant_cfg_without_a_recipe(monkeypatch): - args = _parse(monkeypatch, "--quant_cfg", "nvfp4", "--mlflow", URI) - - assert args.mlflow_experiment == "tester/megatron_bridge_quantize/Qwen3-0.6B-nvfp4" - - -@pytest.mark.parametrize("sep", ["-", "_"]) -def test_multiword_flags_accept_both_spellings(monkeypatch, sep): - args = _parse( - monkeypatch, - "--mlflow", - URI, - f"--mlflow{sep}experiment", - "team/sweep", - f"--mlflow{sep}run{sep}name", - "calib-512", - ) - - assert args.mlflow_experiment == "team/sweep" - assert args.mlflow_run_name == "calib-512" - - def test_the_environment_alone_enables_tracking(monkeypatch): """MLFLOW_TRACKING_URI is MLflow's own variable, so exporting it opts in on its own.""" monkeypatch.setenv("MLFLOW_TRACKING_URI", f"{URI}/") - args = _parse(monkeypatch, "--quant_cfg", "nvfp4") + args = _parse("--quant_cfg", "nvfp4") assert args.mlflow == URI assert not args.mlflow_required # ... but it was not an explicit request @@ -182,10 +111,10 @@ def test_a_bad_uri_is_fatal_only_when_it_was_asked_for(monkeypatch): monkeypatch.setenv("MLFLOW_TRACKING_URI", "file:///local/mlruns") with pytest.warns(UserWarning, match="continuing untracked"): - assert _parse(monkeypatch).mlflow is None + assert _parse().mlflow is None with pytest.raises(SystemExit): - _parse(monkeypatch, "--mlflow", "file:///local/mlruns") + _parse("--mlflow", "file:///local/mlruns") def test_the_printed_arguments_mask_tracking_credentials(monkeypatch): @@ -194,7 +123,7 @@ def test_the_printed_arguments_mask_tracking_credentials(monkeypatch): # Fake credentials: TruffleHog flags any scheme://user:pass@host, and this test exists # precisely to prove they are masked. creds = "https://svc:s3cret@mlflow.example.com" # trufflehog:ignore - args = _parse(monkeypatch, "--mlflow", creds, "--quant_cfg", "nvfp4") + args = _parse("--mlflow", creds, "--quant_cfg", "nvfp4") printed = masked_args(args) @@ -209,9 +138,9 @@ def test_the_printed_arguments_mask_tracking_credentials(monkeypatch): def test_params_track_every_cli_argument(monkeypatch): """Params are derived from the parsed args, so a new flag needs no bookkeeping here.""" monkeypatch.setattr(mlflow_utils.dist, "size", lambda: 8) - args = _parse(monkeypatch, "--quant_cfg", "nvfp4", "--mlflow", URI, "--tp_size", "2") + args = _parse("--quant_cfg", "nvfp4", "--mlflow", URI, "--tp_size", "2") - params, _ = mlflow_utils._run_inputs(args) + params = mlflow_lib.describe_run(args, QUANTIZE, 8)["params"] assert params["hf_model_name_or_path"] == "/models/Qwen3-0.6B" assert params["tp_size"] == 2 @@ -222,13 +151,13 @@ def test_params_track_every_cli_argument(monkeypatch): params ) args.some_future_flag = "future" - assert mlflow_utils._run_inputs(args)[0]["some_future_flag"] == "future" + assert mlflow_lib.describe_run(args, QUANTIZE)["params"]["some_future_flag"] == "future" def test_run_inputs_carry_the_resolved_recipe(monkeypatch): - args = _parse(monkeypatch, "--recipe", RECIPE, "--mlflow", URI) + args = _parse("--recipe", RECIPE, "--mlflow", URI) - _, texts = mlflow_utils._run_inputs(args) + texts = mlflow_lib.describe_run(args, QUANTIZE)["texts"] # $imports are expanded, so the artifact stands alone. recipe = yaml.safe_load(texts["recipe/resolved_recipe.yaml"]) @@ -237,35 +166,22 @@ def test_run_inputs_carry_the_resolved_recipe(monkeypatch): def test_run_inputs_omit_the_recipe_when_unused(monkeypatch): - args = _parse(monkeypatch, "--quant_cfg", "nvfp4", "--mlflow", URI) - - assert mlflow_utils._run_inputs(args)[1] == {} + args = _parse("--quant_cfg", "nvfp4", "--mlflow", URI) - -def test_run_tags_identify_the_produced_checkpoint(monkeypatch, tmp_path): - """checkpoint_path must name what the run *writes*: the export and any QAD run are - pointed at the Megatron checkpoint, so tagging the input would never join the two.""" - export = tmp_path / "Qwen3-0.6B-nvfp4-megatron" - args = _tracked(monkeypatch, export) - - assert mlflow_utils._run_tags(args) == { - "model": "Qwen3-0.6B", - "checkpoint_path": str(export), - "source_checkpoint_path": "/models/Qwen3-0.6B", - } + assert mlflow_lib.describe_run(args, QUANTIZE)["texts"] == {} def test_the_checkpoint_tag_is_absolute(monkeypatch): """A relative --export_megatron_path is useless as a join key.""" - args = _tracked(monkeypatch, "megatron_ckpt") + args = _tracked("megatron_ckpt") - assert Path(mlflow_utils._run_tags(args)["checkpoint_path"]).is_absolute() + assert Path(mlflow_lib.run_tags(args, QUANTIZE)["checkpoint_path"]).is_absolute() def test_run_outputs_name_the_summary(monkeypatch): - args = _tracked(monkeypatch, "/tmp/megatron_ckpt") + args = _tracked("/tmp/megatron_ckpt") - files = mlflow_utils._run_outputs(args) + files = QUANTIZE.outputs(args) assert files["summary/quant_summary.txt"] == Path("/tmp/megatron_ckpt/.quant_summary.txt") @@ -276,11 +192,12 @@ def test_run_outputs_name_the_summary(monkeypatch): def test_non_master_ranks_do_not_open_a_run(monkeypatch, tmp_path): """Under torchrun only the master rank uploads, so the others must not touch the server.""" monkeypatch.setattr(mlflow_utils.dist, "is_master", lambda: False) - args = _tracked(monkeypatch, tmp_path) + args = _tracked(tmp_path) calls = [] - monkeypatch.setattr(mlflow_utils, "_run_inputs", lambda a: calls.append(a) or ({}, {})) + # Patched on the library, which is where tracked_run resolves it. + monkeypatch.setattr(mlflow_lib, "describe_run", lambda a, t, w=1: calls.append(a) or {}) - with mlflow_utils.mlflow_run(args): + with mlflow_utils.mlflow_run(args, QUANTIZE): args.checkpoint_exported = True assert calls == [] @@ -290,11 +207,12 @@ def test_non_master_ranks_do_not_open_a_run(monkeypatch, tmp_path): def test_untracked_runs_do_not_gather_inputs(monkeypatch, tmp_path): """Without --mlflow the recipe must not be re-read: it is parsed again in get_quant_config, and the extra load prints a second '[load_recipe] loading:' line.""" - args = _parse(monkeypatch, "--recipe", RECIPE, "--export_megatron_path", str(tmp_path)) + args = _parse("--recipe", RECIPE, "--export_megatron_path", str(tmp_path)) calls = [] - monkeypatch.setattr(mlflow_utils, "_run_inputs", lambda a: calls.append(a) or ({}, {})) + # Patched on the library, which is where tracked_run resolves it. + monkeypatch.setattr(mlflow_lib, "describe_run", lambda a, t, w=1: calls.append(a) or {}) - with mlflow_utils.mlflow_run(args): + with mlflow_utils.mlflow_run(args, QUANTIZE): pass assert calls == [] @@ -304,9 +222,9 @@ def test_experiment_json_lands_in_the_checkpoint_and_on_the_server( monkeypatch, fake_mlflow, tmp_path ): """The tags point run -> checkpoint; this file points checkpoint -> run.""" - args = _tracked(monkeypatch, tmp_path, "--quant_cfg", "nvfp4") + args = _tracked(tmp_path, "--quant_cfg", "nvfp4") - with mlflow_utils.mlflow_run(args): + with mlflow_utils.mlflow_run(args, QUANTIZE): args.checkpoint_exported = True # stand in for bridge.save_megatron_model written = json.loads((tmp_path / ".experiment.json").read_text()) @@ -321,9 +239,9 @@ def test_a_failed_save_writes_no_pointer_but_still_records_the_run( monkeypatch, fake_mlflow, tmp_path ): """The checkpoint was never written, so nothing on disk may claim this run produced it.""" - args = _tracked(monkeypatch, tmp_path, "--quant_cfg", "nvfp4") + args = _tracked(tmp_path, "--quant_cfg", "nvfp4") - with pytest.raises(RuntimeError), mlflow_utils.mlflow_run(args): + with pytest.raises(RuntimeError), mlflow_utils.mlflow_run(args, QUANTIZE): raise RuntimeError("calibration blew up") assert not (tmp_path / ".experiment.json").exists() @@ -336,9 +254,9 @@ def test_an_untracked_export_drops_an_inherited_pointer(monkeypatch, tmp_path): which would name a run that did not produce these weights.""" inherited = tmp_path / ".experiment.json" inherited.write_text('{"run_id": "stale"}') - args = _parse(monkeypatch, "--export_megatron_path", str(tmp_path)) + args = _parse("--export_megatron_path", str(tmp_path)) - with mlflow_utils.mlflow_run(args): + with mlflow_utils.mlflow_run(args, QUANTIZE): args.checkpoint_exported = True assert not inherited.exists() @@ -358,9 +276,9 @@ def explode(name): fake_mlflow.set_experiment = explode inherited = tmp_path / ".experiment.json" inherited.write_text('{"run_id": "from-an-earlier-run"}') - args = _parse(monkeypatch, "--export_megatron_path", str(tmp_path), "--quant_cfg", "nvfp4") + args = _parse("--export_megatron_path", str(tmp_path), "--quant_cfg", "nvfp4") - with mlflow_utils.mlflow_run(args): + with mlflow_utils.mlflow_run(args, QUANTIZE): args.checkpoint_exported = True assert args.mlflow_required is False @@ -371,9 +289,9 @@ def test_a_failed_untracked_run_leaves_the_directory_alone(monkeypatch, tmp_path """Nothing was exported, so whatever checkpoint is already there keeps its pointer.""" inherited = tmp_path / ".experiment.json" inherited.write_text('{"run_id": "stale"}') - args = _parse(monkeypatch, "--export_megatron_path", str(tmp_path)) + args = _parse("--export_megatron_path", str(tmp_path)) - with pytest.raises(RuntimeError), mlflow_utils.mlflow_run(args): + with pytest.raises(RuntimeError), mlflow_utils.mlflow_run(args, QUANTIZE): raise RuntimeError("calibration blew up") assert inherited.exists() @@ -382,23 +300,28 @@ def test_a_failed_untracked_run_leaves_the_directory_alone(monkeypatch, tmp_path # --- the seam with quantize.py -------------------------------------------------------- -def test_quantize_script_wires_the_tracking(): - """``quantize.py`` needs Megatron to import, so the wiring is checked as text: a renamed - flag or a dropped call would otherwise only surface in the Megatron example lane.""" - source = _SCRIPT.read_text() +# The script's seam, as text: it needs Megatron to import, so a renamed flag or a dropped +# call would otherwise only surface in this lane. The other four join in [2/2]. +_WIRING = { + "QUANTIZE": ( + _SCRIPT, + "with mlflow_run(args, QUANTIZE):", + ("--hf_model_name_or_path", "--export_megatron_path", "--recipe", "--quant_cfg"), + ), +} + + +@pytest.mark.parametrize("tool", list(_WIRING)) +def test_every_script_wires_the_tracking(tool): + script, opens_the_run, flags = _WIRING[tool] + source = script.read_text() - imported = next(line for line in source.splitlines() if line.startswith("from mlflow_utils ")) - for name in ("add_mlflow_args", "mlflow_run", "resolve_mlflow_args"): - assert name in imported - assert "from modelopt.torch.utils.mlflow import masked_args" in source - assert "add_mlflow_args(parser)" in source - assert "resolve_mlflow_args(args, parser)" in source - assert "with mlflow_run(args):" in source + # Registered before parsing, or the flags would not exist on the command line. + assert source.index(f"add_mlflow_args(parser, {tool}") < source.index("parser.parse_args()") + assert f"resolve_mlflow_args(args, parser, {tool})" in source + assert opens_the_run in source # The namespace reaches print_args masked, so a user:token@ URI stays out of the job log. assert "print_args(masked_args(args))" in source - # The provenance pointer is gated on the save having happened. - assert "args.checkpoint_exported = False" in source - assert "args.checkpoint_exported = True" in source - # Every args attribute the tracking reads is a flag quantize.py registers. - for flag in ("--hf_model_name_or_path", "--export_megatron_path", "--recipe", "--quant_cfg"): + # Every args attribute the tracking reads is a flag the script registers. + for flag in flags: assert f'"{flag}"' in source diff --git a/tests/examples/vllm_serve/test_vllm_mlflow_utils.py b/tests/examples/vllm_serve/test_vllm_mlflow_utils.py index 04b97c31580..9097018b508 100644 --- a/tests/examples/vllm_serve/test_vllm_mlflow_utils.py +++ b/tests/examples/vllm_serve/test_vllm_mlflow_utils.py @@ -65,51 +65,9 @@ } -class FakeMlflow: - """Stand-in for the mlflow module, so these tests need no server and no dependency.""" - - def __init__(self): - self.tracking_uri = None - self.experiment = None - self.run_name = None - self.status = None - self.params = {} - self.tags = {} - self.texts = {} - self.metrics = {} - self.artifacts = {} - - def set_tracking_uri(self, uri): - self.tracking_uri = uri - - def set_experiment(self, name): - self.experiment = name - - def start_run(self, run_name=None): - self.run_name = run_name - return SimpleNamespace(info=SimpleNamespace(experiment_id="7", run_id="deadbeef")) - - def log_params(self, params): - self.params.update(params) - - def set_tags(self, tags): - self.tags.update(tags) - - def log_text(self, text, artifact_file): - self.texts[artifact_file] = text - - def log_artifact(self, local_path, artifact_path=None): - self.artifacts[Path(local_path).name] = (artifact_path, Path(local_path).read_text()) - - def log_metrics(self, metrics): - self.metrics.update(metrics) - - def end_run(self, status=None): - self.status = status - - @pytest.fixture(autouse=True) def clean_env(monkeypatch): + """Local rather than shared: this suite also clears the launcher-to-worker variables.""" monkeypatch.setattr(getpass, "getuser", lambda: "tester") for name in _TRACKED_ENV: monkeypatch.delenv(name, raising=False) @@ -121,13 +79,6 @@ def mlflow_utils(monkeypatch): return importlib.import_module("vllm_mlflow_utils") -@pytest.fixture -def fake_mlflow(monkeypatch): - fake = FakeMlflow() - monkeypatch.setitem(sys.modules, "mlflow", fake) - return fake - - def _resolve(mlflow_utils, monkeypatch, model="/ckpts/Qwen3-0.6B", **flags): """Run the launcher's side of the handover, the way vllm_serve_fakequant.py does.""" parser = argparse.ArgumentParser() diff --git a/tests/unit/torch/utils/test_mlflow.py b/tests/unit/torch/utils/test_mlflow.py index da71792e242..9cf601f3ffa 100644 --- a/tests/unit/torch/utils/test_mlflow.py +++ b/tests/unit/torch/utils/test_mlflow.py @@ -14,26 +14,26 @@ # limitations under the License. import argparse +import dataclasses import getpass import io import json import logging import sys -from pathlib import Path -from types import SimpleNamespace import pytest import yaml +from _test_utils.mlflow import FakeMlflow, clean_env # noqa: F401 import modelopt from modelopt.torch.utils.logging import TeeStream from modelopt.torch.utils.mlflow import ( EXPERIMENT_JSON, MlflowRunLogger, + Tool, _git_sha, _redact_argv, add_mlflow_args, - checkpoint_run_tags, command_text, default_experiment_name, drop_experiment_json, @@ -41,7 +41,8 @@ masked_args, resolve_mlflow_args, resolved_recipe_texts, - track_run, + run_tags, + tracked_run, validate_tracking_uri, ) @@ -53,66 +54,6 @@ SHORT_CREDS_URI = "https://u:tok@host" # trufflehog:ignore -class FakeMlflow: - """Stand-in for the mlflow module, so these tests need no server and no dependency.""" - - def __init__(self): - self.tracking_uri = None - self.experiment = None - self.run_name = None - self.status = None - self.params = {} - self.tags = {} - self.texts = {} - self.metrics = {} - self.artifacts = [] - self.artifact_text = {} - # What the server says the run is called, which need not be what was requested. - self.server_run_name = None - - def set_tracking_uri(self, uri): - self.tracking_uri = uri - - def set_experiment(self, name): - self.experiment = name - - def start_run(self, run_name=None): - self.run_name = run_name - return SimpleNamespace( - info=SimpleNamespace( - experiment_id="7", - run_id="deadbeef", - run_name=self.server_run_name or run_name, - ) - ) - - def log_params(self, params): - self.params.update(params) - - def set_tags(self, tags): - self.tags.update(tags) - - def log_text(self, text, artifact_file): - self.texts[artifact_file] = text - - def log_artifact(self, local_path, artifact_path=None): - self.artifacts.append((Path(local_path).name, artifact_path)) - self.artifact_text[Path(local_path).name] = Path(local_path).read_text() - - def log_metrics(self, metrics): - self.metrics.update(metrics) - - def end_run(self, status=None): - self.status = status - - -@pytest.fixture -def fake_mlflow(monkeypatch): - fake = FakeMlflow() - monkeypatch.setitem(sys.modules, "mlflow", fake) - return fake - - def _unreachable(fake): """Make MLflow's own first request fail, the way a dead server does.""" @@ -123,18 +64,6 @@ def explode(*args, **kwargs): return fake -@pytest.fixture(autouse=True) -def clean_env(monkeypatch): - """Pin what the tracking reads from the environment. - - ``resolve_tracking_uri`` consults $MLFLOW_TRACKING_URI, so a developer shell or runner - that exports it -- exactly the population this feature is built for -- would otherwise - flip the tracked/untracked branch under test. Tests that want the variable set it. - """ - monkeypatch.setattr(getpass, "getuser", lambda: "tester") - monkeypatch.delenv("MLFLOW_TRACKING_URI", raising=False) - - def _logger(**kwargs): kwargs.setdefault("experiment_name", "tester/hf_ptq/model-nvfp4") return MlflowRunLogger(URI, **kwargs) @@ -304,9 +233,9 @@ def test_logger_logs_inputs_and_outputs(fake_mlflow, tmp_path, monkeypatch): assert fake_mlflow.tags["modelopt_version"] == modelopt.__version__ # The log keeps its name; the summary is renamed out of its dotfile form. - assert ("hf_ptq.log", "logs") in fake_mlflow.artifacts - assert ("quant_summary.txt", "summary") in fake_mlflow.artifacts - assert not any(name == "moe.html" for name, _ in fake_mlflow.artifacts) + assert fake_mlflow.artifacts["hf_ptq.log"][0] == "logs" + assert fake_mlflow.artifacts["quant_summary.txt"][0] == "summary" + assert "moe.html" not in fake_mlflow.artifacts assert "total_time_s" in fake_mlflow.metrics assert fake_mlflow.status == "FINISHED" @@ -481,7 +410,7 @@ def test_logger_restores_streams_and_reports_failure(fake_mlflow, monkeypatch): assert sys.stdout is stdout and sys.stderr is stderr assert fake_mlflow.status == "FAILED" - assert ("hf_ptq.log", "logs") in fake_mlflow.artifacts + assert fake_mlflow.artifacts["hf_ptq.log"][0] == "logs" def test_logger_never_raises_when_the_server_dies_mid_run(fake_mlflow, capsys): @@ -664,6 +593,19 @@ def test_track_closes_the_run_with_the_right_status(fake_mlflow, monkeypatch): assert fake_mlflow.params == {"qformat": "nvfp4"} +@pytest.mark.parametrize( + ("code", "status"), [(0, "FINISHED"), (None, "FINISHED"), (1, "FAILED"), (2, "FAILED")] +) +def test_a_block_that_exits_cleanly_is_a_finished_run(fake_mlflow, code, status): + """A script that ends by calling sys.exit() rather than returning -- Megatron-Bridge does, + from inside its training loop -- finished if it exited cleanly. Before this shared exit + path, every SystemExit reached the bare ``finally`` and was recorded as FAILED.""" + with pytest.raises(SystemExit), _logger().track(): + raise SystemExit(code) + + assert fake_mlflow.status == status + + def test_track_marks_a_raising_block_failed(fake_mlflow, monkeypatch, tmp_path): monkeypatch.setattr(sys, "argv", ["hf_ptq.py"]) stdout = sys.stdout @@ -680,8 +622,8 @@ def test_track_marks_a_raising_block_failed(fake_mlflow, monkeypatch, tmp_path): assert fake_mlflow.status == "FAILED" assert sys.stdout is stdout # The outputs named upfront are still uploaded, and the traceback is in the log. - assert ("quant_summary.txt", "summary") in fake_mlflow.artifacts - assert "RuntimeError: calibration exploded" in fake_mlflow.artifact_text["hf_ptq.log"] + assert fake_mlflow.artifacts["quant_summary.txt"][0] == "summary" + assert "RuntimeError: calibration exploded" in fake_mlflow.artifacts["hf_ptq.log"][1] def test_failed_run_uploads_the_traceback(fake_mlflow, monkeypatch): @@ -697,7 +639,7 @@ def test_failed_run_uploads_the_traceback(fake_mlflow, monkeypatch): except RuntimeError: logger.finish("FAILED") - log = fake_mlflow.artifact_text["hf_ptq.log"] + log = fake_mlflow.artifacts["hf_ptq.log"][1] assert "calibrating" in log assert "Traceback (most recent call last)" in log assert "RuntimeError: calibration exploded" in log @@ -710,7 +652,7 @@ def test_successful_run_uploads_no_traceback(fake_mlflow, monkeypatch): logger.start() logger.finish("FINISHED") - assert "Traceback" not in fake_mlflow.artifact_text["hf_ptq.log"] + assert "Traceback" not in fake_mlflow.artifacts["hf_ptq.log"][1] def test_only_files_this_run_produced_are_uploaded(fake_mlflow, tmp_path, monkeypatch): @@ -727,7 +669,7 @@ def test_only_files_this_run_produced_are_uploaded(fake_mlflow, tmp_path, monkey fresh.write_text("written by this run") # produced during the run logger.finish("FAILED", files=outputs) - uploaded = [name for name, _ in fake_mlflow.artifacts] + uploaded = list(fake_mlflow.artifacts) assert "moe.html" in uploaded assert "quant_summary.txt" not in uploaded @@ -746,7 +688,7 @@ def test_stale_check_survives_unnormalized_string_paths(fake_mlflow, tmp_path, m logger.start(files=outputs) logger.finish("FAILED", files=outputs) - assert "quant_summary.txt" not in [name for name, _ in fake_mlflow.artifacts] + assert "quant_summary.txt" not in fake_mlflow.artifacts def test_optional_tracking_warns_and_continues_when_the_server_is_unreachable(monkeypatch, capsys): @@ -774,26 +716,29 @@ def test_required_tracking_still_raises(monkeypatch): # --- the CLI surface the example scripts share ----------------------------------------- +_TOOL = Tool( + name="hf_ptq", + tracks="Track this run on an MLflow server (e.g. https:///).", + variant_help="recipe name", + variant=lambda args: "nvfp4", + model=lambda args: "/models/Qwen3-0.6B", + checkpoint=lambda args: "/exports/out", +) + + def _parser(): parser = argparse.ArgumentParser() - add_mlflow_args(parser, "hf_ptq", variant_help="recipe name") + add_mlflow_args(parser, _TOOL) return parser def _resolved(argv, parser=None): parser = parser or _parser() args = parser.parse_args(argv) - resolve_mlflow_args(args, parser, tool="hf_ptq", model="/models/Qwen3-0.6B", variant="nvfp4") + resolve_mlflow_args(args, parser, _TOOL) return args -def test_flags_are_off_by_default(): - args = _resolved([]) - - assert (args.mlflow, args.mlflow_experiment, args.mlflow_run_name) == (None, None, None) - assert args.mlflow_required is False - - def test_the_flag_names_the_experiment_and_normalizes_the_uri(): args = _resolved(["--mlflow", f"{URI}/"]) @@ -808,25 +753,6 @@ def test_an_explicit_experiment_is_left_alone(): assert args.mlflow_experiment == "team/sweep" -@pytest.mark.parametrize("sep", ["-", "_"]) -def test_multiword_flags_accept_both_spellings(sep): - """vLLM's FlexibleArgumentParser rewrites --foo_bar to --foo-bar before matching, so a - flag registered only under the underscored spelling is unreachable from its CLI.""" - args = _resolved([f"--mlflow{sep}experiment", "team/sweep", f"--mlflow{sep}run{sep}name", "r"]) - - assert (args.mlflow_experiment, args.mlflow_run_name) == ("team/sweep", "r") - - -def test_the_environment_alone_enables_tracking(monkeypatch): - monkeypatch.setenv("MLFLOW_TRACKING_URI", f"{URI}/") - - args = _resolved([]) - - assert args.mlflow == URI - assert args.mlflow_required is False # ... but it was not an explicit request - assert args.mlflow_experiment == "tester/hf_ptq/Qwen3-0.6B-nvfp4" - - def test_the_flag_overrides_the_environment(monkeypatch): monkeypatch.setenv("MLFLOW_TRACKING_URI", "https://other.example.com") @@ -974,23 +900,17 @@ def test_dropping_the_pointer_is_idempotent(tmp_path): # --- the pieces every checkpoint-producing run shares ----------------------------------- -def test_checkpoint_run_tags_name_what_the_run_writes(tmp_path): - """An export or an evaluation is pointed at the checkpoint the run produced, so tagging - the input instead would never join the two.""" - tags = checkpoint_run_tags("/models/Qwen3-0.6B", tmp_path / "out") - - assert tags == { - "model": "Qwen3-0.6B", - "checkpoint_path": str(tmp_path / "out"), - "source_checkpoint_path": "/models/Qwen3-0.6B", - } - - -def test_checkpoint_run_tags_resolve_a_relative_path(monkeypatch, tmp_path): - """Export paths commonly default to a relative one, which is useless as a join key.""" - monkeypatch.chdir(tmp_path) - - assert Path(checkpoint_run_tags("m", "exported_model")["checkpoint_path"]).is_absolute() +def _tool(model="/models/Qwen3-0.6B", checkpoint="/exports/out", source=None): + """A Tool that names just what the tags read.""" + return Tool( + name="demo", + tracks="Track it.", + variant_help="v", + variant=lambda args: "v1", + model=lambda args: model, + checkpoint=lambda args: checkpoint, + source=(lambda args: source) if source is not None else None, + ) def test_resolved_recipe_texts_carry_a_self_contained_recipe(): @@ -1016,8 +936,14 @@ def test_masked_args_masks_the_uri_and_nothing_else(): assert args.mlflow == CREDS_URI # the caller's namespace is untouched -def _tracked_logger(**kwargs): - return _logger(**kwargs) +def _run_args(tracked=True): + """The namespace tracked_run reads its destination from.""" + return argparse.Namespace( + mlflow=URI if tracked else None, + mlflow_experiment="tester/hf_ptq/model-nvfp4", + mlflow_run_name=None, + mlflow_required=True, + ) @pytest.mark.parametrize( @@ -1029,25 +955,27 @@ def test_untracked_run_clears_only_what_it_replaced(tmp_path, exported, is_main, """The pointer is dropped exactly when this rank wrote a fresh checkpoint over it.""" stale = tmp_path / EXPERIMENT_JSON stale.write_text('{"run_id": "an-earlier-run"}') - logger = MlflowRunLogger(URI, "e", enabled=False) - with track_run(logger, tmp_path, is_main=is_main, exported=lambda: exported): + with tracked_run( + _run_args(tracked=False), + _tool(checkpoint=tmp_path), + is_main=is_main, + exported=lambda: exported, + ): pass assert stale.exists() is survives -def test_untracked_run_does_not_gather_what_it_will_not_upload(tmp_path): +def test_untracked_run_does_not_gather_what_it_will_not_upload(monkeypatch, tmp_path): """Gathering can re-read a recipe, which an untracked run must not pay for.""" calls = [] - logger = MlflowRunLogger(URI, "e", enabled=False) - - with track_run( - logger, - tmp_path, - is_main=True, - exported=lambda: False, - describe=lambda: calls.append(1) or {}, + monkeypatch.setattr( + sys.modules[tracked_run.__module__], "describe_run", lambda a, t, w=1: calls.append(1) or {} + ) + + with tracked_run( + _run_args(tracked=False), _tool(checkpoint=tmp_path), is_main=True, exported=lambda: False ): pass @@ -1055,17 +983,14 @@ def test_untracked_run_does_not_gather_what_it_will_not_upload(tmp_path): def test_tracked_run_records_the_pointer_and_closes(fake_mlflow, tmp_path): - with track_run( - _logger(), - tmp_path, - is_main=True, - exported=lambda: True, - describe=lambda: {"params": {"model": "m"}}, - ): + args = _run_args() + args.model_for_params = "m" + + with tracked_run(args, _tool(checkpoint=tmp_path), is_main=True, exported=lambda: True): pass assert json.loads((tmp_path / EXPERIMENT_JSON).read_text())["run_id"] == "deadbeef" - assert fake_mlflow.params["model"] == "m" + assert fake_mlflow.params["model_for_params"] == "m" assert fake_mlflow.status == "FINISHED" @@ -1074,19 +999,78 @@ def test_exported_is_read_on_the_way_out(fake_mlflow, tmp_path): record the state before the checkpoint existed.""" state = {"exported": False} - with track_run(_logger(), tmp_path, is_main=True, exported=lambda: state["exported"]): + with tracked_run( + _run_args(), _tool(checkpoint=tmp_path), is_main=True, exported=lambda: state["exported"] + ): state["exported"] = True assert (tmp_path / EXPERIMENT_JSON).exists() +@pytest.mark.parametrize("raises", ["exported", "metrics"]) +def test_a_callback_that_raises_does_not_cost_the_run_its_close(fake_mlflow, tmp_path, raises): + """Both report what the run did, so a run that failed early may never have stashed what + they read. Raising from the exit path would run *instead of* finish(), leaving the run + RUNNING on the server and stdout still pointed at the capture tee.""" + + def boom(*_): + raise AttributeError("'Namespace' object has no attribute 'prune_score'") + + tool = _tool(checkpoint=tmp_path) + if raises == "metrics": + tool = dataclasses.replace(tool, metrics=boom) + + with tracked_run( + _run_args(), tool, is_main=True, exported=boom if raises == "exported" else (lambda: True) + ): + pass + + assert fake_mlflow.status == "FINISHED" + assert not isinstance(sys.stdout, TeeStream) + + +def test_a_tool_that_settles_no_pointer_leaves_the_directory_alone(fake_mlflow, tmp_path): + """For a script that points each of several checkpoints at the run itself: tracked_run + must neither write the pointer nor clear one it did not replace.""" + stale = tmp_path / EXPERIMENT_JSON + stale.write_text('{"run_id": "an-earlier-run"}') + tool = dataclasses.replace(_tool(checkpoint=tmp_path), settles_pointer=False) + + with tracked_run(_run_args(), tool, is_main=True, exported=lambda: True): + pass + + assert json.loads(stale.read_text())["run_id"] == "an-earlier-run" + + def test_a_failed_tracked_run_leaves_no_pointer_but_is_still_recorded(fake_mlflow, tmp_path): with ( pytest.raises(RuntimeError), - track_run(_logger(), tmp_path, is_main=True, exported=lambda: False), + tracked_run(_run_args(), _tool(checkpoint=tmp_path), is_main=True, exported=lambda: False), ): raise RuntimeError("calibration blew up") assert not (tmp_path / EXPERIMENT_JSON).exists() assert json.loads(fake_mlflow.texts["experiment.json"])["run_id"] == "deadbeef" assert fake_mlflow.status == "FAILED" + + +@pytest.mark.parametrize("relative", [True, False], ids=["relative", "absolute"]) +def test_a_source_path_joins_whatever_the_caller_typed(monkeypatch, tmp_path, relative): + """The previous stage tagged its output resolved, so a relative source must match it.""" + produced = tmp_path / "ptq" + produced.mkdir() + monkeypatch.chdir(tmp_path) + source = "ptq" if relative else str(produced) + + upstream = run_tags(argparse.Namespace(), _tool(model="/models/m", checkpoint=produced)) + downstream = run_tags(argparse.Namespace(), _tool(checkpoint=tmp_path / "qad", source=source)) + + assert downstream["source_checkpoint_path"] == upstream["checkpoint_path"] + + +def test_a_hub_model_id_is_not_mistaken_for_a_path(): + """``org/name`` names no directory, so resolving it would invent one under the cwd.""" + tags = run_tags(argparse.Namespace(), _tool(model="Qwen/Qwen3-8B", checkpoint="/out")) + + assert tags["source_checkpoint_path"] == "Qwen/Qwen3-8B" + assert tags["model"] == "Qwen3-8B"