Merge pull request #142 from artokun/fix/transformers5-bertwarper-compat

Fix GroundingDINO bertwarper for transformers 5.x + PyTorch 2.9+
This commit is contained in:
chflame163
2026-02-22 18:32:26 +08:00
committed by GitHub
@@ -10,6 +10,19 @@ from torch import nn
from transformers.modeling_outputs import BaseModelOutputWithPoolingAndCrossAttentions
def _get_head_mask(num_hidden_layers, head_mask=None):
"""Compatibility shim -- transformers 5.x removed get_head_mask from BertModel."""
if head_mask is not None:
if head_mask.dim() == 1:
head_mask = head_mask.unsqueeze(0).unsqueeze(0).unsqueeze(-1).unsqueeze(-1)
head_mask = head_mask.expand(num_hidden_layers, -1, -1, -1, -1)
elif head_mask.dim() == 2:
head_mask = head_mask.unsqueeze(1).unsqueeze(-1).unsqueeze(-1)
else:
head_mask = [None] * num_hidden_layers
return head_mask
class BertModelWarper(nn.Module):
def __init__(self, bert_model):
super().__init__()
@@ -20,9 +33,31 @@ class BertModelWarper(nn.Module):
self.encoder = bert_model.encoder
self.pooler = bert_model.pooler
self.get_extended_attention_mask = bert_model.get_extended_attention_mask
self.invert_attention_mask = bert_model.invert_attention_mask
self.get_head_mask = bert_model.get_head_mask
# transformers 5.x removed get_head_mask from BertModel; use fallback
if hasattr(bert_model, 'get_head_mask'):
self.get_head_mask = bert_model.get_head_mask
else:
self.get_head_mask = lambda head_mask, num_layers: _get_head_mask(num_layers, head_mask)
@staticmethod
def _get_extended_attention_mask(attention_mask, input_shape, device_or_dtype=None):
"""Standalone extended attention mask compatible with all transformers versions.
transformers 5.x changed the 3rd positional arg of get_extended_attention_mask
from device to dtype, causing a TypeError when the old call-site passes a device.
Additionally, newer PyTorch forbids 1.0 - bool_tensor.
This method avoids both issues.
"""
if attention_mask.dim() == 3:
extended = attention_mask[:, None, :, :]
elif attention_mask.dim() == 2:
extended = attention_mask[:, None, None, :]
else:
raise ValueError(f"Wrong shape for attention_mask (shape {attention_mask.shape})")
extended = extended.to(dtype=torch.float32)
extended = (1.0 - extended) * torch.finfo(torch.float32).min
return extended
def forward(
self,
@@ -102,8 +137,8 @@ class BertModelWarper(nn.Module):
# We can provide a self-attention mask of dimensions [batch_size, from_seq_length, to_seq_length]
# ourselves in which case we just need to make it broadcastable to all heads.
extended_attention_mask: torch.Tensor = self.get_extended_attention_mask(
attention_mask, input_shape, device
extended_attention_mask: torch.Tensor = self._get_extended_attention_mask(
attention_mask, input_shape
)
# If a 2D or 3D attention mask is provided for the cross-attention