"""Keep ModernBERT's ROCm efficient-SDPA inputs contiguous. The native 4.57 ModernBERT readout leaves Q/K transposed and V as a strided view of a packed QKV projection. ROCm's efficient backend can miscompute that layout. This model-local guard changes storage only; CPU and forced-MATH calls continue through the original attention implementation. """ from __future__ import annotations from types import MethodType import torch from torch.nn import functional from transformers.models.modernbert.modeling_modernbert import ( ModernBertAttention, ModernBertModel, apply_rotary_pos_emb, ) LAYOUT_POLICY = "rocm_efficient_contiguous_qkv" def contiguous_sdpa_forward( attention, hidden_states, *, attention_mask, sliding_window_mask, position_ids, ): """Native SDPA arithmetic with contiguous post-RoPE Q/K/V storage.""" batch_size = hidden_states.shape[0] qkv = attention.Wqkv(hidden_states).view( batch_size, -1, 3, attention.num_heads, attention.head_dim ) cos, sin = attention.rotary_emb(qkv, position_ids=position_ids) query, key, value = qkv.transpose(3, 1).unbind(dim=2) query, key = apply_rotary_pos_emb(query, key, cos, sin) mask = ( attention_mask if attention.local_attention == (-1, -1) else sliding_window_mask ) output = functional.scaled_dot_product_attention( query.contiguous(), key.contiguous(), value.contiguous(), attn_mask=mask, dropout_p=attention.attention_dropout if attention.training else 0.0, ) output = ( output.transpose(1, 2) .contiguous() .view(batch_size, -1, attention.all_head_size) ) return (attention.out_drop(attention.Wo(output)),) def _guarded_forward(self, hidden_states, output_attentions=False, **kwargs): if ( torch.version.hip is None or hidden_states.device.type != "cuda" or not torch.backends.cuda.mem_efficient_sdp_enabled() or self.config._attn_implementation != "sdpa" or output_attentions ): return self._vela_original_attention_forward( hidden_states, output_attentions=output_attentions, **kwargs ) return contiguous_sdpa_forward( self, hidden_states, attention_mask=kwargs["attention_mask"], sliding_window_mask=kwargs["sliding_window_mask"], position_ids=kwargs["position_ids"], ) def install_rocm_sdpa_layout_guard(encoder): """Install once on this model's attention modules; no global monkey patch.""" if not isinstance(encoder, ModernBertModel): raise ValueError("The SDPA layout guard requires native ModernBERT") for layer in encoder.layers: attention = layer.attn if not isinstance(attention, ModernBertAttention): raise ValueError("Unexpected attention module in native ModernBERT") if getattr(attention, "_vela_sdpa_layout_policy", None) == LAYOUT_POLICY: continue if hasattr(attention, "_vela_original_attention_forward"): raise ValueError("An unknown attention guard is already installed") attention._vela_original_attention_forward = attention.forward attention.forward = MethodType(_guarded_forward, attention) attention._vela_sdpa_layout_policy = LAYOUT_POLICY return encoder