Skip to content
Open
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
24 changes: 24 additions & 0 deletions backends/exllamav3/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -435,6 +435,7 @@ async def create(cls, model_directory: pathlib.Path, hf_model: HFModel, **kwargs
cache_mode_default = "FP16"
self.cache_mode = unwrap(kwargs.get("cache_mode"), cache_mode_default)
self.cache = self.create_cache(self.cache_mode, self.model)
self.log_recurrent_slot_cost()

# Draft cache
if self.use_draft_model:
Expand Down Expand Up @@ -704,6 +705,29 @@ def job_max_rq_tokens(self, max_tokens: int) -> Optional[int]:

return self.max_rq_tokens

def log_recurrent_slot_cost(self):
"""
Logs planned recurrent state storage when drafting is enabled.

Use the layer shapes: history layouts differ by architecture, and some
state tensors live on the CPU rather than the GPU.
"""

recurrent_layers = getattr(self.cache, "recurrent_layers", None)
if not recurrent_layers or not self.cache.max_history:
return

num_slots = self.cache.num_slots
total_bytes = sum(layer.storage_size() for layer in recurrent_layers.values())
slot_word = "slot" if num_slots == 1 else "slots"
hint = " Single-user setups can set max_batch_size: 1." if num_slots > 1 else ""
xlogger.info(
f"Recurrent state storage: {num_slots} {slot_word}, "
f"{total_bytes / 1024**2:.0f} MiB total "
f"({total_bytes / num_slots / 1024**2:.0f} MiB per slot), "
f"max_history: {self.cache.max_history}.{hint}"
)

def create_cache(self, raw_cache_mode: str, model: Model):
# Cast exl2 types to exl3
match raw_cache_mode:
Expand Down
9 changes: 7 additions & 2 deletions common/config_models.py
Original file line number Diff line number Diff line change
Expand Up @@ -372,10 +372,15 @@ class ModelConfig(BaseConfigModel):
None,
description=(
"Set the maximum number of generation jobs that can run concurrently\n"
"The default maximum batch size for transformer architectures is 32. Recurrent\n"
"The default maximum batch size for transformer architectures is 128. Recurrent\n"
"models with linear or sliding attention use more VRAM to support larger batches,\n"
"so the default value is reduced to 4. If you do not require concurrency at all, you\n"
"can reduce it further to minimize VRAM overhead."
"can reduce it further to minimize VRAM overhead.\n"
"Recurrent state storage is reserved per slot and depends on the architecture\n"
"and draft length (for example, 728 MiB per slot for a 27B GDN hybrid with 4\n"
"draft tokens). Planned state storage is logged when history is enabled; it\n"
"does not include weights, the paged KV cache, or other runtime buffers.\n"
"Single-user setups serving a recurrent model can set this to 1."
),
ge=1,
)
Expand Down
7 changes: 6 additions & 1 deletion config_sample.yml
Original file line number Diff line number Diff line change
Expand Up @@ -202,10 +202,15 @@ model:
output_chunking: true

# Set the maximum number of generation jobs that can run concurrently
# The default maximum batch size for transformer architectures is 32. Recurrent
# The default maximum batch size for transformer architectures is 128. Recurrent
# models with linear or sliding attention use more VRAM to support larger batches,
# so the default value is reduced to 4. If you do not require concurrency at all, you
# can reduce it further to minimize VRAM overhead.
# Recurrent state storage is reserved per slot and depends on the architecture
# and draft length (for example, 728 MiB per slot for a 27B GDN hybrid with 4
# draft tokens). Planned state storage is logged when history is enabled; it
# does not include weights, the paged KV cache, or other runtime buffers.
# Single-user setups serving a recurrent model can set this to 1.
max_batch_size:

