From 213008544076c32366027d8c4165ca8950528588 Mon Sep 17 00:00:00 2001 From: Arthur R Longbottom Date: Sat, 21 Feb 2026 10:55:47 -0800 Subject: [PATCH] Fix GroundingDINO bertwarper for transformers 5.x + PyTorch 2.9+ MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit transformers 5.x introduced two breaking changes that affect the bundled GroundingDINO BertModelWarper: 1. `BertModel.get_head_mask()` was removed. Added a standalone `_get_head_mask()` fallback with a `hasattr` check so both transformers 4.x and 5.x work. 2. `BertModel.get_extended_attention_mask()` changed its 3rd positional argument from `device` to `dtype`. The old call-site passes a `torch.device`, which newer transformers interprets as `dtype=torch.device` → TypeError. Replaced with a standalone `_get_extended_attention_mask()` static method that also handles PyTorch 2.9's prohibition of `1.0 - bool_tensor` by explicitly casting to float32 first. Fixes #140 Co-Authored-By: Claude Opus 4.6 --- .../models/GroundingDINO/bertwarper.py | 43 +++++++++++++++++-- 1 file changed, 39 insertions(+), 4 deletions(-) diff --git a/py/local_groundingdino/models/GroundingDINO/bertwarper.py b/py/local_groundingdino/models/GroundingDINO/bertwarper.py index e209a39..2ed4016 100644 --- a/py/local_groundingdino/models/GroundingDINO/bertwarper.py +++ b/py/local_groundingdino/models/GroundingDINO/bertwarper.py @@ -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