Files
DS4Server/tools/mtplx-policy-reference.py
T

71 lines
3.8 KiB
Python

"""Record pinned Qwen server policy functions; no model loading or generation."""
import argparse
import itertools
import json
import os
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("model")
args = parser.parse_args()
print("reference: importing policy functions without model loading", flush=True)
from mtplx.server.openai import _server_runtime_env_overrides
from mtplx import generation as g
from mtplx.session_bank import _lazy_snapshot_enabled
keys = [
"MTPLX_AR_PIPELINE", "MTPLX_FAMILY_CAPTURE_COMMIT", "MTPLX_FUSED_HC_V3",
"MTPLX_FUSED_GDN_INPROJ", "MTPLX_FUSED_GATE_UP", "MTPLX_FUSED_GDN_CONVNORM",
"MTPLX_FUSED_GDN_STEP", "MTPLX_FUSED_CONVNORM_VERIFY", "MTPLX_QSA_GATHER",
"MTPLX_NAX_VERIFY", "MTPLX_SKIP_VERIFY_SNAPSHOT",
]
cases = [{}]
cases += [{key: value} for key in keys for value in ("0", " YES ", "", "unknown")]
cases += [dict((k, v) for k, v in zip(
("MTPLX_COMPILED_GDN", "MTPLX_QWEN4EXP_COMPILE"), pair) if v is not None)
for pair in itertools.product((None, "1", "0", "", "junk"), repeat=2)]
cases += [{key: value} for key in (
"MTPLX_SESSION_LAZY_SNAPSHOT", "MTPLX_SESSION_STORE_ON_PREFILL",
"MTPLX_GDN_BOUNDARY_CAPTURE", "MTPLX_ASYNC_AR", "MTPLX_EVAL_AUDIT",
"MTPLX_LAZY_MTP_HISTORY_APPEND", "MTPLX_DEFER_REPAIR_EVAL",
"MTPLX_PREFILL_EXTERNAL_EMIT_LOGITS",
) for value in ("0", "on", "", "junk")]
cases += [{key: value} for key in (
"MTPLX_SESSION_STORE_ON_PREFILL_MIN_SUFFIX", "MTPLX_SMALL_SUFFIX_FUSED_MAX",
"MTPLX_GDN_BOUNDARY_MAX", "MTPLX_GDN_BOUNDARY_TAIL_INTERVAL",
) for value in ("-3", "0", "17", "1_024", "bad")]
cases += [{"MTPLX_PREFILL_OMLX_EXTERNAL": "1"},
{"MTPLX_PREFILL_STOCK_CACHE_ONLY": "1"},
{"MTPLX_PREFILL_STOCK_CACHE_ONLY": "1", "MTPLX_ALLOW_UNSAFE_PREFILL_STOCK_CACHE_ONLY": "1"},
{"MTPLX_ASYNC_AR": "1", "MTPLX_EVAL_AUDIT": "audit"}]
cases += [{key: value} for key in ("MTPLX_FAMILY_CAPTURE_COMMIT", "MTPLX_DEFER_REPAIR_EVAL")
for value in ("enable", "enabled", "disable", "disabled")]
baseline = {k: v for k, v in os.environ.items() if not k.startswith("MTPLX_")}
for mtp, env in itertools.product((False, True), cases):
os.environ.clear()
os.environ.update(baseline)
os.environ.update(env)
overrides = _server_runtime_env_overrides(argparse.Namespace(
model=args.model, generation_mode="mtp" if mtp else "ar"), None)
os.environ.update(overrides)
# Native app selects bounded prefill; emulate that request-local policy.
os.environ["MTPLX_SUSTAINED_PREFILL"] = "1"
error = None
options = None
try:
with g.prefill_chunk_size_override(2048):
options = dict(chunk=g._prefill_chunk_size(), final_only=g._final_logits_prefill_enabled(),
external_cache_only=g._prefill_external_cache_only_enabled(),
external_emit_logits=g._prefill_external_emit_logits_enabled(),
lazy_kv=_lazy_snapshot_enabled(), store=g._store_on_prefill_env_enabled(),
store_min=g._store_on_prefill_min_suffix(), fused_max=g._small_suffix_fused_max(),
boundaries=g._gdn_boundary_capture_enabled(), boundary_max=g._gdn_boundary_max_count(),
boundary_tail=g._gdn_boundary_tail_interval(), pipeline=g._env_truthy("MTPLX_AR_PIPELINE"),
sync_eval=not g._env_truthy("MTPLX_ASYNC_AR") or bool(os.environ.get("MTPLX_EVAL_AUDIT")),
capture=g._family_capture_commit_enabled() if mtp else None,
lazy_history=g._env_truthy("MTPLX_LAZY_MTP_HISTORY_APPEND"),
defer_repair=g._defer_repair_eval() if mtp else None)
except ValueError as exc:
error = str(exc)
print(json.dumps(dict(event="policy_reference", mtp=mtp, environment=env,
overrides=overrides, options=options, error=error)), flush=True)