KeeeeepGoing's picture
Upload folder using huggingface_hub
1d813fd verified
Raw
History Blame Contribute Delete
8.45 kB
import torch
from typing import Optional, Tuple, Dict, Any, List
class HybridMambaAttentionDynamicCache:
def __init__(self, config, batch_size, dtype=torch.bfloat16, device=None):
self.dtype = dtype
self.layers_block_type = config.layers_block_type
self.device = device
self.has_previous_state = False
self.max_batch_size = batch_size
self.sequence_len_offset = 0
self.batch_size_offset = 0
self.transformer_layers = []
for i in range(len(self.layers_block_type)):
name = self.layers_block_type[i]
layer_type = name.split("_")[0] if "_" in name else name
if layer_type in ("Transformer", "HelixSeekMLA"):
self.transformer_layers.append(i)
num_layers = len(config.layers_block_type)
self.conv_states = [None for _ in range(num_layers)]
self.recurrent_states = [None for _ in range(num_layers)]
self.key_cache = [None for _ in range(num_layers)]
self.value_cache = [None for _ in range(num_layers)]
def update(
self,
key_states: torch.Tensor,
value_states: torch.Tensor,
layer_idx: int,
cache_kwargs: Optional[Dict[str, Any]] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
if self.key_cache[layer_idx] is None:
self.key_cache[layer_idx] = key_states
self.value_cache[layer_idx] = value_states
else:
self.key_cache[layer_idx] = torch.cat([self.key_cache[layer_idx], key_states], dim=2)
self.value_cache[layer_idx] = torch.cat([self.value_cache[layer_idx], value_states], dim=2)
return self.key_cache[layer_idx], self.value_cache[layer_idx]
def reorder_cache(self, beam_idx: torch.LongTensor):
for layer_idx in range(len(self.key_cache)):
if self.key_cache[layer_idx] is not None:
device = self.key_cache[layer_idx].device
beam_idx_device = beam_idx.to(device)
self.key_cache[layer_idx] = self.key_cache[layer_idx].index_select(0, beam_idx_device)
self.value_cache[layer_idx] = self.value_cache[layer_idx].index_select(0, beam_idx_device)
if self.conv_states[layer_idx] is not None:
device = self.conv_states[layer_idx][0].device
beam_idx_device = beam_idx.to(device)
q_conv, k_conv, v_conv = self.conv_states[layer_idx]
self.conv_states[layer_idx] = (
q_conv.index_select(0, beam_idx_device),
k_conv.index_select(0, beam_idx_device),
v_conv.index_select(0, beam_idx_device),
)
if self.recurrent_states[layer_idx] is not None:
device = self.recurrent_states[layer_idx].device
beam_idx_device = beam_idx.to(device)
self.recurrent_states[layer_idx] = self.recurrent_states[layer_idx].index_select(0, beam_idx_device)
def get_seq_length(self, layer_idx: Optional[int] = 0) -> int:
if self.transformer_layers:
if layer_idx is None or layer_idx not in self.transformer_layers:
layer_idx = self.transformer_layers[0]
if layer_idx >= len(self.key_cache) or self.key_cache[layer_idx] is None:
return 0
return self.key_cache[layer_idx].shape[-2]
for states in (self.conv_states, self.recurrent_states):
for s in states:
if s is not None:
if isinstance(s, tuple):
s = s[0]
return s.shape[1] if s.dim() >= 2 else s.shape[0]
return 0
def get_max_length(self) -> Optional[int]:
return None
def __len__(self) -> int:
return len(self.layers_block_type)
def to_legacy_cache(self):
raise NotImplementedError("HybridMambaAttentionDynamicCache does not have a legacy cache equivalent.")
@classmethod
def from_legacy_cache(cls, past_key_values: Optional[Tuple[Tuple[torch.FloatTensor]]] = None):
raise NotImplementedError("HybridMambaAttentionDynamicCache does not have a legacy cache equivalent.")
class HelixSeekDeltaDynamicCache:
is_compileable = False
def __init__(self, config, batch_size=1, dtype=torch.bfloat16, device=None):
self.config = config
self.dtype = dtype
self.device = device
self.batch_size = batch_size
self._has_previous_state = False
num_layers = len(config.layers_block_type)
layer_types = []
for i in range(num_layers):
layer_name = config.layers_block_type[i] if hasattr(config, 'layers_block_type') else ""
if layer_name.startswith("HelixSeekDelta"):
layer_types.append("linear_attention")
else:
layer_types.append("full_attention")
self.layer_types = layer_types
self.transformer_layers = [
i for i in range(num_layers) if self.layer_types[i] == "full_attention"
]
linear_layers = [i for i in range(num_layers) if self.layer_types[i] == "linear_attention"]
self.last_linear_layer = linear_layers[-1] if linear_layers else -1
self.conv_states = [None for _ in range(num_layers)]
self.recurrent_states = [None for _ in range(num_layers)]
self.key_cache = [None for _ in range(num_layers)]
self.value_cache = [None for _ in range(num_layers)]
def __len__(self):
return len(self.layer_types)
def update(
self,
key_states: torch.Tensor,
value_states: torch.Tensor,
layer_idx: int,
cache_kwargs: Optional[Dict[str, Any]] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
self.key_cache[layer_idx] = key_states if self.key_cache[layer_idx] is None else torch.cat([self.key_cache[layer_idx], key_states], dim=2)
self.value_cache[layer_idx] = value_states if self.value_cache[layer_idx] is None else torch.cat([self.value_cache[layer_idx], value_states], dim=2)
return self.key_cache[layer_idx], self.value_cache[layer_idx]
def reorder_cache(self, beam_idx: torch.LongTensor):
for layer_idx in range(len(self.key_cache)):
if self.key_cache[layer_idx] is not None:
device = self.key_cache[layer_idx].device
beam_idx_device = beam_idx.to(device)
self.key_cache[layer_idx] = self.key_cache[layer_idx].index_select(0, beam_idx_device)
self.value_cache[layer_idx] = self.value_cache[layer_idx].index_select(0, beam_idx_device)
if self.conv_states[layer_idx] is not None:
device = self.conv_states[layer_idx][0].device
beam_idx_device = beam_idx.to(device)
q_conv, k_conv, v_conv = self.conv_states[layer_idx]
self.conv_states[layer_idx] = (
q_conv.index_select(0, beam_idx_device),
k_conv.index_select(0, beam_idx_device),
v_conv.index_select(0, beam_idx_device),
)
if self.recurrent_states[layer_idx] is not None:
device = self.recurrent_states[layer_idx].device
beam_idx_device = beam_idx.to(device)
self.recurrent_states[layer_idx] = self.recurrent_states[layer_idx].index_select(0, beam_idx_device)
def get_seq_length(self, layer_idx: Optional[int] = 0) -> int:
if layer_idx is None:
layer_idx = 0
if not self.transformer_layers:
layer_idx = 0
elif layer_idx not in self.transformer_layers:
layer_idx = self.transformer_layers[0]
if layer_idx is None:
return 0
if len(self.key_cache) <= layer_idx or self.key_cache[layer_idx] is None:
return 0
return self.key_cache[layer_idx].shape[-2]
def get_mask_sizes(self, cache_position: torch.Tensor, layer_idx: int) -> Tuple[int, int]:
kv_offset = 0
query_length = cache_position.shape[0]
past_seen_tokens = self.get_seq_length(layer_idx)
kv_length = query_length + past_seen_tokens
return kv_length, kv_offset
@property
def has_previous_state(self):
if self.last_linear_layer == -1:
return False
return self.conv_states[self.last_linear_layer] is not None