71 lines
3.8 KiB
Python
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)
|