feat(moa): per-reference-model max_tokens override

MoA reference_max_tokens is preset-level — one cap for all reference
models. When mixing a verbose model with a terse one, a single cap is
either too tight for the terse model or too loose for the verbose one.

Now each reference slot can optionally carry its own max_tokens:

  reference_models:
    - provider: openrouter
      model: deepseek/deepseek-v4-pro
      max_tokens: ***        # per-slot cap, overrides preset-level
    - provider: openai-codex
      model: gpt-5.5
      # no max_tokens → falls back to preset-level reference_max_tokens

_clean_slot (moa_config.py) preserves an optional max_tokens field on
the slot dict, coerced via _coerce_int_or_none. _run_reference
(moa_loop.py) reads slot-level max_tokens first, falling back to the
preset-level cap passed by the caller. Slots without the field are
unaffected — backward compatible.

Type hints on slot-handling functions updated from dict[str, str] to
dict[str, Any] to reflect the now-heterogeneous slot shape.
This commit is contained in:
Rain 2026-07-07 18:18:23 +02:00 committed by Teknium
parent ead9d7b256
commit bc7212cf93
5 changed files with 178 additions and 5 deletions

View file

@ -281,7 +281,7 @@ def _maybe_apply_moa_cache_control(
def _run_reference(
slot: dict[str, str],
slot: dict[str, Any],
ref_messages: list[dict[str, Any]],
*,
temperature: float | None = None,
@ -333,11 +333,16 @@ def _run_reference(
# (their caching is automatic; markers are ignored harmlessly, but we
# only decorate when the policy says the route honors them).
messages = _maybe_apply_moa_cache_control(messages, runtime)
# Per-slot max_tokens takes precedence over the preset-level
# reference_max_tokens passed in by the caller. This lets each
# reference model have its own output cap independently.
_slot_max_tokens: int | None = slot.get("max_tokens")
_effective_max_tokens = _slot_max_tokens if _slot_max_tokens is not None else max_tokens
response = call_llm(
task="moa_reference",
messages=messages,
temperature=temperature,
max_tokens=max_tokens,
max_tokens=_effective_max_tokens,
reasoning_config=_slot_reasoning_config(slot),
**runtime,
)
@ -398,7 +403,7 @@ def _run_reference(
def _run_references_parallel(
reference_models: list[dict[str, str]],
reference_models: list[dict[str, Any]],
ref_messages: list[dict[str, Any]],
*,
temperature: float | None = None,
@ -683,8 +688,8 @@ def aggregate_moa_context(
*,
user_prompt: str,
api_messages: list[dict[str, Any]],
reference_models: list[dict[str, str]],
aggregator: dict[str, str],
reference_models: list[dict[str, Any]],
aggregator: dict[str, Any],
temperature: float | None = None,
aggregator_temperature: float | None = None,
reference_max_tokens: int | None = None,

View file

@ -0,0 +1,2 @@
matarbot
# PR #60391 salvage

View file

@ -105,6 +105,14 @@ def _clean_slot(slot: Any) -> dict[str, Any] | None:
effort = _clean_reasoning_effort(slot.get("reasoning_effort"))
if effort:
clean["reasoning_effort"] = effort
# Optional per-slot max_tokens: overrides the preset-level
# reference_max_tokens for this specific reference model. None (the
# default) = no cap, so existing slots are unaffected. Allows tuning
# each advisor's output length independently — useful when one model
# is verbose and another is terse.
slot_mt = _coerce_int_or_none(slot.get("max_tokens"))
if slot_mt is not None:
clean["max_tokens"] = slot_mt
return clean

View file

@ -0,0 +1,82 @@
"""Tests for per-slot max_tokens in MoA reference calls.
Verifies that a ``max_tokens`` field on a reference slot dict takes
precedence over the preset-level ``reference_max_tokens``, and that
slot-level max_tokens=None falls back to the preset-level cap.
"""
from __future__ import annotations
from unittest.mock import MagicMock, patch
import pytest
class TestRunReferenceSlotMaxTokens:
"""_run_reference should prefer slot-level max_tokens over preset-level."""
def test_slot_max_tokens_overrides_preset_level(self):
"""When slot has max_tokens, it overrides the preset-level cap."""
from agent.moa_loop import _run_reference
captured_kwargs: dict = {}
def fake_call_llm(**kwargs):
captured_kwargs.update(kwargs)
mock_resp = MagicMock()
mock_resp.choices = [MagicMock(message=MagicMock(content="advice"))]
mock_resp.usage = None
return mock_resp
slot = {"provider": "openrouter", "model": "deepseek/deepseek-v4-pro", "max_tokens": 600}
with patch("agent.moa_loop._slot_runtime", return_value={"provider": "openrouter", "model": "deepseek/deepseek-v4-pro"}), \
patch("agent.moa_loop.call_llm", side_effect=fake_call_llm), \
patch("agent.moa_loop._maybe_apply_moa_cache_control", side_effect=lambda msgs, rt: msgs):
_run_reference(slot, [{"role": "user", "content": "hi"}], max_tokens=2000)
assert captured_kwargs.get("max_tokens") == 600
def test_slot_max_tokens_absent_falls_back_to_preset(self):
"""When slot has no max_tokens, the preset-level cap is used."""
from agent.moa_loop import _run_reference
captured_kwargs: dict = {}
def fake_call_llm(**kwargs):
captured_kwargs.update(kwargs)
mock_resp = MagicMock()
mock_resp.choices = [MagicMock(message=MagicMock(content="advice"))]
mock_resp.usage = None
return mock_resp
slot = {"provider": "openrouter", "model": "deepseek/deepseek-v4-pro"}
with patch("agent.moa_loop._slot_runtime", return_value={"provider": "openrouter", "model": "deepseek/deepseek-v4-pro"}), \
patch("agent.moa_loop.call_llm", side_effect=fake_call_llm), \
patch("agent.moa_loop._maybe_apply_moa_cache_control", side_effect=lambda msgs, rt: msgs):
_run_reference(slot, [{"role": "user", "content": "hi"}], max_tokens=2000)
assert captured_kwargs.get("max_tokens") == 2000
def test_both_none_means_uncapped(self):
"""When neither slot nor preset has max_tokens, it's None (uncapped)."""
from agent.moa_loop import _run_reference
captured_kwargs: dict = {}
def fake_call_llm(**kwargs):
captured_kwargs.update(kwargs)
mock_resp = MagicMock()
mock_resp.choices = [MagicMock(message=MagicMock(content="advice"))]
mock_resp.usage = None
return mock_resp
slot = {"provider": "openrouter", "model": "deepseek/deepseek-v4-pro"}
with patch("agent.moa_loop._slot_runtime", return_value={"provider": "openrouter", "model": "deepseek/deepseek-v4-pro"}), \
patch("agent.moa_loop.call_llm", side_effect=fake_call_llm), \
patch("agent.moa_loop._maybe_apply_moa_cache_control", side_effect=lambda msgs, rt: msgs):
_run_reference(slot, [{"role": "user", "content": "hi"}], max_tokens=None)
assert captured_kwargs.get("max_tokens") is None

View file

@ -461,3 +461,79 @@ def test_validate_moa_payload_rejects_non_dict():
assert validate_moa_payload(None)
assert validate_moa_payload([1, 2])
assert validate_moa_payload({"presets": {"p": "not-a-dict"}})
# ── Per-slot max_tokens ────────────────────────────────────────────────────
def test_slot_max_tokens_preserved():
"""A max_tokens field on a reference slot survives normalization."""
cfg = normalize_moa_config(
{
"presets": {
"p": {
"reference_models": [
{"provider": "openrouter", "model": "deepseek/deepseek-v4-pro", "max_tokens": 600},
{"provider": "openai-codex", "model": "gpt-5.5"},
],
"aggregator": {"provider": "openrouter", "model": "anthropic/claude-opus-4.8"},
}
}
}
)
refs = cfg["presets"]["p"]["reference_models"]
assert refs[0] == {"provider": "openrouter", "model": "deepseek/deepseek-v4-pro", "max_tokens": 600}
assert refs[1] == {"provider": "openai-codex", "model": "gpt-5.5"}
def test_slot_max_tokens_coerced_from_string():
"""Hand-edited YAML string '600' coerces to int on a slot."""
cfg = normalize_moa_config(
{
"presets": {
"p": {
"reference_models": [
{"provider": "openrouter", "model": "deepseek/deepseek-v4-pro", "max_tokens": "600"},
],
}
}
}
)
refs = cfg["presets"]["p"]["reference_models"]
assert refs[0]["max_tokens"] == 600
def test_slot_max_tokens_invalid_dropped():
"""Non-positive / non-numeric slot max_tokens is dropped (slot kept)."""
for bad in (0, -5, "abc", "", None):
cfg = normalize_moa_config(
{
"presets": {
"p": {
"reference_models": [
{"provider": "openrouter", "model": "deepseek/deepseek-v4-pro", "max_tokens": bad},
],
}
}
}
)
ref = cfg["presets"]["p"]["reference_models"][0]
assert "max_tokens" not in ref, bad
assert ref["provider"] == "openrouter"
def test_slot_max_tokens_absent_by_default():
"""Slots without max_tokens don't get the field — backward compat."""
cfg = normalize_moa_config(
{
"presets": {
"p": {
"reference_models": [
{"provider": "openrouter", "model": "deepseek/deepseek-v4-pro"},
],
}
}
}
)
ref = cfg["presets"]["p"]["reference_models"][0]
assert "max_tokens" not in ref