diff --git a/src/mimo_audio_tokenizer/modeling_audio_tokenizer.py b/src/mimo_audio_tokenizer/modeling_audio_tokenizer.py index 1a089a3..e9234c6 100644 --- a/src/mimo_audio_tokenizer/modeling_audio_tokenizer.py +++ b/src/mimo_audio_tokenizer/modeling_audio_tokenizer.py @@ -4,11 +4,16 @@ import numpy as np import torch import torch.nn as nn -from flash_attn import flash_attn_varlen_func from torch.nn import functional as F from transformers.activations import ACT2FN from transformers.modeling_utils import PreTrainedModel +# flash_attn is only available on CUDA - make import optional +try: + from flash_attn import flash_attn_varlen_func +except ImportError: + flash_attn_varlen_func = None + from .configuration_audio_tokenizer import MiMoAudioTokenizerConfig from .modeling_rope_utils import ROPE_INIT_FUNCTIONS, dynamic_rope_update, apply_rotary_pos_emb from .quantization import ResidualVectorQuantizer @@ -270,6 +275,63 @@ def __init__(self, embed_dim, num_heads, window_size=(-1, -1), causal=False): self.causal = causal + def _standard_attention( + self, + query_states: torch.Tensor, + key_states: torch.Tensor, + value_states: torch.Tensor, + seq_len: torch.Tensor, + ): + """Standard PyTorch attention implementation for CPU fallback. + + Input shapes: [total_seq_len, num_heads, head_dim] (packed format) + Output shape: [total_seq_len, num_heads, head_dim] (packed format) + """ + total_seq_len, num_heads, head_dim = query_states.shape + batch_size = seq_len.shape[0] + + # Process each sequence in the batch separately + outputs = [] + offset = 0 + for i, length in enumerate(seq_len): + length_item = length.item() + q = query_states[offset:offset + length_item] # [seq_len, num_heads, head_dim] + k = key_states[offset:offset + length_item] + v = value_states[offset:offset + length_item] + + # Transpose for attention: [num_heads, seq_len, head_dim] + q = q.transpose(0, 1) + k = k.transpose(0, 1) + v = v.transpose(0, 1) + + # Compute attention scores + scale = head_dim ** -0.5 + attn_weights = torch.matmul(q, k.transpose(-2, -1)) * scale # [num_heads, seq_len, seq_len] + + # Apply causal mask if needed + if self.causal: + causal_mask = torch.triu( + torch.ones(length_item, length_item, device=query_states.device), + diagonal=1 + ).bool() + attn_weights = attn_weights.masked_fill( + causal_mask.unsqueeze(0), float('-inf') + ) + + # Softmax and apply to values + attn_weights = F.softmax(attn_weights, dim=-1) + attn_output = torch.matmul(attn_weights, v) # [num_heads, seq_len, head_dim] + + # Transpose back: [seq_len, num_heads, head_dim] + attn_output = attn_output.transpose(0, 1) + outputs.append(attn_output) + + offset += length_item + + # Concatenate all outputs + output = torch.cat(outputs, dim=0) + return output + def forward( self, hidden_states: torch.Tensor, @@ -291,21 +353,29 @@ def forward( query_states = apply_rotary_pos_emb(query_states, cos, sin) key_states = apply_rotary_pos_emb(key_states, cos, sin) - cu_len = F.pad(torch.cumsum(seq_len, dim=0), (1, 0), "constant", 0).to( - torch.int32 - ) - max_seqlen = torch.max(seq_len).to(torch.int32).detach() - attn_output = flash_attn_varlen_func( - query_states, - key_states, - value_states, - cu_len, - cu_len, - max_seqlen, - max_seqlen, - causal=self.causal, - window_size=self.window_size, - ) + # Check if we're on CUDA and flash_attn is available + if hidden_states.device.type == "cuda" and flash_attn_varlen_func is not None: + cu_len = F.pad(torch.cumsum(seq_len, dim=0), (1, 0), "constant", 0).to( + torch.int32 + ) + max_seqlen = torch.max(seq_len).to(torch.int32).detach() + attn_output = flash_attn_varlen_func( + query_states, + key_states, + value_states, + cu_len, + cu_len, + max_seqlen, + max_seqlen, + causal=self.causal, + window_size=self.window_size, + ) + else: + # CPU fallback using standard PyTorch attention + attn_output = self._standard_attention( + query_states, key_states, value_states, seq_len + ) + attn_output = attn_output.reshape(bsz, self.embed_dim) attn_output = self.out_proj(attn_output) return attn_output