[NPU] Fix TypeError in get_state_buf_infos when index_head_dim is None on MLA (#25383)

This commit is contained in:
Kurkur
2026-05-19 09:09:11 +08:00
committed by GitHub
parent d028697d17
commit d90bc65e30
@@ -354,6 +354,8 @@ class NPUMLATokenToKVPool(MLATokenToKVPool):
)
def get_state_buf_infos(self):
if self.index_head_dim is None:
return [], [], []
data_ptrs = [self.index_k_buffer[i].data_ptr() for i in range(self.layer_num)]
data_lens = [self.index_k_buffer[i].nbytes for i in range(self.layer_num)]
item_lens = [self.index_k_buffer[i][0].nbytes for i in range(self.layer_num)]