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.configuration_qwen2 import Qwen2Config
|
||||||
from transformers.models.qwen2.modeling_qwen2 import Qwen2Model
|
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.layers.quantization.base_config import QuantizationConfig
|
||||||
from sglang.srt.runtime_context import get_server_args
|
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__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -477,6 +469,22 @@ def get_position_ids(lengths):
|
|||||||
LAYER_NORM = {"LayerNorm": nn.LayerNorm}
|
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):
|
class AudioEncoderAttention(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -492,10 +500,18 @@ class AudioEncoderAttention(nn.Module):
|
|||||||
self.window_size = window_size
|
self.window_size = window_size
|
||||||
self.causal = causal
|
self.causal = causal
|
||||||
|
|
||||||
self.k_proj = nn.Linear(embed_dim, embed_dim, bias=False)
|
self.attn = VisionAttention(
|
||||||
self.v_proj = nn.Linear(embed_dim, embed_dim, bias=True)
|
embed_dim=embed_dim,
|
||||||
self.q_proj = nn.Linear(embed_dim, embed_dim, bias=True)
|
num_heads=num_heads,
|
||||||
self.out_proj = nn.Linear(embed_dim, embed_dim, bias=True)
|
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(
|
def forward(
|
||||||
self,
|
self,
|
||||||
@@ -504,51 +520,13 @@ class AudioEncoderAttention(nn.Module):
|
|||||||
max_seqlen: int,
|
max_seqlen: int,
|
||||||
rope_position_embeddings=None,
|
rope_position_embeddings=None,
|
||||||
):
|
):
|
||||||
bsz, _ = hidden_states.size()
|
out = self.attn(
|
||||||
|
hidden_states,
|
||||||
query_states = self.q_proj(hidden_states).view(
|
cu_seqlens=cu_seqlens,
|
||||||
bsz, self.num_heads, self.head_dim
|
position_embeddings=rope_position_embeddings,
|
||||||
|
max_seqlen=max_seqlen,
|
||||||
)
|
)
|
||||||
key_states = self.k_proj(hidden_states).view(bsz, self.num_heads, self.head_dim)
|
return out.squeeze(0)
|
||||||
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
|
|
||||||
|
|
||||||
|
|
||||||
class AudioEncoderTransformerLayer(nn.Module):
|
class AudioEncoderTransformerLayer(nn.Module):
|
||||||
@@ -1141,6 +1119,60 @@ class MiMoV2AudioConfig:
|
|||||||
return config
|
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:
|
class AudioEncoderMixin:
|
||||||
"""LM model mixin that adds MiMo audio encoder components.
|
"""LM model mixin that adds MiMo audio encoder components.
|
||||||
|
|
||||||
@@ -1262,8 +1294,7 @@ class AudioEncoderMixin:
|
|||||||
f"No model weights found in {path} "
|
f"No model weights found in {path} "
|
||||||
"(expected model.safetensors or pytorch_model.bin)"
|
"(expected model.safetensors or pytorch_model.bin)"
|
||||||
)
|
)
|
||||||
# strict=False: upstream ckpt also carries decoder/vocoder weights
|
state_dict = _remap_audio_tokenizer_state_dict(state_dict)
|
||||||
# that this encoder-only MiMoAudioTokenizer doesn't materialize.
|
|
||||||
model.load_state_dict(state_dict, strict=False)
|
model.load_state_dict(state_dict, strict=False)
|
||||||
model = model.to(device=device, dtype=torch.bfloat16)
|
model = model.to(device=device, dtype=torch.bfloat16)
|
||||||
model.eval()
|
model.eval()
|
||||||
|
|||||||
@@ -1277,6 +1277,25 @@ class MiMoV2ForCausalLM(nn.Module, AudioEncoderMixin):
|
|||||||
if name.startswith("audio_encoder."):
|
if name.startswith("audio_encoder."):
|
||||||
name = name[len("audio_encoder.") :]
|
name = name[len("audio_encoder.") :]
|
||||||
name = self.remap_audio_weight_name(name)
|
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:
|
if name not in params_dict:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
f"Audio param {name} not found in params_dict, skipping"
|
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.managers.schedule_batch import MultimodalInputs
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
|
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.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 import MiMoForCausalLM
|
||||||
from sglang.srt.models.mimo_audio import AudioEncoderMixin, MiMoAudioEncoderConfig
|
from sglang.srt.models.mimo_audio import AudioEncoderMixin, MiMoAudioEncoderConfig
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
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
|
MiMoV2ASRConfig = Any
|
||||||
|
|
||||||
# Top-level audio sub-module name prefixes (after AUDIO_WEIGHT_REMAP). Loaded
|
# Top-level audio sub-module name prefixes (after AUDIO_WEIGHT_REMAP). Loaded
|
||||||
@@ -84,7 +53,6 @@ class MiMoV2ASRForCausalLM(MiMoForCausalLM, AudioEncoderMixin):
|
|||||||
quant_config=None,
|
quant_config=None,
|
||||||
prefix: str = "",
|
prefix: str = "",
|
||||||
) -> None:
|
) -> None:
|
||||||
_maybe_override_audio_attn_for_blackwell()
|
|
||||||
super().__init__(config, quant_config=quant_config, prefix=prefix)
|
super().__init__(config, quant_config=quant_config, prefix=prefix)
|
||||||
self.build_audio_encoder(MiMoAudioEncoderConfig(**config.audio_config))
|
self.build_audio_encoder(MiMoAudioEncoderConfig(**config.audio_config))
|
||||||
|
|
||||||
|
|||||||
@@ -279,7 +279,7 @@ class MiMoVisionTransformer(nn.Module):
|
|||||||
num_heads=num_heads,
|
num_heads=num_heads,
|
||||||
hidden_act=vision_config.hidden_act,
|
hidden_act=vision_config.hidden_act,
|
||||||
norm_layer=norm_layer,
|
norm_layer=norm_layer,
|
||||||
attn_implementation="flash_attention_3",
|
attn_implementation=None,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix(f"blocks.{i}", prefix),
|
prefix=add_prefix(f"blocks.{i}", prefix),
|
||||||
use_sink=(
|
use_sink=(
|
||||||
|
|||||||
Reference in New Issue
Block a user