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
16 changes: 16 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,22 @@ adhere to [Semantic Versioning](https://semver.org/spec/v2.0.0.html).
The server always decoded up to 32 requests together, which slows every
stream past the width where total throughput stops growing.

### Fixed
- Sampled speculative verify and the server's unfiltered sampler computed
logprob math at the activation dtype; float16 rows reached categorical
unwidened. Both now widen to float32 first (greedy paths unchanged).
- MiniMax MSA indexer projections stay F32 under float16 activations; the
bfloat16 exact-bit narrowing does not exist at float16.
- Non-streamed replies lost the leading indentation of their first content
line: the reasoning splitters (ATEM/harmony and the think-tag path all
served models use) stripped all content whitespace. Verbatim code answers
came back misindented on line one. Content now keeps first-line indent
and still drops leading blank lines and trailing whitespace.
- gemma-4 applied the final logit softcap at the activation dtype, rounding
every logit by up to ~0.12 nats at bfloat16 and flipping near-tie top-1
picks. The softcap now computes in float32 and emits float32 logits, like
muse-glimmer. GMLX_G4_SOFTCAP_F32=0 restores the old behavior.

## [0.3.2] - 2026-08-13

### Changed
Expand Down
5 changes: 5 additions & 0 deletions gmlx/loader.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,7 @@
from .sparse_sdpa import install_sparse_sdpa
from .arrays_cache_fix import install_arrays_cache_fix
from .gemma4_sync import install_gemma4_nosync
from .softcap_f32 import install_gemma4_softcap_f32
from .quantized_sdpa_fix import install_quantized_sdpa_mask_fix
from .rope_batch_fix import install_rope_batch_fix
from .rotating_cache_fix import install_rotating_cache_fix
Expand Down Expand Up @@ -2930,6 +2931,8 @@ def _install_and_load(
log("[install] gemma-4 host-sync-free masks/rope offsets active")
if install_gemma4_batched_sdpa() and _gemma4_target(model):
log("[install] gemma-4 hd512 batched-decode row route active")
if install_gemma4_softcap_f32() and _gemma4_target(model):
log("[install] gemma-4 float32 logit softcap active")
if install_cascade_sdpa() and install_cascade_stamp():
log("[install] shared-prefix cascade decode route active")
if install_sparse_sdpa():
Expand Down Expand Up @@ -3391,6 +3394,8 @@ def load_model(
_log("[install] gemma-4 host-sync-free masks/rope offsets active")
if install_gemma4_batched_sdpa() and _gemma4_target(model):
_log("[install] gemma-4 hd512 batched-decode row route active")
if install_gemma4_softcap_f32() and _gemma4_target(model):
_log("[install] gemma-4 float32 logit softcap active")
if install_cascade_sdpa() and install_cascade_stamp():
_log("[install] shared-prefix cascade decode route active")
if install_sparse_sdpa():
Expand Down
13 changes: 7 additions & 6 deletions gmlx/minimax_m3_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -637,12 +637,13 @@ def rename(k):
and v.dtype == mx.float32
):
# Some MSA GGUFs store the indexer projections as F32. The
# source checkpoint is BF16, so narrowing to the activation
# dtype halves the memory and keeps the indexer matmul at the
# graph's width; at bfloat16 it recovers the original bits
# exactly. Norms stay F32 (tiny, and the norm math runs in F32
# regardless).
v = v.astype(activation_dtype())
# source checkpoint is BF16, so narrowing to bfloat16 recovers
# the original bits exactly and halves the memory. float16
# cannot represent those bits; keep F32 there (MLX promotes
# the indexer matmul). Norms stay F32 (tiny, and the norm math
# runs in F32 regardless).
if activation_dtype() != mx.float16:
v = v.astype(activation_dtype())
renamed[rename(k)] = v
weights = renamed

Expand Down
12 changes: 11 additions & 1 deletion gmlx/reasoning.py
Original file line number Diff line number Diff line change
Expand Up @@ -275,10 +275,20 @@ def split_harmony_reply(text: str, *,
spans = filt.feed(text)
spans += filt.flush()
reasoning = "".join(t for t, m in spans if m == _REASON).strip()
content = "".join(t for t, m in spans if m == _ANSWER).strip()
content = trim_content_ws("".join(t for t, m in spans if m == _ANSWER))
return (reasoning or None, content)


def trim_content_ws(text: str) -> str:
"""Trim reply content: drop leading blank lines and trailing whitespace,
keep the first content line's indent (a full strip eats verbatim-code
indentation)."""
lines = (text or "").split("\n")
while lines and not lines[0].strip():
lines.pop(0)
return "\n".join(lines).rstrip()


_DIM = "\x1b[2m"
_RESET = "\x1b[0m"
_CLEAR_EOL = "\x1b[K"
Expand Down
42 changes: 39 additions & 3 deletions gmlx/server_patches/chat_behavior.py
Original file line number Diff line number Diff line change
Expand Up @@ -259,7 +259,38 @@ def _make_thinking_budget_criteria(self, args, input_ids):
def _strip_xtml_sections(content: str) -> str:
for m in _XTML_SECTION_MARKERS:
content = content.replace(m, "")
return content.strip()
return _trim_content_ws(content)


def _trim_content_ws(text: str) -> str:
from ..reasoning import trim_content_ws
return trim_content_ws(text)


def _split_thinking_keep_indent(rs, cls, text, start_token, end_token):
"""rs._split_thinking with content whitespace preserved. Same marker
walk as stock; only the content-side .strip() calls are narrowed
(upstream body seam-fingerprinted via app._split_thinking_text)."""
if not text:
return None, text
for start_marker, end_marker in cls._build_open_close_markers(
start_token, end_token):
start = text.find(start_marker)
end = text.find(end_marker, start if start >= 0 else 0)
if 0 <= start < end:
reasoning = text[start + len(start_marker):end].strip()
content = rs._strip_content_markers(
text[:start] + text[end + len(end_marker):])
return reasoning or None, _trim_content_ws(content)
if end_marker in text:
reasoning, content = text.split(end_marker, 1)
reasoning = rs._clean_reasoning(reasoning, start_marker)
return (reasoning or None,
_trim_content_ws(rs._strip_content_markers(content)))
if start_marker in text:
reasoning = rs._clean_reasoning(text, start_marker)
return reasoning or None, ""
return None, _trim_content_ws(rs._strip_content_markers(text))

_LAST_RENDERED_PROMPT: contextvars.ContextVar = contextvars.ContextVar(
"kq_last_rendered_prompt", default=None)
Expand Down Expand Up @@ -370,8 +401,10 @@ def _split_thinking_text(text, thinking_start_token=None,
# so this gate cannot misfire on it.)
from ..reasoning import split_harmony_reply
return split_harmony_reply(text, start_in_header=in_header)
reasoning, content = split(
text, thinking_start_token, thinking_end_token)
# Stock walk with content whitespace preserved: the stock
# splitter's .strip() eats first-line code indent.
reasoning, content = _split_thinking_keep_indent(
rs, cls, text, thinking_start_token, thinking_end_token)
if reasoning is None and content and retire_key.truncated_thinking(
text, cls._build_open_close_markers(
thinking_start_token, thinking_end_token),
Expand All @@ -383,6 +416,9 @@ def _split_thinking_text(text, thinking_start_token=None,

_split_thinking_text.__dict__[_STREAM_SEED_FLAG] = True
app._split_thinking_text = _split_thinking_text
# The /v1/responses output builder calls the responses_state module
# global directly; give it the same splitter.
rs._split_thinking = _split_thinking_text


# ignore-eos: forced-length decode (server-level)
Expand Down
6 changes: 4 additions & 2 deletions gmlx/server_patches/sampling.py
Original file line number Diff line number Diff line change
Expand Up @@ -283,7 +283,9 @@ def __call__(self, logprobs):
# them so the returned ids keep the caller's shape (matches top_p_sampling).
import mlx.core as mx
if not self._has_filter:
return mx.random.categorical(logprobs * (1.0 / self.temperature), axis=-1)
# f32: float16 logprobs would otherwise reach categorical unwidened
lp = logprobs.astype(mx.float32)
return mx.random.categorical(lp * (1.0 / self.temperature), axis=-1)
lead = logprobs.shape[:-1]
lp2 = logprobs.reshape(-1, logprobs.shape[-1])
masked, part, order = self._filtered(lp2)
Expand All @@ -300,7 +302,7 @@ def _cat(row, key):
return mx.random.categorical(row, key=key)

if not self._has_filter:
scaled = logprobs * (1.0 / self.temperature)
scaled = logprobs.astype(mx.float32) * (1.0 / self.temperature)
return mx.vmap(_cat, in_axes=(0, 0))(scaled, keys)
# only this tiny categorical-over-k is per-row keyed; the filter is batched.
masked, part, order = self._filtered(logprobs)
Expand Down
36 changes: 36 additions & 0 deletions gmlx/softcap_f32.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,36 @@
"""Float32 final-logit softcap for gemma-4.

Upstream computes tanh(x/cap)*cap at the activation dtype. Each 16-bit
logit then rounds by ~cap*eps: ~0.12 nats at bfloat16 with cap 30, enough
to flip near-tie top-1 picks. Compute in float32 and emit float32 logits,
matching the muse-glimmer softcap path. Covers the mlx-lm gemma4_text and
mlx-vlm gemma4 seams; install-once, GMLX_G4_SOFTCAP_F32=0 disables.
"""
from __future__ import annotations

import mlx.core as mx

from .envflags import env_bool

_installed = False


def _logit_softcap_f32(softcap, x):
return mx.tanh(x.astype(mx.float32) / softcap) * softcap


def install_gemma4_softcap_f32() -> bool:
global _installed
if not env_bool("GMLX_G4_SOFTCAP_F32", True):
return False
if _installed:
return True
from mlx_lm.models import gemma4_text as _lm
_lm.logit_softcap = _logit_softcap_f32
try:
from mlx_vlm.models.gemma4 import language as _vlm
_vlm.logit_softcap = _logit_softcap_f32
except ImportError:
pass
_installed = True
return True
3 changes: 3 additions & 0 deletions gmlx/speculative.py
Original file line number Diff line number Diff line change
Expand Up @@ -580,6 +580,7 @@ def _coupled_walk(lm, verify, draft_tokens: mx.array, sampler, budget: int,
pq_arr = _pq_graph(logits[:n_draft], pq)
else:
_pq_stats["misaligned"] += 1
logits = logits.astype(mx.float32) # 16-bit logprob math skews sampled verify
logprobs = logits - mx.logsumexp(logits, axis=-1, keepdims=True)
if sampler is None:
target = mx.argmax(logprobs, axis=-1) # [n_pos]
Expand Down Expand Up @@ -649,6 +650,7 @@ def _coupled_walk_batch(
target = verify.target_tokens # [B, n_pos]
else:
logits = lm.speculative_logits_from_hidden(verify.hidden) # [B, n_pos, V]
logits = logits.astype(mx.float32) # 16-bit logprob math skews sampled verify
logprobs = logits - mx.logsumexp(logits, axis=-1, keepdims=True)
if sampler is None:
target = mx.argmax(logprobs, axis=-1) # [B, n_pos]
Expand Down Expand Up @@ -1817,6 +1819,7 @@ def _gated_step(inputs):
logits = getattr(out, "logits", out)[:, -1, :]
if greedy:
return mx.argmax(logits, axis=-1).astype(token_dtype)
logits = logits.astype(mx.float32) # 16-bit logprob math skews sampled verify
logprobs = logits - mx.logsumexp(logits, axis=-1, keepdims=True)
return sampler(logprobs).reshape(-1).astype(token_dtype)

Expand Down
16 changes: 13 additions & 3 deletions tests/test_long_context.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,9 @@
real attention bug diverges in the first few characters.

Both are ``integration`` + ``slow`` and skip unless the env points at real
models (see ``conftest``). Select one arch with ``-k qwen2`` and/or shrink the
models (see ``conftest``). Both sweep the 16-bit activation dtypes (float16 is
the pre-Apple9 auto choice; its narrower range must hold at depth too). Select
one arch with ``-k qwen2``, one dtype with ``-k float16``, and/or shrink the
length with ``KQUANT_LONGCTX_TOKENS=4096`` for a quick run; the defaults sweep
every arch whose GGUF is present at >=16k and can take minutes on large models.
"""
Expand Down Expand Up @@ -48,6 +50,10 @@

TARGET = int(os.environ.get("KQUANT_LONGCTX_TOKENS", "16384"))

# Both 16-bit widths. float16 is the auto choice on pre-Apple9 GPUs and has
# a narrower exponent range, so depth signatures must hold under both.
ACTIVATION_DTYPES = ("bfloat16", "float16")

# Arches whose mlx-lm model has a known upper-context limitation vs the GGUF
# reference, so token-for-token parity is only expected up to this length:
# gemma2 - mlx-lm's gemma2 has no sliding-window attention; output is exact
Expand Down Expand Up @@ -206,8 +212,10 @@ def _llama_complete(binary, model_path, prompt_text, n, ctx):


# tests
@pytest.mark.parametrize("dtype", ACTIVATION_DTYPES)
@pytest.mark.parametrize("arch", CANDIDATE_ARCHES)
def test_long_decode_integrity(arch, gguf_index):
def test_long_decode_integrity(arch, dtype, gguf_index, monkeypatch):
monkeypatch.setenv("GMLX_ACTIVATION_DTYPE", dtype)
path = _require(gguf_index, arch)
model, config, tok = _load(path)
vocab = int(config["vocab_size"])
Expand Down Expand Up @@ -251,10 +259,12 @@ def _suppress_eos(_tokens, logits):
f"- degeneration")


@pytest.mark.parametrize("dtype", ACTIVATION_DTYPES)
@pytest.mark.parametrize("arch", CANDIDATE_ARCHES)
def test_long_prefill_parity(arch, gguf_index, llamacpp_bin):
def test_long_prefill_parity(arch, dtype, gguf_index, llamacpp_bin, monkeypatch):
from mlx_lm.generate import stream_generate

monkeypatch.setenv("GMLX_ACTIVATION_DTYPE", dtype)
if arch in PARITY_SKIP:
pytest.skip(f"{arch}: {PARITY_SKIP[arch]}")
path = _require(gguf_index, arch)
Expand Down
14 changes: 14 additions & 0 deletions tests/test_reasoning.py
Original file line number Diff line number Diff line change
Expand Up @@ -539,6 +539,20 @@ def test_split_harmony_reply_shapes():
assert split_harmony_reply(plain) == (None, "Six.")


def test_split_harmony_reply_keeps_first_line_code_indent():
from gmlx.reasoning import split_harmony_reply
full = ('<|channel|>analysis<|message|>Recall the body.'
"<|end|><|start|>assistant<|channel|>final<|message|>"
'\n """\n Docstring.\n')
r, c = split_harmony_reply(full)
assert r == "Recall the body."
assert c == ' """\n Docstring.'
# whitespace-only content collapses to empty
ws = ('<|channel|>analysis<|message|>x'
"<|end|><|start|>assistant<|channel|>final<|message|>\n \n")
assert split_harmony_reply(ws) == ("x", "")


# --- muse-glimmer ATEM channel ------------------------------------------------
#
# Routing is on the message HEADER, not on a "to=self" marker: any marker whose
Expand Down
33 changes: 33 additions & 0 deletions tests/test_server_patches.py
Original file line number Diff line number Diff line change
Expand Up @@ -330,6 +330,39 @@ def test_nonstream_split_strips_xtml_section_markers():
sp_chat._LAST_RENDERED_PROMPT.reset(tok)


def test_nonstream_split_keeps_first_line_code_indent():
"""The stock splitter .strip()s content, which deletes the first line's
leading indent from verbatim-code replies. The seeded splitter trims
newlines only, for every marker branch and the no-marker fallthrough."""
app_mod = importlib.import_module("mlx_vlm.server.app")
rs = importlib.import_module("mlx_vlm.server.responses_state")
sp.install_stream_thinking_seed()
split = app_mod._split_thinking_text
tok = sp_chat._LAST_RENDERED_PROMPT.set(None)
try:
# open+close pair
r, c = split("<think>plan</think>\n\n if x:\n y()")
assert r == "plan"
assert c == " if x:\n y()"
# close-only (prompt-opened block)
r, c = split("plan\n</think>\n\n\tdoc = doc || document;")
assert r == "plan"
assert c == "\tdoc = doc || document;"
# no marker at all
assert split(" indented") == (None, " indented")
# whitespace-only content collapses to empty
assert split("<think>plan</think>\n \n") == ("plan", "")
# XTML section strip keeps the indent too
text = ("plan<|close|>think<|sep|><|open|>response<|sep|>"
" return 4;<|close|>response<|sep|><|end_of_msg|>")
r, c = split(text, "<|open|>think<|sep|>", "<|close|>think<|sep|>")
assert c == " return 4;"
# the /v1/responses module global got the same splitter
assert rs._split_thinking is split
finally:
sp_chat._LAST_RENDERED_PROMPT.reset(tok)


def test_install_thinking_budget_fix_applies_and_idempotent():
# Fail-loud guard: asserts the seam bound to the REAL mlx-vlm symbol, so a
# rename of ResponseGenerator._make_thinking_budget_criteria turns into a CI
Expand Down