Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
102 changes: 86 additions & 16 deletions src/mimo_audio_tokenizer/modeling_audio_tokenizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand All @@ -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
Expand Down