# Tokens between recurrent state checkpoints near the end of the prompt and
Expand Down
4 changes: 1 addition & 3 deletions docs/02.-Server-options.md
Original file line number Diff line number Diff line change
Expand Up @@ -85,11 +85,9 @@ Note: Most of the options here will only apply on initial model load/startup (ep
| cache_mode | String ("FP16") | Cache mode for the model. Specify the pair `k_bits,v_bits` where each is an integer from 2-8 (e.g. `8,8`).<br><br>The legacy values FP16, Q8, Q6 and Q4 are also accepted. |
| cache_size | Int (max_seq_len) | Size of the K/V cache<br><br>Note: If using CFG, the cache size should be 2 * max_seq_len. |
| chunk_size | Int (2048) | Amount of tokens per chunk with ingestion. A lower value reduces VRAM usage at the cost of ingestion speed. |
| max_batch_size | Int (None) | The absolute maximum amount of prompts to process at one time. This value is automatically adjusted based on cache size. |
| max_batch_size | Int (None) | The maximum number of prompts to process at one time. The default is 128 for transformer models and 4 for recurrent models; the effective batch size can be lower depending on cache capacity.<br><br>Recurrent state storage is reserved per slot and depends on the architecture and draft length (for example, 728 MiB per slot for a 27B GDN hybrid with 4 draft tokens). Planned state storage is logged when history is enabled; it excludes weights, the paged KV cache, and other runtime buffers, and can include CPU-side state. Single-user setups serving a recurrent model can set this to 1. |
| recurrent_checkpoint_interval | Int (None) | Tokens between recurrent state checkpoints near the end of the prompt and during generation, for models with recurrent (linear or sliding attention) layers. Must be a multiple of 256. Blank uses the engine's per-architecture default (2048 for most models). |

| recurrent_checkpoint_interval_pp | Int (None) | Tokens between recurrent state checkpoints during prompt ingestion, more than `2 * chunk_size` from the end of the prompt. Must be a multiple of 256; blank uses the engine default of 32768. Recurrent states cannot be rolled back, so editing an earlier part of a cached prompt replays from the last checkpoint before the edit. A denser grid makes that cost proportional to the distance from the edit to the end of the prompt, for one recurrent state of system RAM per checkpoint (bounded by `sysmem_recurrent_cache`) and a few percent slower cold prefill. |

| prompt_template | String (None) | Name of a jinja2 chat template to apply for this model. Must be located in the `templates` directory. |
| vision | Bool (False) | Enable vision support for the provided model (if it exists). |
| sampling | Mapping (None) | Sampler overrides for this model, with the same syntax as the top-level `sampling` section: an optional `override_preset` plus inline sampler overrides. Applied on top of the global section while the model is loaded; put it in the model folder's `tabby_config.yml`. |
Expand Down
92 changes: 92 additions & 0 deletions tests/test_recurrent_slot_cost.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,92 @@
import unittest
from types import SimpleNamespace
from unittest.mock import patch

from exllamav3.modules.gated_delta_net import GDNLayerState
from exllamav3.modules.ple import PLELayerState
from exllamav3.modules.sliding_attn import SWALayerState

from backends.exllamav3.model import ExllamaV3Container


def gdn_states(slots, history):
module = SimpleNamespace(
fdim_qkv=10240,
conv_kernel_size=4,
num_v_heads=48,
k_head_dim=128,
v_head_dim=128,
)
return {i: GDNLayerState(module, slots, history, 0) for i in range(48)}


class RecurrentSlotCostTests(unittest.TestCase):
def log(self, cache):
container = ExllamaV3Container.__new__(ExllamaV3Container)
container.cache = cache
with patch("backends.exllamav3.model.xlogger.info") as info:
container.log_recurrent_slot_cost()
return info

def test_gdn_four_slots(self):
states = gdn_states(4, 4)
self.assertTrue(all(s.recurrent_state.is_meta for s in states.values()))
info = self.log(SimpleNamespace(recurrent_layers=states, num_slots=4, max_history=4))
info.assert_called_once_with(
"Recurrent state storage: 4 slots, 2910 MiB total "
"(728 MiB per slot), max_history: 4. "
"Single-user setups can set max_batch_size: 1."
)

def test_gdn_single_slot_has_no_reduction_hint(self):
info = self.log(
SimpleNamespace(recurrent_layers=gdn_states(1, 4), num_slots=1, max_history=4)
)
info.assert_called_once_with(
"Recurrent state storage: 1 slot, 728 MiB total (728 MiB per slot), max_history: 4."
)

def test_no_history_is_silent(self):
self.log(
SimpleNamespace(recurrent_layers=gdn_states(4, 0), num_slots=4, max_history=0)
).assert_not_called()

def test_no_recurrent_layers_is_silent(self):
for cache in (SimpleNamespace(), SimpleNamespace(recurrent_layers={})):
with self.subTest(cache=cache):
self.log(cache).assert_not_called()

def test_sliding_attention_uses_storage_not_history_times_checkpoint(self):
module = SimpleNamespace(
kv_state_size=4352, sliding_window=4096, num_kv_heads=4, head_dim=128
)
state = SWALayerState(module, 4, 4, 0)
self.assertTrue(state.k_state.is_meta)
self.assertNotEqual(state.storage_size(), 4 * 5 * state.get_checkpoint_size())
info = self.log(SimpleNamespace(recurrent_layers={0: state}, num_slots=4, max_history=4))
info.assert_called_once_with(
"Recurrent state storage: 4 slots, 34 MiB total "
"(8 MiB per slot), max_history: 4. "
"Single-user setups can set max_batch_size: 1."
)

def test_mixed_state_storage_is_not_labeled_as_vram(self):
states = gdn_states(4, 4)
module = SimpleNamespace(
conv_state_len=4,
ple_embedding=SimpleNamespace(context_len=4),
hc_mult=4,
hidden_size=256,
)
states[48] = PLELayerState(module, 4, 4, 0)
info = self.log(SimpleNamespace(recurrent_layers=states, num_slots=4, max_history=4))
info.assert_called_once()
message = info.call_args.args[0]
total_mib = sum(state.storage_size() for state in states.values()) / 1024**2
self.assertIn(f"{total_mib:.0f} MiB total", message)
self.assertNotIn("VRAM", message)
self.assertNotIn("states", message)


if __name__ == "__main__":
unittest.main()
Loading