From ed361ae7f0ba5495c79c73a4c4f785facea56f84 Mon Sep 17 00:00:00 2001 From: Jimmy Shong <69131491+Jiminator@users.noreply.github.com> Date: Wed, 29 Jul 2026 20:03:00 -0700 Subject: [PATCH] Fix attention backends for models with per-layer head counts (num_attention_heads_per_layer) (#32625) --- python/sglang/srt/configs/model_config.py | 8 ++++++++ .../srt/layers/attention/flashinfer_backend.py | 12 ++++++++++-- python/sglang/srt/layers/attention/triton_backend.py | 6 ++++-- .../attention_methods/dense_attention.py | 3 +++ .../attention_methods/dsa_attention.py | 3 +++ .../attention_methods/dsv4_attention.py | 3 +++ .../attention_methods/dual_chunk_attention.py | 3 +++ .../attention_methods/gdn_attention.py | 3 +++ .../attention_methods/kda_attention.py | 3 +++ .../attention_methods/lightning_attention.py | 3 +++ .../attention_methods/mamba2_attention.py | 3 +++ .../attention_methods/mla_attention.py | 3 +++ 12 files changed, 49 insertions(+), 4 deletions(-) diff --git a/python/sglang/srt/configs/model_config.py b/python/sglang/srt/configs/model_config.py index 5194d8dba..37566da40 100644 --- a/python/sglang/srt/configs/model_config.py +++ b/python/sglang/srt/configs/model_config.py @@ -1093,6 +1093,14 @@ class ModelConfig: # equal to the number of attention heads. return self.hf_text_config.num_attention_heads + def get_max_num_attention_heads(self) -> int: + """Max per-layer query head count; num_attention_heads unless the + model sets num_attention_heads_per_layer.""" + per_layer = getattr(self.hf_text_config, "num_attention_heads_per_layer", None) + if per_layer: + return max(per_layer) + return self.num_attention_heads + def get_num_kv_heads(self, tensor_parallel_size) -> int: """Returns the number of KV heads per GPU.""" total_num_kv_heads = self.get_total_num_kv_heads() diff --git a/python/sglang/srt/layers/attention/flashinfer_backend.py b/python/sglang/srt/layers/attention/flashinfer_backend.py index 14b06e7ce..02bc8e6de 100644 --- a/python/sglang/srt/layers/attention/flashinfer_backend.py +++ b/python/sglang/srt/layers/attention/flashinfer_backend.py @@ -1468,8 +1468,12 @@ class FlashInferAttnBackend(AttentionBackend): class FlashInferIndicesUpdaterDecode: def __init__(self, model_runner: ModelRunner, attn_backend: FlashInferAttnBackend): # Parse Constants + # Plan with the max per-layer head count: FlashInfer bakes num_qo_heads + # into the plan, and layers running more heads than planned are + # silently corrupted. Over-planning is safe. self.num_qo_heads = ( - model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size + model_runner.model_config.get_max_num_attention_heads() + // get_parallel().attn_tp_size ) self.num_kv_heads = model_runner.model_config.get_num_kv_heads( get_parallel().attn_tp_size @@ -1736,8 +1740,12 @@ class FlashInferIndicesUpdaterDecode: class FlashInferIndicesUpdaterPrefill: def __init__(self, model_runner: ModelRunner, attn_backend: FlashInferAttnBackend): # Parse Constants + # Plan with the max per-layer head count: FlashInfer bakes num_qo_heads + # into the plan, and layers running more heads than planned are + # silently corrupted. Over-planning is safe. self.num_qo_heads = ( - model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size + model_runner.model_config.get_max_num_attention_heads() + // get_parallel().attn_tp_size ) self.num_kv_heads = model_runner.model_config.get_num_kv_heads( get_parallel().attn_tp_size diff --git a/python/sglang/srt/layers/attention/triton_backend.py b/python/sglang/srt/layers/attention/triton_backend.py index dc6cbdd71..925ac1acf 100644 --- a/python/sglang/srt/layers/attention/triton_backend.py +++ b/python/sglang/srt/layers/attention/triton_backend.py @@ -178,7 +178,8 @@ class TritonAttnBackend(AttentionBackend): self.dcp_size = get_parallel().attn_dcp_size self.dcp_rank = get_parallel().attn_dcp_rank self.num_head = ( - model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size + model_runner.model_config.get_max_num_attention_heads() + // get_parallel().attn_tp_size ) * self.dcp_size self.num_kv_head = model_runner.model_config.get_num_kv_heads( get_parallel().attn_tp_size @@ -1858,7 +1859,8 @@ class TritonMultiStepDraftBackend: ) self.max_context_len = self.attn_backends[0].max_context_len self.num_head = ( - model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size + model_runner.model_config.get_max_num_attention_heads() + // get_parallel().attn_tp_size ) self.device = model_runner.device # Cached variables for generate_draft_decode_kv_indices diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/dense_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/dense_attention.py index b86d93ffb..fb1ae2702 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/dense_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/dense_attention.py @@ -295,6 +295,9 @@ class TinyModelConfig: assert self.num_attention_heads % tp_size == 0 return self.num_attention_heads // tp_size + def get_max_num_attention_heads(self) -> int: + return self.num_attention_heads + def get_num_kv_heads(self, tp_size: int) -> int: assert self.num_key_value_heads % tp_size == 0 return self.num_key_value_heads // tp_size diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/dsa_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/dsa_attention.py index 7fc5f7cb6..c097add8f 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/dsa_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/dsa_attention.py @@ -263,6 +263,9 @@ class TinyDSAModelConfig: self.hf_text_config = self.hf_config self.linear_attn_registry_result = None + def get_max_num_attention_heads(self) -> int: + return self.num_attention_heads + class DSAMockModelRunner(ModelRunner): def __init__( diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/dsv4_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/dsv4_attention.py index 626b6776c..2ce09fa28 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/dsv4_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/dsv4_attention.py @@ -287,6 +287,9 @@ class TinyDSV4ModelConfig: self.hf_text_config = self.hf_config self.linear_attn_registry_result = None + def get_max_num_attention_heads(self) -> int: + return self.num_attention_heads + class MockDSV4ModelRunner: """Minimal runner exposing what `DeepseekV4AttnBackend.__init__` reads. diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/dual_chunk_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/dual_chunk_attention.py index 7802aef15..a6e44d24d 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/dual_chunk_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/dual_chunk_attention.py @@ -295,6 +295,9 @@ class TinyDualChunkModelConfig: assert self.num_attention_heads % tp_size == 0 return self.num_attention_heads // tp_size + def get_max_num_attention_heads(self) -> int: + return self.num_attention_heads + def get_num_kv_heads(self, tp_size: int) -> int: assert self.num_key_value_heads % tp_size == 0 return self.num_key_value_heads // tp_size diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/gdn_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/gdn_attention.py index 743b4b314..cabc12937 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/gdn_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/gdn_attention.py @@ -188,6 +188,9 @@ class TinyGDNModelConfig: self.hf_text_config = self.hf_config self.linear_attn_registry_result = None + def get_max_num_attention_heads(self) -> int: + return self.num_attention_heads + def get_num_kv_heads(self, tp_size: int) -> int: assert self.num_key_value_heads % tp_size == 0 return self.num_key_value_heads // tp_size diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/kda_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/kda_attention.py index 481e724a6..219e67cb6 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/kda_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/kda_attention.py @@ -193,6 +193,9 @@ class TinyKDAModelConfig: self.hf_text_config = self.hf_config self.linear_attn_registry_result = None + def get_max_num_attention_heads(self) -> int: + return self.num_attention_heads + def get_num_kv_heads(self, tp_size: int) -> int: assert self.num_key_value_heads % tp_size == 0 return self.num_key_value_heads // tp_size diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/lightning_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/lightning_attention.py index bbb0235e9..43e4cdf8b 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/lightning_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/lightning_attention.py @@ -203,6 +203,9 @@ class TinyLightningModelConfig: self.hf_text_config = self.hf_config self.linear_attn_registry_result = None + def get_max_num_attention_heads(self) -> int: + return self.num_attention_heads + def get_num_kv_heads(self, tp_size: int) -> int: assert self.num_key_value_heads % tp_size == 0 return self.num_key_value_heads // tp_size diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/mamba2_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/mamba2_attention.py index d1f91810d..f23c3cd4e 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/mamba2_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/mamba2_attention.py @@ -289,6 +289,9 @@ class TinyMamba2ModelConfig: self.hf_text_config = self.hf_config self.linear_attn_registry_result = None + def get_max_num_attention_heads(self) -> int: + return self.num_attention_heads + def get_num_kv_heads(self, tp_size: int) -> int: assert self.num_key_value_heads % tp_size == 0 return self.num_key_value_heads // tp_size diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/mla_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/mla_attention.py index 3fd845321..e7debefa5 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/mla_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/mla_attention.py @@ -200,6 +200,9 @@ class TinyMLAModelConfig: assert self.num_attention_heads % tp_size == 0 return self.num_attention_heads // tp_size + def get_max_num_attention_heads(self) -> int: + return self.num_attention_heads + def get_num_kv_heads(self, tp_size: int) -> int: return 1