Fix MiMo-V2 on Blackwell: FA3 fallback and TP-aware audio weight loading (#31343)
This commit is contained in:
@@ -20,17 +20,9 @@ from transformers.modeling_utils import PreTrainedModel
|
||||
from transformers.models.qwen2.configuration_qwen2 import Qwen2Config
|
||||
from transformers.models.qwen2.modeling_qwen2 import Qwen2Model
|
||||
|
||||
from sglang.srt.layers.attention.vision import VisionAttention
|
||||
from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
from sglang.srt.utils import is_cuda
|
||||
|
||||
if is_cuda():
|
||||
from sgl_kernel.flash_attn import flash_attn_varlen_func
|
||||
else:
|
||||
|
||||
def flash_attn_varlen_func(*args, **kwargs):
|
||||
raise RuntimeError("MiMoAudioTokenizer requires CUDA to run.")
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -477,6 +469,22 @@ def get_position_ids(lengths):
|
||||
LAYER_NORM = {"LayerNorm": nn.LayerNorm}
|
||||
|
||||
|
||||
def _audio_rope_applier(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
position_embeddings: Tuple[torch.Tensor, torch.Tensor],
|
||||
x_shape,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
cos, sin = position_embeddings
|
||||
cos = cos.unsqueeze(1)
|
||||
sin = sin.unsqueeze(1)
|
||||
x1_q, x2_q = q[..., : q.shape[-1] // 2], q[..., q.shape[-1] // 2 :]
|
||||
x1_k, x2_k = k[..., : k.shape[-1] // 2], k[..., k.shape[-1] // 2 :]
|
||||
q_embed = q * cos + torch.cat((-x2_q, x1_q), dim=-1) * sin
|
||||
k_embed = k * cos + torch.cat((-x2_k, x1_k), dim=-1) * sin
|
||||
return q_embed, k_embed
|
||||
|
||||
|
||||
class AudioEncoderAttention(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
@@ -492,10 +500,18 @@ class AudioEncoderAttention(nn.Module):
|
||||
self.window_size = window_size
|
||||
self.causal = causal
|
||||
|
||||
self.k_proj = nn.Linear(embed_dim, embed_dim, bias=False)
|
||||
self.v_proj = nn.Linear(embed_dim, embed_dim, bias=True)
|
||||
self.q_proj = nn.Linear(embed_dim, embed_dim, bias=True)
|
||||
self.out_proj = nn.Linear(embed_dim, embed_dim, bias=True)
|
||||
self.attn = VisionAttention(
|
||||
embed_dim=embed_dim,
|
||||
num_heads=num_heads,
|
||||
projection_size=embed_dim,
|
||||
use_qkv_parallel=True,
|
||||
qkv_bias=True,
|
||||
proj_bias=True,
|
||||
flatten_batch=True,
|
||||
window_size=window_size,
|
||||
customized_position_embedding_applier=_audio_rope_applier,
|
||||
prefix="attn",
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -504,51 +520,13 @@ class AudioEncoderAttention(nn.Module):
|
||||
max_seqlen: int,
|
||||
rope_position_embeddings=None,
|
||||
):
|
||||
bsz, _ = hidden_states.size()
|
||||
|
||||
query_states = self.q_proj(hidden_states).view(
|
||||
bsz, self.num_heads, self.head_dim
|
||||
out = self.attn(
|
||||
hidden_states,
|
||||
cu_seqlens=cu_seqlens,
|
||||
position_embeddings=rope_position_embeddings,
|
||||
max_seqlen=max_seqlen,
|
||||
)
|
||||
key_states = self.k_proj(hidden_states).view(bsz, self.num_heads, self.head_dim)
|
||||
value_states = self.v_proj(hidden_states).view(
|
||||
bsz, self.num_heads, self.head_dim
|
||||
)
|
||||
|
||||
if rope_position_embeddings is not None:
|
||||
cos, sin = rope_position_embeddings
|
||||
query_states, key_states = self.apply_rotary_pos_emb(
|
||||
query_states, key_states, cos, sin
|
||||
)
|
||||
|
||||
attn_output = flash_attn_varlen_func(
|
||||
query_states,
|
||||
key_states,
|
||||
value_states,
|
||||
cu_seqlens,
|
||||
cu_seqlens,
|
||||
max_seqlen,
|
||||
max_seqlen,
|
||||
causal=self.causal,
|
||||
window_size=self.window_size,
|
||||
)
|
||||
|
||||
attn_output = attn_output.reshape(bsz, self.embed_dim)
|
||||
attn_output = self.out_proj(attn_output)
|
||||
return attn_output
|
||||
|
||||
@staticmethod
|
||||
def _rotate_half(x):
|
||||
x1 = x[..., : x.shape[-1] // 2]
|
||||
x2 = x[..., x.shape[-1] // 2 :]
|
||||
return torch.cat((-x2, x1), dim=-1)
|
||||
|
||||
@classmethod
|
||||
def apply_rotary_pos_emb(cls, q, k, cos, sin, unsqueeze_dim=1):
|
||||
cos = cos.unsqueeze(unsqueeze_dim)
|
||||
sin = sin.unsqueeze(unsqueeze_dim)
|
||||
q_embed = (q * cos) + (cls._rotate_half(q) * sin)
|
||||
k_embed = (k * cos) + (cls._rotate_half(k) * sin)
|
||||
return q_embed, k_embed
|
||||
return out.squeeze(0)
|
||||
|
||||
|
||||
class AudioEncoderTransformerLayer(nn.Module):
|
||||
@@ -1141,6 +1119,60 @@ class MiMoV2AudioConfig:
|
||||
return config
|
||||
|
||||
|
||||
def _remap_audio_tokenizer_state_dict(state_dict: dict) -> dict:
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
|
||||
tp_size = get_parallel().attn_tp_size
|
||||
tp_rank = get_parallel().attn_tp_rank
|
||||
|
||||
remapped = {}
|
||||
qkv_parts: dict[str, dict] = {}
|
||||
for key, value in state_dict.items():
|
||||
if "self_attn.out_proj" in key:
|
||||
new_key = key.replace("self_attn.out_proj", "self_attn.attn.proj")
|
||||
if tp_size > 1 and key.endswith(".weight"):
|
||||
chunk_size = value.shape[-1] // tp_size
|
||||
value = value[..., tp_rank * chunk_size : (tp_rank + 1) * chunk_size]
|
||||
remapped[new_key] = value
|
||||
elif (
|
||||
"self_attn.q_proj" in key
|
||||
or "self_attn.k_proj" in key
|
||||
or "self_attn.v_proj" in key
|
||||
):
|
||||
suffix = ".weight" if key.endswith(".weight") else ".bias"
|
||||
base = key.rsplit(".", 1)[0]
|
||||
qkv_key = (
|
||||
base.replace("self_attn.q_proj", "self_attn.attn.qkv_proj")
|
||||
.replace("self_attn.k_proj", "self_attn.attn.qkv_proj")
|
||||
.replace("self_attn.v_proj", "self_attn.attn.qkv_proj")
|
||||
+ suffix
|
||||
)
|
||||
if qkv_key not in qkv_parts:
|
||||
qkv_parts[qkv_key] = {}
|
||||
if "q_proj" in key:
|
||||
qkv_parts[qkv_key]["q"] = value
|
||||
elif "k_proj" in key:
|
||||
qkv_parts[qkv_key]["k"] = value
|
||||
elif "v_proj" in key:
|
||||
qkv_parts[qkv_key]["v"] = value
|
||||
else:
|
||||
remapped[key] = value
|
||||
for qkv_key, parts in qkv_parts.items():
|
||||
q = parts.get("q")
|
||||
k = parts.get("k")
|
||||
v = parts.get("v")
|
||||
if q is not None and v is not None:
|
||||
if k is None:
|
||||
k = torch.zeros_like(q)
|
||||
if tp_size > 1:
|
||||
chunk_size = q.shape[0] // tp_size
|
||||
q = q[tp_rank * chunk_size : (tp_rank + 1) * chunk_size]
|
||||
k = k[tp_rank * chunk_size : (tp_rank + 1) * chunk_size]
|
||||
v = v[tp_rank * chunk_size : (tp_rank + 1) * chunk_size]
|
||||
remapped[qkv_key] = torch.cat([q, k, v], dim=0)
|
||||
return remapped
|
||||
|
||||
|
||||
class AudioEncoderMixin:
|
||||
"""LM model mixin that adds MiMo audio encoder components.
|
||||
|
||||
@@ -1262,8 +1294,7 @@ class AudioEncoderMixin:
|
||||
f"No model weights found in {path} "
|
||||
"(expected model.safetensors or pytorch_model.bin)"
|
||||
)
|
||||
# strict=False: upstream ckpt also carries decoder/vocoder weights
|
||||
# that this encoder-only MiMoAudioTokenizer doesn't materialize.
|
||||
state_dict = _remap_audio_tokenizer_state_dict(state_dict)
|
||||
model.load_state_dict(state_dict, strict=False)
|
||||
model = model.to(device=device, dtype=torch.bfloat16)
|
||||
model.eval()
|
||||
|
||||
@@ -1277,6 +1277,25 @@ class MiMoV2ForCausalLM(nn.Module, AudioEncoderMixin):
|
||||
if name.startswith("audio_encoder."):
|
||||
name = name[len("audio_encoder.") :]
|
||||
name = self.remap_audio_weight_name(name)
|
||||
if "input_local_transformer" not in name:
|
||||
name = name.replace("self_attn.out_proj", "self_attn.attn.proj")
|
||||
audio_stacked = False
|
||||
for param_name, weight_name, shard_id in [
|
||||
("self_attn.attn.qkv_proj", "self_attn.q_proj", "q"),
|
||||
("self_attn.attn.qkv_proj", "self_attn.k_proj", "k"),
|
||||
("self_attn.attn.qkv_proj", "self_attn.v_proj", "v"),
|
||||
]:
|
||||
if weight_name in name:
|
||||
name = name.replace(weight_name, param_name)
|
||||
if name not in params_dict:
|
||||
break
|
||||
param = params_dict[name]
|
||||
weight_loader = param.weight_loader
|
||||
weight_loader(param, loaded_weight, shard_id)
|
||||
audio_stacked = True
|
||||
break
|
||||
if audio_stacked:
|
||||
continue
|
||||
if name not in params_dict:
|
||||
logger.warning(
|
||||
f"Audio param {name} not found in params_dict, skipping"
|
||||
|
||||
@@ -18,43 +18,12 @@ from sglang.srt.managers.mm_utils import (
|
||||
from sglang.srt.managers.schedule_batch import MultimodalInputs
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
|
||||
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
||||
from sglang.srt.models import mimo_audio as _mimo_audio_module
|
||||
from sglang.srt.models.mimo import MiMoForCausalLM
|
||||
from sglang.srt.models.mimo_audio import AudioEncoderMixin, MiMoAudioEncoderConfig
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _maybe_override_audio_attn_for_blackwell() -> None:
|
||||
"""Swap mimo_audio.flash_attn_varlen_func to upstream FA2 on GPUs that
|
||||
sgl-kernel's FA3 doesn't support.
|
||||
|
||||
sgl-kernel FA3 only covers sm80/86/89/90 — on Blackwell consumer cards
|
||||
(sm_120 / RTX 50xx) its varlen kernel raises NotImplementedError. ASR is
|
||||
small enough to be deployed on those GPUs, so when FA3 isn't supported
|
||||
we replace the module-level reference with upstream flash-attn (FA2),
|
||||
which works on sm_120. No-op on supported GPUs (FA3 stays).
|
||||
|
||||
MiMo-V2 (the heavy multimodal model) is only deployed on H100/A100, so
|
||||
this override never triggers in its hot path.
|
||||
"""
|
||||
try:
|
||||
from sgl_kernel.flash_attn import is_fa3_supported
|
||||
except ImportError:
|
||||
return
|
||||
if is_fa3_supported():
|
||||
return
|
||||
try:
|
||||
from flash_attn import flash_attn_varlen_func
|
||||
except ImportError as e:
|
||||
raise RuntimeError(
|
||||
"MiMo-V2-ASR audio encoder needs upstream flash-attn on this GPU "
|
||||
"(sgl-kernel FA3 doesn't support sm_120). Install with "
|
||||
"`pip install flash-attn --no-build-isolation`."
|
||||
) from e
|
||||
_mimo_audio_module.flash_attn_varlen_func = flash_attn_varlen_func
|
||||
|
||||
|
||||
MiMoV2ASRConfig = Any
|
||||
|
||||
# Top-level audio sub-module name prefixes (after AUDIO_WEIGHT_REMAP). Loaded
|
||||
@@ -84,7 +53,6 @@ class MiMoV2ASRForCausalLM(MiMoForCausalLM, AudioEncoderMixin):
|
||||
quant_config=None,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
_maybe_override_audio_attn_for_blackwell()
|
||||
super().__init__(config, quant_config=quant_config, prefix=prefix)
|
||||
self.build_audio_encoder(MiMoAudioEncoderConfig(**config.audio_config))
|
||||
|
||||
|
||||
@@ -279,7 +279,7 @@ class MiMoVisionTransformer(nn.Module):
|
||||
num_heads=num_heads,
|
||||
hidden_act=vision_config.hidden_act,
|
||||
norm_layer=norm_layer,
|
||||
attn_implementation="flash_attention_3",
|
||||
attn_implementation=None,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix(f"blocks.{i}", prefix),
|
||||
use_sink=(
|
||||
|
||||
Reference in New Issue
Block a user