Add optional MLflow tracking to the vLLM fake-quant server (#2120)

### What does this PR do?

Type of change: new feature

Wires `examples/vllm_serve/vllm_serve_fakequant.py` up to
`modelopt.torch.utils.mlflow` via `--mlflow <tracking-uri>`, the same
way #2023 did for `hf_ptq.py`, so a fake-quant serve records **what it
actually quantized** and an evaluation of that endpoint can be traced
back to a recipe. Without the flag, behavior is unchanged — every hook
is gated on it.

Three design points worth review:

1. **The run is recorded in the vLLM worker, not the launcher.**
`vllm_serve_fakequant.py` is the API-server frontend; the engine and its
workers are separate processes whose stdout it never sees, so a run
opened there would capture none of the calibration. The launcher instead
only settles the tracking configuration — validating the URI, naming the
experiment, recording the command the user actually typed — and
publishes it through the environment, which is how every other setting
in this example (`QUANT_CFG`, `RECIPE_PATH`, …) already reaches the
workers. Global rank 0 opens the run, so a TP-8 serve produces one run.

2. **The run covers load-through-warm-up, not the server's lifetime.**
It opens *before the weights load*, so an unreachable server or a
missing token fails in seconds rather than after a load and a full
calibration, and it closes `FINISHED` once the model is quantized and
warmed up. A run that stayed open for the serving lifetime would never
close cleanly on SIGTERM.

3. **`recipe/quant_cfg.yaml` is only written on the preset path.** With
`RECIPE_PATH`, `get_quant_config` returns the recipe's `quantize`
section unchanged and `resolved_recipe.yaml` already carries it. With
`QUANT_CFG`/`KV_QUANT_CFG` it is the *only* record of what ran: the
params carry the preset names, while the config reaching `mtq.quantize`
is those two deep-copied, merged, and — for an MLA model — extended at
runtime with `*kv_c_bmm_quantizer` / `*k_pe_bmm_quantizer` by inspecting
the loaded model.

Uploaded artifacts:

| Artifact | Contents |
| --- | --- |
| `command.txt` | The launcher's invocation, copy-pasteable, credentials
masked |
| `version.txt` | The ModelOpt version that ran |
| `recipe/resolved_recipe.yaml` | `RECIPE_PATH` with its `$import`s
expanded |
| `recipe/quant_cfg.yaml` | Merged `QUANT_CFG`/`KV_QUANT_CFG` + MLA
fixup (preset path only) |
| `logs/<script>.log` | The rank-0 worker's stdout/stderr, including a
crash traceback |
| `summary/quant_summary.txt` | The per-quantizer summary |

Plus the quantization *and* serving settings as searchable params, and
`user` / `hostname` / `modelopt_version` / `git_sha` / `vllm_version`
tags. The `checkpoint_path` tag matches the one `hf_ptq.py` sets, so a
checkpoint's PTQ run and every serve of it join up.

Two small library additions, both consumed by the new example module:

- `command_text(argv=None)` — records another process's invocation,
since a spawned worker's own `sys.argv` is vLLM plumbing rather than
anything a user typed.
- `MlflowRunLogger.log_text()` — uploads a value settled midway through
a run, so a crash during calibration still keeps the config that caused
it.

The example `Dockerfile` installs the `mlflow` extra; the client remains
optional and is imported only once tracking is enabled.

### Usage

```bash
RECIPE_PATH=<recipe.yaml> python vllm_serve_fakequant.py <model_path> -tp 8 \
  --host 0.0.0.0 --port 8000 \
  --mlflow https://<your-mlflow-server>/
```

```
[mlflow] tracking to https://<your-mlflow-server>, experiment $USER/vllm_serve_fakequant/<model>-<recipe>
(Worker_TP0) [mlflow] run: https://<your-mlflow-server>/#/experiments/19/runs/1c6679448f25...
```

`--mlflow-experiment` / `--mlflow-run-name` override the defaults.
`$MLFLOW_TRACKING_URI` enables tracking on its own and is best-effort;
an explicit `--mlflow` overrides it and fails loudly.

> This is the **quantization** tracking server. It is unrelated to any
server an evaluation harness exports its scores to — NeMo Evaluator
Launcher has its own `export.mlflow.tracking_uri`. The README calls this
out.

### Testing

**Unit — 87 passing**
(`tests/examples/vllm_serve/test_vllm_mlflow_utils.py`, 33 new;
`tests/unit/torch/utils/test_mlflow.py`, +5). `vllm_mlflow_utils`
deliberately imports no vLLM, so the whole launcher→worker handover is
covered without a GPU, a server, or the mlflow client.

**End to end on aws-cmh** (4× GB300, `simple_evals.gpqa_diamond`,
Nemotron-3.5-Lightning-30B-A3B-BF16 fake-quantized with
`general/ptq/nvfp4_mlp_only-kv_fp8_cast`): run `FINISHED` in 261.5 s,
opened by `Worker_TP0` only, all artifacts present and verified by
content — `command.txt` held the launcher's invocation rather than the
worker's spawn argv, and `resolved_recipe.yaml` was 6797 B against 1845
B of source. 104 quantizers enabled (92 NVFP4 dynamic block-16 expert
weight/input with calibrated amax, 12 FP8 KV bmm). The eval then ran to
completion against the served endpoint, 22/22 requests HTTP 200.

Two bugs the hardware run caught, both fixed here with regression tests:

- `--mlflow_run_name` was rejected. vLLM's
`FlexibleArgumentParser.parse_args` rewrites **every** `--foo_bar` to
`--foo-bar` before matching, so a flag registered only under the
underscored spelling is unreachable from its CLI. Both spellings are now
registered. A unit test on a plain `ArgumentParser` could not have
caught this.
- `recipe/quant_cfg.yaml` uploaded a Python `repr` blob under a `.yaml`
name: a recipe's `quantize` is a `QuantizeConfig`, `yaml.safe_dump`
raises `RepresenterError` on it, and the old JSON fallback stringified
the object. `_dump_yaml` now unwraps pydantic via
`model_dump(mode="json")` and raises otherwise, with the caller
downgrading that to a warning so a bad config cannot take down a serve.

**Known coverage gap:** the preset (`QUANT_CFG`/`KV_QUANT_CFG`) path —
the only one that now writes `recipe/quant_cfg.yaml` — is covered by
unit test but has not been exercised on hardware; the canary used
`RECIPE_PATH`. Likewise the case where `$MLFLOW_TRACKING_URI` is present
*inside* the deployment container and `--mlflow` overrides it is
unit-tested only: NeMo Evaluator Launcher forwards only declared env
vars, so the eval server's URI never entered the container in the
canary.

### Before your PR is "*Ready for review*"

- Is this change backward compatible?: ✅ — new optional flags only; no
`--mlflow` means no behavior change.
- If you copied code from any other sources or added a new PIP
dependency, did you follow guidance in `CONTRIBUTING.md`: ✅ — no new
dependency. Uses the existing optional `nvidia-modelopt[mlflow]` extra
(`mlflow-skinny`, Apache-2.0) added in #2023; the example `Dockerfile`
now installs it. No code copied from other sources.
- Did you write any new necessary tests?: ✅ — 38 new tests.
- Did you update
[Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?:
✅ — 0.47 Misc.
- Did you get Claude approval on this PR?: ❌ — `/claude review` not yet
run.

### Additional Information

Follows #2023, which added `MlflowRunLogger` and the `hf_ptq.py`
integration.

Note for anyone tracking from an OCI cluster:
`mlflow-modelopt.nvidia.com` is unreachable from oci-nrt and oci-hsg.
TCP 443 completes and the connection is then reset on the first
application byte, regardless of SNI or protocol, one RTT away — the PDX
PaaS ingress appears to apply a source-IP policy, and the OCI clusters
egress from Oracle-owned addresses (`155.248.190.0`, `168.110.199.1`)
rather than NVIDIA's. gcp-nrt, aws-cmh and cw-dfw all reach it. This is
an infrastructure matter, not a property of this change, but it
determines where the feature is usable today.

🤖 Generated with [Claude Code](https://claude.com/claude-code)


<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->
## Summary by CodeRabbit

* **New Features**
* Added optional MLflow tracking for vLLM fake-quantization serving
runs.
* Records serving, quantization, worker, and invocation metadata,
including configuration and summary artifacts.
* Supports tracking URI, credentials, environment, and command-line
configuration.
  * Added command and text artifact logging for active MLflow runs.
* **Documentation**
* Documented setup, configuration, recorded artifacts, lifecycle, and
fallback behavior.
  * Updated the example container to include MLflow support.
* **Tests**
* Added comprehensive coverage for tracking configuration, logging,
failures, and disabled tracking.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->

---------

Signed-off-by: Chenjie Luo <chenjiel@nvidia.com>
This commit is contained in:
Chenjie Luo
2026-08-13 00:00:21 +00:00
committed by GitHub
parent f3e88df9ea
commit b96841db3e
9 changed files with 1117 additions and 20 deletions
+1
View File
@@ -19,6 +19,7 @@ Changelog
- Add ``modelopt.torch.utils.mlflow.MlflowRunLogger`` for recording a script run on an MLflow tracking server: the invocation, the ModelOpt version, the run log (captured by teeing ``stdout``/``stderr``) and any caller-supplied artifacts, with configuration as searchable params. ``mlflow`` is an optional dependency, imported only when tracking is enabled.
- Add ``--mlflow <tracking-uri>`` to ``examples/hf_ptq/hf_ptq.py`` (MLflow's own ``MLFLOW_TRACKING_URI`` is honoured too). A tracked run records the invocation, the resolved recipe (``$import``\ s expanded), the run log and the quantization summaries, with every command-line argument as a searchable param; failed runs are recorded with their traceback. The experiment defaults to ``$USER/hf_ptq/<checkpoint basename>-<recipe name or --qformat>`` and can be overridden with ``--mlflow_experiment`` / ``--mlflow_run_name``.
- Add ``--mlflow <tracking-uri>`` to ``examples/vllm_serve/vllm_serve_fakequant.py`` (MLflow's own ``MLFLOW_TRACKING_URI`` is honoured too), so a fake-quant serve records what it quantized and an evaluation of that endpoint can be traced back to a recipe. A tracked run uploads the launcher command, the resolved ``RECIPE_PATH`` (or the merged ``QUANT_CFG``/``KV_QUANT_CFG`` when presets are used), the worker log and the quantizer summary; the experiment defaults to ``$USER/vllm_serve_fakequant/<model basename>-<recipe name or quantization config>`` and can be overridden with ``--mlflow-experiment`` / ``--mlflow-run-name``.
**Backward Breaking Changes**
+3 -2
View File
@@ -18,9 +18,10 @@ COPY . Model-Optimizer
# Remove .git directory to reduce image size
RUN rm -rf Model-Optimizer/.git
# Install modelopt from local source with all dependencies
# Install modelopt from local source with all dependencies. `mlflow` is the optional
# tracking client used by --mlflow; it is not part of `all`.
RUN cd Model-Optimizer && \
pip install -e ".[all,dev-test]"
pip install -e ".[all,dev-test,mlflow]"
# Llama4 requires this
RUN pip install flash-attn==2.7.4.post1 --no-build-isolation
+61
View File
@@ -62,6 +62,67 @@ Step 4 (Optional): using lm_eval to run evaluation
lm_eval --model local-completions --tasks gsm8k --model_args model=<model_name>,base_url=http://127.0.0.1:8000/v1/completions,num_concurrent=1,max_retries=3,tokenized_requests=False,batch_size=128,tokenizer_backend=None
```
## Tracking a serve with MLflow
Pass `--mlflow <tracking-uri>`, or set MLflow's own `MLFLOW_TRACKING_URI`, to record what
this server actually quantized, so the numbers an evaluation produces can be traced back to
a recipe:
```bash
RECIPE_PATH=<PATH_TO_RECIPE> python vllm_serve_fakequant.py <model_path> -tp 8 \
--host 0.0.0.0 --port 8000 \
--mlflow https://<your-mlflow-server>/
```
This is the *quantization* tracking server. It is unrelated to any tracking server an
evaluation harness exports its scores to — NeMo Evaluator Launcher, for instance, has its
own `export.mlflow.tracking_uri`. Keep the two separate.
Quantization runs in the vLLM **worker**, not in `vllm_serve_fakequant.py`, so that is where
the run is recorded: the launcher validates the URI and hands the settings to the workers
through the environment, and global rank 0 opens the run. It opens *before the weights
load*, so a bad URI or a missing token fails within seconds rather than after a load and a
full calibration, and it closes `FINISHED` once the model is quantized and warmed up —
serving itself is not tracked.
<details>
<summary>Uploaded artifacts</summary>
| Artifact | Contents |
| --- | --- |
| `command.txt` | The launcher's full invocation, copy-pasteable, with credentials masked |
| `version.txt` | The ModelOpt version that ran |
| `recipe/resolved_recipe.yaml` | `RECIPE_PATH` with its `$import`s expanded, so it stands alone |
| `recipe/quant_cfg.yaml` | `QUANT_CFG` and `KV_QUANT_CFG` merged, plus any MLA fixup — only when no recipe is used, since a recipe's config is already in `resolved_recipe.yaml` |
| `logs/<script>.log` | The rank-0 worker's Python stdout/stderr, including the traceback if it crashed |
| `summary/quant_summary.txt` | The per-quantizer summary |
</details>
The quantization settings from the table above are logged as searchable params, alongside
the serving settings (`tensor_parallel_size`, `max_model_len`, `dtype`, `kv_cache_dtype`, …)
and `user` / `hostname` / `modelopt_version` / `git_sha` / `vllm_version` tags. The
`checkpoint_path` tag is the checkpoint being served, which is the same key
`examples/hf_ptq/hf_ptq.py` tags its runs with — so the PTQ run that produced a checkpoint
and every serve of it can be found together.
Other flags:
- `--mlflow_experiment` — defaults to
`$USER/vllm_serve_fakequant/<model basename>-<recipe name>`, falling back to
`$QUANT_CFG`/`$KV_QUANT_CFG` when no recipe is used.
- `--mlflow_run_name` — defaults to the UTC start time, `YYYYmmdd-HHMMSS`.
- `$MLFLOW_TRACKING_URI` enables tracking on its own; `--mlflow` overrides it. A URI taken
from the environment is best-effort — if the client is missing or the server is
unreachable the server warns and serves untracked. An explicit `--mlflow` fails loudly
instead.
Tracking needs the client: `pip install nvidia-modelopt[mlflow]` (already in this example's
`Dockerfile`). Authentication uses MLflow's own environment variables
(`MLFLOW_TRACKING_TOKEN`, or `MLFLOW_TRACKING_USERNAME` / `MLFLOW_TRACKING_PASSWORD`); with
`--distributed-executor-backend ray` those are forwarded to the workers along with the
tracking settings, since a Ray worker starts with a clean environment.
## Load QAT/PTQ model and serve in vLLM (WIP)
Step 1: export the model with bf16 weights and quantizer state. To export the model:
+34 -13
View File
@@ -21,6 +21,7 @@ from typing import Any
import torch
from transformers import AutoTokenizer
from vllm.v1.worker.gpu_worker import Worker as BaseWorker
from vllm_mlflow_utils import FakeQuantMlflowTracker
from vllm_ptq_utils import calibrate_fun, get_quant_config
from vllm_reload_utils import (
convert_dict_to_vllm,
@@ -51,7 +52,7 @@ quant_config: dict[str, Any] = {
}
def _fakequant_run_prolog_worker(self) -> None:
def _fakequant_run_prolog_worker(self, mlflow_tracker: FakeQuantMlflowTracker) -> None:
trust_remote_code = os.environ.get("TRUST_REMOTE_CODE", "false").lower() == "true"
tokenizer = AutoTokenizer.from_pretrained(
@@ -124,6 +125,8 @@ def _fakequant_run_prolog_worker(self) -> None:
calibrate_loop = calibrate_fun(calib_dataloader, self)
quant_cfg = get_quant_config(quant_config, model)
# Before calibration, which is the run this artifact is most wanted for if it dies.
mlflow_tracker.log_quant_config(quant_cfg)
with disable_compilation(model):
print("Quantizing model...")
@@ -141,6 +144,7 @@ def _fakequant_run_prolog_worker(self) -> None:
if not torch.distributed.is_initialized() or torch.distributed.get_rank() == 0:
mtq.print_quant_summary(model)
mlflow_tracker.log_quant_summary(model)
mtq.fold_weight(model)
for name, module in model.named_modules():
@@ -152,23 +156,40 @@ def _fakequant_run_prolog_worker(self) -> None:
class FakeQuantWorker(BaseWorker):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
# Inert unless a tracking URI was published and this is the rank-0 worker.
self.mlflow_tracker = FakeQuantMlflowTracker(self, quant_config)
def load_model(self, *args, **kwargs) -> None:
# The run opens here, before the weights load: an unreachable tracking server or a
# missing token then fails in seconds instead of after the load and calibration,
# and the log it captures covers both.
self.mlflow_tracker.start()
with self.mlflow_tracker.fail_on_error():
return super().load_model(*args, **kwargs)
@torch.inference_mode()
def determine_available_memory(self) -> int:
model = self.model_runner.model
if hasattr(model, "unwrap"):
model = model.unwrap()
with disable_compilation(model):
with self.mlflow_tracker.fail_on_error(), disable_compilation(model):
return super().determine_available_memory()
def compile_or_warm_up_model(self) -> float:
if (
quant_config["quant_cfg"]
or quant_config["kv_quant_cfg"]
or quant_config["modelopt_state_path"]
or quant_config["recipe_path"]
):
_fakequant_run_prolog_worker(self)
# Must return the base worker's compilation time (seconds). Returning None
# breaks vLLM V1 executor: initialize_from_config does max(compilation_times)
# across TP workers.
return super().compile_or_warm_up_model()
with self.mlflow_tracker.fail_on_error():
if (
quant_config["quant_cfg"]
or quant_config["kv_quant_cfg"]
or quant_config["modelopt_state_path"]
or quant_config["recipe_path"]
):
_fakequant_run_prolog_worker(self, self.mlflow_tracker)
# Must return the base worker's compilation time (seconds). Returning None
# breaks vLLM V1 executor: initialize_from_config does max(compilation_times)
# across TP workers.
compilation_time = super().compile_or_warm_up_model()
# The model is quantized and warmed up; everything after this is serving.
self.mlflow_tracker.finish("FINISHED")
return compilation_time
+417
View File
@@ -0,0 +1,417 @@
# 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.
"""MLflow tracking for the vLLM fake-quant server, mirroring ``examples/hf_ptq``.
The quantization this example performs happens inside the vLLM **worker** process, not in
``vllm_serve_fakequant.py``: the launcher is the API-server frontend, and the engine and its
workers are separate processes whose output it never sees. So the launcher only settles the
tracking configuration -- validating the URI, naming the experiment, recording the command
the user actually typed -- and publishes it through the environment, the same way every
other setting in this example reaches the workers. Rank 0 opens the run and uploads the
recipe, the effective quantization config, the calibration log and the quantizer summary.
The run covers weight load through warm-up and closes ``FINISHED`` when the server is ready
to serve, rather than staying open for the server's whole lifetime.
Nothing here imports vLLM, so the tracking can be exercised without it.
"""
import argparse
import contextlib
import os
import shutil
import tempfile
import warnings
from collections.abc import Iterator
from pathlib import Path
from typing import Any
from urllib.parse import urlparse
import yaml
import modelopt.torch.quantization as mtq
from modelopt.recipe import load_recipe
from modelopt.torch.utils.mlflow import (
MlflowRunLogger,
command_text,
default_experiment_name,
validate_tracking_uri,
)
TOOL_NAME = "vllm_serve_fakequant"
# Written by the launcher, read by the workers. The two MLflow-owned names are MLflow's own,
# so a shell that already exports them opts in without touching the command line.
TRACKING_URI_ENV = "MLFLOW_TRACKING_URI"
EXPERIMENT_ENV = "MLFLOW_EXPERIMENT_NAME"
RUN_NAME_ENV = "MODELOPT_MLFLOW_RUN_NAME"
REQUIRED_ENV = "MODELOPT_MLFLOW_REQUIRED"
COMMAND_ENV = "MODELOPT_MLFLOW_COMMAND"
# Everything the rank-0 worker needs in its environment to reach the tracking server. The
# credentials are never set here, only forwarded when the launching shell exported them --
# without that, a Ray worker authenticates as nobody and the run fails to open.
MLFLOW_ENV_VARS = frozenset(
{
TRACKING_URI_ENV,
EXPERIMENT_ENV,
RUN_NAME_ENV,
REQUIRED_ENV,
COMMAND_ENV,
"MLFLOW_TRACKING_TOKEN",
"MLFLOW_TRACKING_USERNAME",
"MLFLOW_TRACKING_PASSWORD",
"MLFLOW_TRACKING_INSECURE_TLS",
"MLFLOW_HTTP_REQUEST_MAX_RETRIES",
}
)
def add_mlflow_args(parser: argparse.ArgumentParser) -> None:
"""Add the MLflow tracking flags to the launcher's parser.
The multi-word flags are registered under both spellings. vLLM's
``FlexibleArgumentParser`` rewrites every ``--foo_bar`` on the command line to
``--foo-bar`` before matching, so the dashed spelling is the one that has to exist for
the flag to be reachable at all; the underscored spelling is what ``hf_ptq`` uses and
keeps these usable with a plain ``argparse.ArgumentParser``.
"""
parser.add_argument(
"--mlflow",
default=None,
help=(
"Track this server's calibration on an MLflow server "
"(e.g. https://<your-mlflow-server>/), 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. "
"MLflow's own $MLFLOW_TRACKING_URI enables tracking without this flag, which "
"overrides it. A URI taken from the environment is best-effort: if it is "
"unusable the server warns and serves untracked."
),
)
parser.add_argument(
"--mlflow-experiment",
"--mlflow_experiment",
default=None,
help=(
"MLflow experiment name. Default: "
f"$USER/{TOOL_NAME}/<model basename>-<recipe name, or the quantization config>."
),
)
parser.add_argument(
"--mlflow-run-name",
"--mlflow_run_name",
default=None,
help="MLflow run name. Default: the UTC start time as YYYYmmdd-HHMMSS.",
)
def resolve_mlflow_args(args: argparse.Namespace, parser: argparse.ArgumentParser) -> None:
"""Settle the tracking configuration and publish it to the worker processes.
Validating here rather than in the worker is what makes a typo in the URI fail at launch
instead of after the weights are on the GPUs. As in ``hf_ptq``, only ``--mlflow`` is a
deliberate request and therefore fatal when unusable; ``$MLFLOW_TRACKING_URI`` is
commonly exported for unrelated tooling and must not take a serve down with it.
"""
required = args.mlflow is not None
uri = args.mlflow or os.environ.get(TRACKING_URI_ENV) or None
if uri:
try:
uri = validate_tracking_uri(uri)
except ValueError as e:
if required:
parser.error(f"--mlflow: {e}")
warnings.warn(f"Ignoring ${TRACKING_URI_ENV}, continuing untracked: {e}")
uri = None
if not uri:
# A rejected URI must not reach the workers, which would try it again and fail there.
os.environ.pop(TRACKING_URI_ENV, None)
return
os.environ[TRACKING_URI_ENV] = uri
os.environ[REQUIRED_ENV] = "1" if required else "0"
# The workers' own sys.argv is vLLM's spawn plumbing; this is the command a user ran.
os.environ[COMMAND_ENV] = command_text()
os.environ[EXPERIMENT_ENV] = (
args.mlflow_experiment
or os.environ.get(EXPERIMENT_ENV)
or default_experiment_name(TOOL_NAME, args.model, quant_variant())
)
if args.mlflow_run_name:
os.environ[RUN_NAME_ENV] = args.mlflow_run_name
print(
f"[mlflow] tracking to {_without_credentials(uri)}, experiment {os.environ[EXPERIMENT_ENV]}"
)
def _without_credentials(uri: str) -> str:
"""Strip any ``user:token@`` from *uri*, for printing.
``MlflowRunLogger`` masks the same thing in every URI it prints or uploads, and this
line ends up in the worker log that the run itself uploads, so it has to match.
"""
parsed = urlparse(uri)
return parsed._replace(netloc=parsed.netloc.rpartition("@")[2]).geturl()
def quant_variant() -> str:
"""What distinguishes this serve of the model, for the default experiment name.
Read from the environment rather than taken as an argument because that is where this
example's quantization settings live, and both the launcher and a worker that was
started directly (``vllm serve --worker-cls fakequant_worker.FakeQuantWorker``) need it.
"""
if recipe := os.environ.get("RECIPE_PATH"):
return Path(recipe.rstrip("/")).stem
quant_cfg = os.environ.get("QUANT_CFG")
kv_quant_cfg = os.environ.get("KV_QUANT_CFG")
if quant_cfg or kv_quant_cfg:
return "-".join(cfg for cfg in (quant_cfg, kv_quant_cfg) if cfg)
if os.environ.get("MODELOPT_STATE_PATH"):
return "modelopt_state"
if os.environ.get("QUANT_FILE_PATH"):
return "quantizer_state"
return "unquantized"
class FakeQuantMlflowTracker:
"""Records one vLLM fake-quant worker's calibrate-and-serve as an MLflow run.
Inert unless the launcher published a tracking URI *and* this is the global rank-0
worker, so the worker needs no branching: every method is a no-op otherwise, and the
server behaves exactly as it did before tracking existed.
"""
def __init__(self, worker: Any, quant_config: dict[str, Any]):
"""Configure the run from the environment; nothing contacts the server yet."""
uri = os.environ.get(TRACKING_URI_ENV) or None
self._quant_config = quant_config
self._staging: Path | None = None
self._files: dict[str, Path] = {}
self._closed = False
self._worker = worker
self._logger = MlflowRunLogger(
uri or "",
os.environ.get(EXPERIMENT_ENV) or _fallback_experiment(worker),
run_name=os.environ.get(RUN_NAME_ENV) or None,
enabled=bool(uri) and getattr(worker, "rank", 0) == 0,
required=os.environ.get(REQUIRED_ENV) == "1",
)
@property
def enabled(self) -> bool:
"""Whether this worker records the run."""
return self._logger.enabled
def start(self) -> None:
"""Open the run and upload what is already known.
Called before the weights load so an unreachable server or a missing token fails in
seconds rather than after a multi-minute load followed by a calibration.
"""
if not self._logger.enabled:
return
self._staging = Path(tempfile.mkdtemp(prefix="modelopt-vllm-mlflow-"))
self._files = {"summary/quant_summary.txt": self._staging / ".quant_summary.txt"}
try:
self._logger.start(
params={**self._quant_config, **_vllm_params(self._worker)},
tags=_run_tags(self._worker),
texts=self._start_texts(),
files=self._files,
)
except BaseException:
# An explicit --mlflow is fatal here by design; leave no staging directory behind.
self._discard_staging()
raise
if not self._logger.enabled:
# A URI from the environment is best-effort: start() reports an unusable server
# by disabling itself rather than raising, and every later method -- finish()
# included -- returns before reaching the cleanup. So clean up here instead.
self._discard_staging()
def log_quant_config(self, quant_cfg: Any) -> None:
"""Upload the merged ``QUANT_CFG``/``KV_QUANT_CFG`` config, when that is what ran.
Only for the preset path, where this is the sole record of what was applied: the
params carry the preset *names*, while the config that reaches ``mtq.quantize`` is
those two deep-copied, merged, and -- for an MLA model -- extended at runtime with
``*kv_c_bmm_quantizer`` / ``*k_pe_bmm_quantizer`` by inspecting the loaded model.
A recipe run skips it: ``get_quant_config`` returns the recipe's ``quantize``
section unchanged, which ``recipe/resolved_recipe.yaml`` already carries.
Uploaded as soon as it is known rather than at the end, because a run that dies
during calibration is exactly the one this artifact is wanted for.
"""
if not self._logger.enabled or self._quant_config.get("recipe_path"):
return
try:
text = _dump_yaml(quant_cfg)
except Exception as e:
# Never let an unserializable config take down a serve that would have worked.
print(f"[mlflow] WARNING: could not serialize the quantization config: {e}")
return
self._logger.log_text("recipe/quant_cfg.yaml", text)
def log_quant_summary(self, model: Any) -> None:
"""Stage the per-quantizer summary for upload; a no-op when untracked."""
if not self._logger.enabled or self._staging is None:
return
# Writes .quant_summary.txt and prints only its path, so the console copy the caller
# already printed is not repeated.
mtq.print_quant_summary(model, output_dir=str(self._staging))
def finish(self, status: str) -> None:
"""Upload the log and the summary, and close the run with *status*.
Only the first call has an effect: vLLM drives the worker through several guarded
steps, and a run already closed as ``FINISHED`` must not be reopened or downgraded
by a failure in whatever the server does next.
"""
if not self._logger.enabled or self._closed:
return
self._closed = True
try:
self._logger.finish(status, files=self._files)
finally:
self._discard_staging()
def _discard_staging(self) -> None:
if self._staging is not None:
shutil.rmtree(self._staging, ignore_errors=True)
self._staging = None
@contextlib.contextmanager
def fail_on_error(self) -> Iterator[None]:
"""Close the run as ``FAILED`` if the wrapped step raises, then re-raise.
vLLM drives the worker through several calls; without this a failure in any of them
would leave the run ``RUNNING`` forever, with no log attached.
"""
try:
yield
except BaseException:
self.finish("FAILED")
raise
def _start_texts(self) -> dict[str, str]:
texts = {}
if command := os.environ.get(COMMAND_ENV):
texts["command.txt"] = command
if recipe_path := self._quant_config.get("recipe_path"):
# The resolved recipe, not the source file: a recipe may be a directory or use
# $imports, and only the resolved form stands alone.
texts["recipe/resolved_recipe.yaml"] = _dump_yaml(
load_recipe(recipe_path).model_dump(mode="json")
)
return texts
def _dump_yaml(value: Any) -> str:
"""YAML for an artifact, unwrapping a pydantic model first.
A recipe's ``quantize`` section is a ``QuantizeConfig``, which ``yaml.safe_dump`` cannot
represent. Raising here is deliberate: a ``repr`` fallback would upload an unparseable
one-line blob under a ``.yaml`` name, which looks like a successful artifact until
someone tries to read it. The caller turns the failure into a warning instead.
"""
if hasattr(value, "model_dump"):
value = value.model_dump(mode="json")
return yaml.safe_dump(value, sort_keys=False)
def _fallback_experiment(worker: Any) -> str:
"""Name the experiment from the worker itself, for a worker started without the launcher."""
model = _model_config_value(worker, "model") or "unknown"
return default_experiment_name(TOOL_NAME, str(model), quant_variant())
def _model_config_value(worker: Any, name: str, default: Any = None) -> Any:
config = getattr(getattr(worker, "vllm_config", None), "model_config", None)
return getattr(config, name, default)
def _vllm_params(worker: Any) -> dict[str, Any]:
"""Serving settings worth searching on, best-effort across vLLM versions.
Every field is read through ``getattr`` with a default: these configs are reshuffled
between vLLM releases, and a renamed attribute must not take down a serve that would
otherwise have worked.
"""
vllm_config = getattr(worker, "vllm_config", None)
model_config = getattr(vllm_config, "model_config", None)
parallel_config = getattr(vllm_config, "parallel_config", None)
cache_config = getattr(vllm_config, "cache_config", None)
params = {
"model": _stringify(getattr(model_config, "model", None)),
"served_model_name": _stringify(getattr(model_config, "served_model_name", None)),
"dtype": _stringify(getattr(model_config, "dtype", None)),
"max_model_len": getattr(model_config, "max_model_len", None),
"enforce_eager": getattr(model_config, "enforce_eager", None),
"quantization": _stringify(getattr(model_config, "quantization", None)),
"tensor_parallel_size": getattr(parallel_config, "tensor_parallel_size", None),
"pipeline_parallel_size": getattr(parallel_config, "pipeline_parallel_size", None),
"data_parallel_size": getattr(parallel_config, "data_parallel_size", None),
"world_size": getattr(parallel_config, "world_size", None),
"kv_cache_dtype": _stringify(getattr(cache_config, "cache_dtype", None)),
"vllm_version": _vllm_version(),
}
return {k: v for k, v in params.items() if v is not None}
def _run_tags(worker: Any) -> dict[str, str]:
"""Tags shared with ``hf_ptq``, so a checkpoint's PTQ run and the serves of it join up.
``checkpoint_path`` is the checkpoint being served, which is the same key ``hf_ptq``
tags with the checkpoint it *writes* and the one an evaluation is pointed at
(NEL's ``deployment.checkpoint_path``).
"""
model = _stringify(_model_config_value(worker, "model")) or "unknown"
tags = {
"tool": TOOL_NAME,
"model": Path(model).name,
"checkpoint_path": _resolved(model),
"vllm_version": _vllm_version(),
}
if served := _stringify(_model_config_value(worker, "served_model_name")):
tags["served_model_name"] = served
return tags
def _resolved(model: str) -> str:
"""The served checkpoint as an absolute path, or unchanged for a Hugging Face model id."""
return str(Path(model).resolve()) if os.path.exists(model) else model
def _stringify(value: Any) -> str | None:
"""Flatten a config value to a string; vLLM stores some of these as one-element lists."""
if value is None:
return None
if isinstance(value, (list, tuple)):
return ",".join(str(v) for v in value) or None
return str(value)
def _vllm_version() -> str:
try:
import vllm
return str(vllm.__version__)
except ImportError:
return "unknown"
@@ -59,6 +59,7 @@ import vllm
from packaging import version
from vllm.entrypoints.openai.api_server import run_server
from vllm.entrypoints.openai.cli_args import make_arg_parser
from vllm_mlflow_utils import MLFLOW_ENV_VARS, add_mlflow_args, resolve_mlflow_args
vllm_version = version.parse(vllm.__version__)
if vllm_version <= version.parse("0.11.0"):
@@ -68,6 +69,8 @@ else:
# Env vars to copy from the driver to Ray workers (must match fakequant_worker / vllm_ptq_utils).
# The MLflow ones are settled by resolve_mlflow_args() below, after this list is published:
# Ray reads the values when it creates the actors, so naming them here is enough.
additional_env_vars = {
"QUANT_DATASET",
"QUANT_CALIB_SIZE",
@@ -78,6 +81,7 @@ additional_env_vars = {
"CALIB_BATCH_SIZE",
"RECIPE_PATH",
"TRUST_REMOTE_CODE",
*MLFLOW_ENV_VARS,
}
try:
@@ -102,6 +106,7 @@ def main():
parser = FlexibleArgumentParser(description="vLLM model server with quantization support")
parser.add_argument("model", type=str, help="The path or name of the model to serve")
parser = make_arg_parser(parser)
add_mlflow_args(parser)
# Ensure workers can import our custom worker module when using spawn
repo_root = str(Path(__file__).resolve().parent)
if repo_root not in sys.path:
@@ -116,6 +121,9 @@ def main():
# Parse arguments
args = parser.parse_args()
# Settled before the engine starts, so an unusable tracking URI fails here rather than
# in a worker that has already loaded the weights.
resolve_mlflow_args(args, parser)
# Run the server
uvloop.run(run_server(args))
+30 -5
View File
@@ -42,7 +42,13 @@ from urllib.parse import urlparse
import modelopt
from modelopt.torch.utils.logging import TeeStream
__all__ = ["MlflowRunLogger", "current_user", "default_experiment_name", "validate_tracking_uri"]
__all__ = [
"MlflowRunLogger",
"command_text",
"current_user",
"default_experiment_name",
"validate_tracking_uri",
]
# MLflow experiment names are stored in a VARCHAR(256) column by the SQL-backed stores. The
# per-component cap stops one pathological component from crowding out the others; the name
@@ -188,9 +194,14 @@ def _git_sha() -> str:
return "unknown"
def _command_text() -> str:
"""The invocation, as a copy-pasteable line."""
lines = [shlex.join([sys.executable, *_redact_argv(sys.argv)])]
def command_text(argv: list[str] | None = None) -> str:
"""The invocation, as a copy-pasteable line, with credentials masked.
*argv* defaults to this process's own ``sys.argv``. Pass another process's argv when the
run is opened somewhere the user never typed a command -- a worker subprocess, whose own
``sys.argv`` is an implementation detail rather than a reproducible invocation.
"""
lines = [shlex.join([sys.executable, *_redact_argv(sys.argv if argv is None else argv)])]
world_size = int(os.environ.get("WORLD_SIZE", "1"))
if world_size > 1:
lines += [
@@ -343,6 +354,20 @@ class MlflowRunLogger:
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.
For a value that is only settled midway through the run and is worth having even if
the run later crashes -- the quantization config a calibration is about to apply,
say. :meth:`start` and :meth:`finish` cover everything known at the two ends.
"""
if not self.enabled or self._run is None:
return
try:
self._log_texts({artifact_path: text})
except Exception as e:
print(f"[mlflow] WARNING: could not upload {artifact_path}: {e}")
def _abort_run(self) -> None:
"""End a run that failed before :meth:`start` returned, so it is not left RUNNING."""
if self._run is None:
@@ -431,7 +456,7 @@ class MlflowRunLogger:
# The version is a tag as well, for searching; the artifact travels with the run.
self._log_texts(
{
"command.txt": _command_text(),
"command.txt": command_text(),
"version.txt": f"{modelopt.__version__}\n",
**(texts or {}),
}
@@ -0,0 +1,514 @@
# 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.
"""MLflow tracking for the vLLM fake-quant server.
``vllm_mlflow_utils`` deliberately imports no vLLM, so the whole launcher-to-worker
handover is exercised here without a GPU, a running server, or the mlflow client.
"""
import argparse
import getpass
import importlib
import os
import sys
from pathlib import Path
from types import SimpleNamespace
import pytest
import yaml
_EXAMPLES_DIR = Path(__file__).resolve().parents[3] / "examples" / "vllm_serve"
URI = "https://mlflow.example.com"
# Fake credentials for the redaction test. TruffleHog's URI detector flags any
# scheme://user:pass@host, so the marker sits on the definition; this test exists
# precisely to prove such credentials are masked.
CREDS_URI = "https://alice:s3cret@mlflow.example.com" # trufflehog:ignore
# What the launcher publishes plus what it reads, cleared between tests so one case cannot
# leak tracking configuration into the next.
_TRACKED_ENV = (
"MLFLOW_TRACKING_URI",
"MLFLOW_EXPERIMENT_NAME",
"MODELOPT_MLFLOW_RUN_NAME",
"MODELOPT_MLFLOW_REQUIRED",
"MODELOPT_MLFLOW_COMMAND",
"RECIPE_PATH",
"QUANT_CFG",
"KV_QUANT_CFG",
"MODELOPT_STATE_PATH",
"QUANT_FILE_PATH",
)
QUANT_CONFIG = {
"dataset": "cnn_dailymail",
"calib_size": 512,
"quant_cfg": "NVFP4_DEFAULT_CFG",
"kv_quant_cfg": None,
"quant_file_path": None,
"modelopt_state_path": None,
"calib_batch_size": 1,
"recipe_path": None,
}
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):
monkeypatch.setattr(getpass, "getuser", lambda: "tester")
for name in _TRACKED_ENV:
monkeypatch.delenv(name, raising=False)
@pytest.fixture
def mlflow_utils(monkeypatch):
monkeypatch.syspath_prepend(str(_EXAMPLES_DIR))
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()
parser.add_argument("model")
mlflow_utils.add_mlflow_args(parser)
argv = [model, *(token for k, v in flags.items() for token in (f"--{k}", v))]
monkeypatch.setattr(sys, "argv", ["vllm_serve_fakequant.py", *argv])
args = parser.parse_args(argv)
mlflow_utils.resolve_mlflow_args(args, parser)
return args
def _worker(rank=0, model="/ckpts/Qwen3-0.6B", **model_config):
"""A stand-in for the vLLM worker, holding only what the tracker reads off it."""
return SimpleNamespace(
rank=rank,
vllm_config=SimpleNamespace(
model_config=SimpleNamespace(model=model, **model_config),
parallel_config=SimpleNamespace(tensor_parallel_size=8, world_size=8),
cache_config=SimpleNamespace(cache_dtype="auto"),
),
)
# --- launcher side -------------------------------------------------------------------
def test_flag_publishes_tracking_config_to_the_environment(mlflow_utils, monkeypatch):
_resolve(mlflow_utils, monkeypatch, mlflow=f"{URI}/")
assert os.environ["MLFLOW_TRACKING_URI"] == URI # trailing slash stripped
assert os.environ["MODELOPT_MLFLOW_REQUIRED"] == "1"
assert (
os.environ["MLFLOW_EXPERIMENT_NAME"] == "tester/vllm_serve_fakequant/Qwen3-0.6B-unquantized"
)
# The command a user typed, not the worker subprocess's own argv.
assert "vllm_serve_fakequant.py" in os.environ["MODELOPT_MLFLOW_COMMAND"]
def test_experiment_name_follows_the_quantization_settings(mlflow_utils, monkeypatch):
monkeypatch.setenv("QUANT_CFG", "NVFP4_DEFAULT_CFG")
monkeypatch.setenv("KV_QUANT_CFG", "FP8_KV_CFG")
_resolve(mlflow_utils, monkeypatch, mlflow=URI)
assert (
os.environ["MLFLOW_EXPERIMENT_NAME"]
== "tester/vllm_serve_fakequant/Qwen3-0.6B-NVFP4_DEFAULT_CFG-FP8_KV_CFG"
)
def test_recipe_names_the_experiment_when_set(mlflow_utils, monkeypatch):
monkeypatch.setenv("RECIPE_PATH", "/recipes/nvfp4_default-kv_fp8_cast.yaml")
monkeypatch.setenv("QUANT_CFG", "NVFP4_DEFAULT_CFG") # the recipe is authoritative
_resolve(mlflow_utils, monkeypatch, mlflow=URI)
assert (
os.environ["MLFLOW_EXPERIMENT_NAME"]
== "tester/vllm_serve_fakequant/Qwen3-0.6B-nvfp4_default-kv_fp8_cast"
)
def test_explicit_experiment_and_run_name_win(mlflow_utils, monkeypatch):
_resolve(
mlflow_utils,
monkeypatch,
mlflow=URI,
mlflow_experiment="team/sweep",
mlflow_run_name="calib-512",
)
assert os.environ["MLFLOW_EXPERIMENT_NAME"] == "team/sweep"
assert os.environ["MODELOPT_MLFLOW_RUN_NAME"] == "calib-512"
@pytest.mark.parametrize("sep", ["-", "_"])
def test_multiword_flags_accept_both_spellings(mlflow_utils, monkeypatch, 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."""
parser = argparse.ArgumentParser()
parser.add_argument("model")
mlflow_utils.add_mlflow_args(parser)
args = parser.parse_args(
[
"/ckpts/m",
"--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"
@pytest.mark.parametrize(
("env", "expected"),
[
({"RECIPE_PATH": "/r/w4a8.yaml"}, "w4a8"),
({"QUANT_CFG": "NVFP4_DEFAULT_CFG"}, "NVFP4_DEFAULT_CFG"),
({"KV_QUANT_CFG": "FP8_KV_CFG"}, "FP8_KV_CFG"),
({"MODELOPT_STATE_PATH": "/x/vllm_fq_modelopt_state.pth"}, "modelopt_state"),
({"QUANT_FILE_PATH": "/x/quantizer_state.pth"}, "quantizer_state"),
({}, "unquantized"),
],
)
def test_quant_variant_reads_the_examples_own_settings(mlflow_utils, monkeypatch, env, expected):
for name, value in env.items():
monkeypatch.setenv(name, value)
assert mlflow_utils.quant_variant() == expected
def test_explicit_flag_with_a_bad_uri_fails_at_launch(mlflow_utils, monkeypatch):
with pytest.raises(SystemExit):
_resolve(mlflow_utils, monkeypatch, mlflow="mlflow.example.com")
def test_bad_uri_from_the_environment_serves_untracked(mlflow_utils, monkeypatch):
monkeypatch.setenv("MLFLOW_TRACKING_URI", "sqlite:///mlflow.db")
with pytest.warns(UserWarning, match="continuing untracked"):
_resolve(mlflow_utils, monkeypatch)
# Cleared, so a worker does not retry the URI the launcher already rejected.
assert "MLFLOW_TRACKING_URI" not in os.environ
def test_printed_uri_carries_no_credentials(mlflow_utils, monkeypatch, capsys):
"""This line lands in the worker log that the run uploads, so it has to be masked the
same way MlflowRunLogger masks every URI it prints."""
_resolve(mlflow_utils, monkeypatch, mlflow=CREDS_URI)
printed = capsys.readouterr().out
assert "s3cret" not in printed and "alice" not in printed
assert "https://mlflow.example.com" in printed
# Only the display is masked; the workers still need the credentials to authenticate.
assert os.environ["MLFLOW_TRACKING_URI"] == CREDS_URI
def test_environment_uri_is_best_effort_not_required(mlflow_utils, monkeypatch):
monkeypatch.setenv("MLFLOW_TRACKING_URI", URI)
_resolve(mlflow_utils, monkeypatch)
assert os.environ["MODELOPT_MLFLOW_REQUIRED"] == "0"
def test_no_uri_publishes_nothing(mlflow_utils, monkeypatch):
_resolve(mlflow_utils, monkeypatch)
assert "MLFLOW_TRACKING_URI" not in os.environ
assert "MODELOPT_MLFLOW_COMMAND" not in os.environ
def test_ray_copy_list_covers_what_the_launcher_sets(mlflow_utils):
"""Ray copies only the names on this list, so a name the worker needs must be on it."""
assert {
mlflow_utils.TRACKING_URI_ENV,
mlflow_utils.EXPERIMENT_ENV,
mlflow_utils.RUN_NAME_ENV,
mlflow_utils.REQUIRED_ENV,
mlflow_utils.COMMAND_ENV,
} <= mlflow_utils.MLFLOW_ENV_VARS
# Credentials are never set here, only forwarded; without them a worker authenticates
# as nobody and the run fails to open.
assert "MLFLOW_TRACKING_TOKEN" in mlflow_utils.MLFLOW_ENV_VARS
# --- worker side ---------------------------------------------------------------------
def test_tracker_is_inert_without_a_tracking_uri(mlflow_utils):
tracker = mlflow_utils.FakeQuantMlflowTracker(_worker(), QUANT_CONFIG)
assert not tracker.enabled
# Every entry point is safe to call unconditionally, so the worker needs no branching.
tracker.start()
tracker.log_quant_config({"quant_cfg": {}})
tracker.log_quant_summary(object())
tracker.finish("FINISHED")
def test_only_rank_zero_records_the_run(mlflow_utils, monkeypatch):
monkeypatch.setenv("MLFLOW_TRACKING_URI", URI)
assert mlflow_utils.FakeQuantMlflowTracker(_worker(rank=0), QUANT_CONFIG).enabled
assert not mlflow_utils.FakeQuantMlflowTracker(_worker(rank=3), QUANT_CONFIG).enabled
def test_worker_names_the_experiment_when_started_without_the_launcher(mlflow_utils, monkeypatch):
"""`vllm serve --worker-cls fakequant_worker.FakeQuantWorker` sets no experiment name."""
monkeypatch.setenv("MLFLOW_TRACKING_URI", URI)
monkeypatch.setenv("QUANT_CFG", "NVFP4_DEFAULT_CFG")
tracker = mlflow_utils.FakeQuantMlflowTracker(_worker(), QUANT_CONFIG)
assert (
tracker._logger.experiment_name
== "tester/vllm_serve_fakequant/Qwen3-0.6B-NVFP4_DEFAULT_CFG"
)
def test_start_uploads_the_launchers_command_and_the_serving_settings(
mlflow_utils, monkeypatch, fake_mlflow
):
monkeypatch.setenv("MLFLOW_TRACKING_URI", URI)
monkeypatch.setenv("MLFLOW_EXPERIMENT_NAME", "tester/vllm_serve_fakequant/Qwen3-0.6B-nvfp4")
monkeypatch.setenv(
"MODELOPT_MLFLOW_COMMAND", "python3 vllm_serve_fakequant.py /ckpts/x -tp 8\n"
)
tracker = mlflow_utils.FakeQuantMlflowTracker(
_worker(served_model_name="qwen", max_model_len=4096), QUANT_CONFIG
)
tracker.start()
tracker.finish("FINISHED")
assert fake_mlflow.experiment == "tester/vllm_serve_fakequant/Qwen3-0.6B-nvfp4"
assert fake_mlflow.texts["command.txt"].startswith("python3 vllm_serve_fakequant.py")
# The quantization settings and the serving settings are both searchable.
assert fake_mlflow.params["quant_cfg"] == "NVFP4_DEFAULT_CFG"
assert fake_mlflow.params["calib_size"] == 512
assert fake_mlflow.params["tensor_parallel_size"] == 8
assert fake_mlflow.params["max_model_len"] == 4096
# The join key with the hf_ptq run that produced the checkpoint being served.
assert fake_mlflow.tags["checkpoint_path"] == "/ckpts/Qwen3-0.6B"
assert fake_mlflow.tags["model"] == "Qwen3-0.6B"
assert fake_mlflow.tags["tool"] == "vllm_serve_fakequant"
assert fake_mlflow.tags["served_model_name"] == "qwen"
assert fake_mlflow.status == "FINISHED"
def test_quant_config_is_uploaded_before_calibration(mlflow_utils, monkeypatch, fake_mlflow):
monkeypatch.setenv("MLFLOW_TRACKING_URI", URI)
tracker = mlflow_utils.FakeQuantMlflowTracker(_worker(), QUANT_CONFIG)
tracker.start()
tracker.log_quant_config({"quant_cfg": {"*weight_quantizer": {"num_bits": 4}}})
uploaded = fake_mlflow.texts["recipe/quant_cfg.yaml"]
tracker.finish("FINISHED")
# Present while the run was still open, so a crash during calibration keeps it.
assert yaml.safe_load(uploaded) == {"quant_cfg": {"*weight_quantizer": {"num_bits": 4}}}
def test_a_best_effort_start_failure_leaves_no_staging_directory(
mlflow_utils, monkeypatch, fake_mlflow
):
"""start() reports an unusable server by disabling itself rather than raising, and every
later method returns before the cleanup -- so the temp dir would outlive the process."""
monkeypatch.setenv("MLFLOW_TRACKING_URI", URI) # from the environment, so not required
def explode(*args, **kwargs):
raise ConnectionError("no route to host")
fake_mlflow.set_experiment = explode
tracker = mlflow_utils.FakeQuantMlflowTracker(_worker(), QUANT_CONFIG)
tracker.start()
assert not tracker.enabled # downgraded to untracked, the serve carries on
assert tracker._staging is None
tracker.finish("FINISHED") # still safe to call
def test_an_explicit_flag_start_failure_leaves_no_staging_directory(
mlflow_utils, monkeypatch, fake_mlflow
):
monkeypatch.setenv("MLFLOW_TRACKING_URI", URI)
monkeypatch.setenv("MODELOPT_MLFLOW_REQUIRED", "1") # --mlflow, so failure is fatal
def explode(*args, **kwargs):
raise ConnectionError("no route to host")
fake_mlflow.set_experiment = explode
tracker = mlflow_utils.FakeQuantMlflowTracker(_worker(), QUANT_CONFIG)
with pytest.raises(ConnectionError):
tracker.start()
assert tracker._staging is None
def test_a_recipe_run_does_not_duplicate_the_config(mlflow_utils, monkeypatch, fake_mlflow):
"""get_quant_config returns the recipe's quantize section unchanged, and
recipe/resolved_recipe.yaml already carries it -- properly serialized."""
monkeypatch.setenv("MLFLOW_TRACKING_URI", URI)
tracker = mlflow_utils.FakeQuantMlflowTracker(
_worker(), {**QUANT_CONFIG, "recipe_path": "/r/nvfp4.yaml", "quant_cfg": None}
)
tracker._logger.start() # bypass _start_texts, which would load the recipe from disk
tracker.log_quant_config(object())
tracker._logger.finish("FINISHED")
assert "recipe/quant_cfg.yaml" not in fake_mlflow.texts
def test_a_pydantic_config_is_dumped_as_yaml_not_repr(mlflow_utils, monkeypatch, fake_mlflow):
"""yaml.safe_dump cannot represent a QuantizeConfig; a repr fallback would upload an
unparseable blob under a .yaml name and look like it worked."""
monkeypatch.setenv("MLFLOW_TRACKING_URI", URI)
class FakeQuantizeConfig:
def model_dump(self, mode=None):
return {"quant_cfg": [{"quantizer_name": "*weight_quantizer", "enable": True}]}
tracker = mlflow_utils.FakeQuantMlflowTracker(_worker(), QUANT_CONFIG)
tracker.start()
tracker.log_quant_config(FakeQuantizeConfig())
uploaded = fake_mlflow.texts["recipe/quant_cfg.yaml"]
tracker.finish("FINISHED")
assert yaml.safe_load(uploaded) == {
"quant_cfg": [{"quantizer_name": "*weight_quantizer", "enable": True}]
}
def test_an_unserializable_config_warns_instead_of_killing_the_serve(
mlflow_utils, monkeypatch, fake_mlflow, capsys
):
monkeypatch.setenv("MLFLOW_TRACKING_URI", URI)
tracker = mlflow_utils.FakeQuantMlflowTracker(_worker(), QUANT_CONFIG)
tracker.start()
tracker.log_quant_config({"cfg": object()}) # no representer, no model_dump
tracker.finish("FINISHED")
assert "could not serialize the quantization config" in capsys.readouterr().out
assert "recipe/quant_cfg.yaml" not in fake_mlflow.texts
def test_quant_summary_is_uploaded_from_the_staging_directory(
mlflow_utils, monkeypatch, fake_mlflow
):
monkeypatch.setenv("MLFLOW_TRACKING_URI", URI)
tracker = mlflow_utils.FakeQuantMlflowTracker(_worker(), QUANT_CONFIG)
# Stand in for mtq.print_quant_summary(model, output_dir=...), which is what writes it.
def write_summary(model, output_dir):
Path(output_dir, ".quant_summary.txt").write_text("2 TensorQuantizers found in model\n")
monkeypatch.setattr(
importlib.import_module("modelopt.torch.quantization"),
"print_quant_summary",
write_summary,
)
tracker.start()
tracker.log_quant_summary(object())
tracker.finish("FINISHED")
artifact_path, content = fake_mlflow.artifacts["quant_summary.txt"]
assert artifact_path == "summary"
assert "2 TensorQuantizers" in content
def test_a_run_with_no_summary_uploads_none(mlflow_utils, monkeypatch, fake_mlflow):
"""A reload from MODELOPT_STATE_PATH on a non-zero rank writes no summary."""
monkeypatch.setenv("MLFLOW_TRACKING_URI", URI)
tracker = mlflow_utils.FakeQuantMlflowTracker(_worker(), QUANT_CONFIG)
tracker.start()
tracker.finish("FINISHED")
assert "quant_summary.txt" not in fake_mlflow.artifacts
def test_a_failed_step_closes_the_run_and_reraises(mlflow_utils, monkeypatch, fake_mlflow):
monkeypatch.setenv("MLFLOW_TRACKING_URI", URI)
tracker = mlflow_utils.FakeQuantMlflowTracker(_worker(), QUANT_CONFIG)
tracker.start()
with pytest.raises(RuntimeError, match="out of memory"), tracker.fail_on_error():
raise RuntimeError("out of memory")
assert fake_mlflow.status == "FAILED"
# The log is attached even though the worker never reached the end of warm-up.
assert any(name.endswith(".log") for name in fake_mlflow.artifacts)
def test_a_closed_run_is_not_reopened_by_a_later_step(mlflow_utils, monkeypatch, fake_mlflow):
"""vLLM drives the worker through several guarded steps after warm-up finishes."""
monkeypatch.setenv("MLFLOW_TRACKING_URI", URI)
tracker = mlflow_utils.FakeQuantMlflowTracker(_worker(), QUANT_CONFIG)
tracker.start()
tracker.finish("FINISHED")
with pytest.raises(RuntimeError), tracker.fail_on_error():
raise RuntimeError("the server died an hour later")
assert fake_mlflow.status == "FINISHED"
+49
View File
@@ -28,6 +28,7 @@ from modelopt.torch.utils.mlflow import (
MlflowRunLogger,
_git_sha,
_redact_argv,
command_text,
default_experiment_name,
validate_tracking_uri,
)
@@ -307,6 +308,54 @@ def test_command_flags_the_invisible_torchrun_wrapper(fake_mlflow, monkeypatch):
assert "WORLD_SIZE=8" in command and "not part of sys.argv" in command
def test_command_can_record_another_processs_invocation(monkeypatch):
"""A worker's own sys.argv is spawn plumbing, so the caller can supply the real one."""
monkeypatch.setattr(sys, "argv", ["-c", "from multiprocessing.spawn import spawn_main"])
command = command_text(["vllm_serve_fakequant.py", "/ckpts/model", "--api-key", "sk-secret"])
assert "vllm_serve_fakequant.py /ckpts/model" in command
assert "sk-secret" not in command
assert "spawn_main" not in command
def test_log_text_uploads_while_the_run_is_open(fake_mlflow):
"""For a value settled midway through, which a later crash would otherwise lose."""
logger = _logger()
logger.start()
logger.log_text("recipe/quant_cfg.yaml", "quant_cfg: {}\n")
uploaded_before_finish = fake_mlflow.texts.get("recipe/quant_cfg.yaml")
logger.finish("FAILED")
assert uploaded_before_finish == "quant_cfg: {}\n"
def test_log_text_is_inert_outside_an_open_run(fake_mlflow):
logger = _logger(enabled=False)
logger.log_text("recipe/quant_cfg.yaml", "quant_cfg: {}\n")
logger = _logger() # enabled, but never started
logger.log_text("recipe/quant_cfg.yaml", "quant_cfg: {}\n")
assert fake_mlflow.texts == {}
def test_log_text_never_raises_when_the_upload_fails(fake_mlflow, capsys):
"""Losing one artifact must not take down the quantization that produced it."""
logger = _logger()
logger.start()
def explode(*args, **kwargs):
raise ConnectionError("no route to host")
fake_mlflow.log_text = explode
logger.log_text("recipe/quant_cfg.yaml", "quant_cfg: {}\n")
logger.finish("FINISHED")
assert "could not upload recipe/quant_cfg.yaml" in capsys.readouterr().out
def test_capture_includes_preconfigured_library_logging(fake_mlflow, monkeypatch):
"""transformers/huggingface_hub bind sys.stderr at import, long before capture starts."""
monkeypatch.setattr(sys, "argv", ["hf_ptq.py"])