Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
127 changes: 28 additions & 99 deletions examples/hf_ptq/example_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__)

Expand Down Expand Up @@ -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://<your-mlflow-server>/), "
"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://<your-mlflow-server>/), "
"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",
}
8 changes: 4 additions & 4 deletions examples/hf_ptq/hf_ptq.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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,
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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].")
Expand Down
120 changes: 15 additions & 105 deletions examples/megatron_bridge/mlflow_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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://<your-mlflow-server>/), "
"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
Loading
Loading