Fix IndexError in Triton backend with pipeline parallelism (#30340)
Co-authored-by: Claude Sonnet 4.5 <noreply@anthropic.com> Co-authored-by: Shangming Cai <csmthu@gmail.com>
This commit is contained in:
co-authored by
Claude Sonnet 4.5
Shangming Cai
parent
7af3d000f2
commit
fe52b49827
@@ -205,9 +205,11 @@ class TritonAttnBackend(AttentionBackend):
|
|||||||
self.v_head_dim = model_runner.token_to_kv_pool.get_v_head_dim()
|
self.v_head_dim = model_runner.token_to_kv_pool.get_v_head_dim()
|
||||||
self.swa_v_head_dim = None
|
self.swa_v_head_dim = None
|
||||||
else:
|
else:
|
||||||
self.v_head_dim = model_runner.token_to_kv_pool.get_value_buffer(0).shape[
|
# Use start_layer instead of 0 to handle pipeline parallelism.
|
||||||
-1
|
# In PP, start_layer may be > 0, so layer 0 isn't in this stage's buffer.
|
||||||
]
|
self.v_head_dim = model_runner.token_to_kv_pool.get_value_buffer(
|
||||||
|
model_runner.token_to_kv_pool.start_layer
|
||||||
|
).shape[-1]
|
||||||
self.swa_v_head_dim = None
|
self.swa_v_head_dim = None
|
||||||
self.max_context_len = model_runner.model_config.context_len
|
self.max_context_len = model_runner.model_config.context_len
|
||||||
self.device = model_runner.device
|
self.device = model_runner.device
|
||||||
|
|||||||
@@ -3869,7 +3869,11 @@ class HybridLinearKVPool(KVCache):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def get_v_head_dim(self):
|
def get_v_head_dim(self):
|
||||||
return self.full_kv_pool.get_value_buffer(0).shape[-1]
|
# Use start_layer to handle pipeline parallelism where layer 0
|
||||||
|
# may not be present in this stage's buffer.
|
||||||
|
return self.full_kv_pool.get_value_buffer(self.full_kv_pool.start_layer).shape[
|
||||||
|
-1
|
||||||
|
]
|
||||||
|
|
||||||
def set_mla_kv_buffer(
|
def set_mla_kv_buffer(
|
||||||
self,
|
self,
|
||||||
@@ -4991,4 +4995,6 @@ class MiniMaxSparseKVPool(KVCache):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def get_v_head_dim(self):
|
def get_v_head_dim(self):
|
||||||
return self.main_pool.get_value_buffer(0).shape[-1]
|
# Use start_layer to handle pipeline parallelism where layer 0
|
||||||
|
# may not be present in this stage's buffer.
|
||||||
|
return self.main_pool.get_value_buffer(self.main_pool.start_layer).shape[-1]
|
||||||
|
|||||||
Reference in New Issue
Block a user