modeling_florence2.py
124.5 KB · 2846 lines · python Raw
1 # coding=utf-8
2 # Copyright 2024 Microsoft and the HuggingFace Inc. team. All rights reserved.
3 #
4 # Licensed under the Apache License, Version 2.0 (the "License");
5 # you may not use this file except in compliance with the License.
6 # You may obtain a copy of the License at
7 #
8 # http://www.apache.org/licenses/LICENSE-2.0
9 #
10 # Unless required by applicable law or agreed to in writing, software
11 # distributed under the License is distributed on an "AS IS" BASIS,
12 # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13 # See the License for the specific language governing permissions and
14 # limitations under the License.
15
16 """ PyTorch Florence-2 model."""
17 from dataclasses import dataclass
18 from typing import List, Optional, Tuple, Union
19
20 import math
21 import torch
22 import torch.utils.checkpoint
23 from torch import nn
24 import torch.nn.functional as F
25 import torch.utils.checkpoint as checkpoint
26 from torch.nn import CrossEntropyLoss
27 from collections import OrderedDict
28 from einops import rearrange
29 from timm.layers import DropPath, trunc_normal_
30
31 from transformers.modeling_utils import PreTrainedModel
32 from transformers.generation.utils import GenerationMixin
33 from transformers.utils import (
34 ModelOutput,
35 add_start_docstrings,
36 add_start_docstrings_to_model_forward,
37 is_flash_attn_2_available,
38 logging,
39 replace_return_docstrings,
40 is_flash_attn_2_available,
41 is_flash_attn_greater_or_equal_2_10,
42 )
43 from .configuration_florence2 import Florence2Config
44 from .configuration_florence2 import Florence2LanguageConfig
45 from .configuration_florence2 import Florence2VisionConfig
46
47
48 from transformers.activations import ACT2FN
49 from transformers.modeling_attn_mask_utils import (
50 _prepare_4d_attention_mask,
51 _prepare_4d_attention_mask_for_sdpa,
52 _prepare_4d_causal_attention_mask,
53 _prepare_4d_causal_attention_mask_for_sdpa,
54 )
55 from transformers.modeling_outputs import (
56 BaseModelOutput,
57 BaseModelOutputWithPastAndCrossAttentions,
58 Seq2SeqLMOutput,
59 Seq2SeqModelOutput,
60 )
61
62
63 if is_flash_attn_2_available():
64 from flash_attn.bert_padding import index_first_axis, pad_input, unpad_input # noqa
65
66 logger = logging.get_logger(__name__)
67
68 _CONFIG_FOR_DOC = "Florence2Config"
69
70 class LearnedAbsolutePositionEmbedding2D(nn.Module):
71 """
72 This module learns positional embeddings up to a fixed maximum size.
73 """
74
75 def __init__(self, embedding_dim=256, num_pos=50):
76 super().__init__()
77 self.row_embeddings = nn.Embedding(num_pos, embedding_dim // 2)
78 self.column_embeddings = nn.Embedding(num_pos, embedding_dim - (embedding_dim // 2))
79
80 def forward(self, pixel_values):
81 """
82 pixel_values: (batch_size, height, width, num_channels)
83 returns: (batch_size, height, width, embedding_dim * 2)
84 """
85 if len(pixel_values.shape) != 4:
86 raise ValueError('pixel_values must be a 4D tensor')
87 height, width = pixel_values.shape[1:3]
88 width_values = torch.arange(width, device=pixel_values.device)
89 height_values = torch.arange(height, device=pixel_values.device)
90 x_emb = self.column_embeddings(width_values)
91 y_emb = self.row_embeddings(height_values)
92 # (height, width, embedding_dim * 2)
93 pos = torch.cat([x_emb.unsqueeze(0).repeat(height, 1, 1), y_emb.unsqueeze(1).repeat(1, width, 1)], dim=-1)
94 # (embedding_dim * 2, height, width)
95 pos = pos.permute(2, 0, 1)
96 pos = pos.unsqueeze(0)
97 # (batch_size, embedding_dim * 2, height, width)
98 pos = pos.repeat(pixel_values.shape[0], 1, 1, 1)
99 # (batch_size, height, width, embedding_dim * 2)
100 pos = pos.permute(0, 2, 3, 1)
101 return pos
102
103 class PositionalEmbeddingCosine1D(nn.Module):
104 """
105 This class implements a very simple positional encoding. It follows closely
106 the encoder from the link below:
107 https://pytorch.org/tutorials/beginner/translation_transformer.html
108
109 Args:
110 embed_dim: The dimension of the embeddings.
111 dropout_prob: The dropout probability.
112 max_seq_len: The maximum length to precompute the positional encodings.
113 """
114 def __init__(
115 self,
116 embed_dim: int = 512,
117 max_seq_len: int = 1024) -> None:
118 super(PositionalEmbeddingCosine1D, self).__init__()
119 self.embed_dim = embed_dim
120 self.max_seq_len = max_seq_len
121 # Generate the sinusoidal arrays.
122 factor = math.log(10000)
123 denominator = torch.exp(
124 -factor * torch.arange(0, self.embed_dim, 2) / self.embed_dim)
125 # Matrix where rows correspond to a positional embedding as a function
126 # of the position index (i.e., the row index).
127 frequencies = \
128 torch.arange(0, self.max_seq_len) \
129 .reshape(self.max_seq_len, 1) * denominator
130 pos_idx_to_embed = torch.zeros((self.max_seq_len, self.embed_dim))
131 # Populate uneven entries.
132 pos_idx_to_embed[:, 0::2] = torch.sin(frequencies)
133 pos_idx_to_embed[:, 1::2] = torch.cos(frequencies)
134 # Save the positional embeddings in a constant buffer.
135 self.register_buffer("pos_idx_to_embed", pos_idx_to_embed)
136
137 def forward(self, seq_embeds: torch.Tensor) -> torch.Tensor:
138 """
139 Args:
140 seq_embeds: The sequence embeddings in order. Allowed size:
141 1. [T, D], where T is the length of the sequence, and D is the
142 frame embedding dimension.
143 2. [B, T, D], where B is the batch size and T and D are the
144 same as above.
145
146 Returns a tensor of with the same dimensions as the input: i.e.,
147 [1, T, D] or [T, D].
148 """
149 shape_len = len(seq_embeds.shape)
150 assert 2 <= shape_len <= 3
151 len_seq = seq_embeds.size(-2)
152 assert len_seq <= self.max_seq_len
153 pos_embeds = self.pos_idx_to_embed[0:seq_embeds.size(-2), :]
154 # Adapt pre-computed positional embeddings to the input.
155 if shape_len == 3:
156 pos_embeds = pos_embeds.view(
157 (1, pos_embeds.size(0), pos_embeds.size(1)))
158 return pos_embeds
159
160
161 class LearnedAbsolutePositionEmbedding1D(nn.Module):
162 """
163 Learnable absolute positional embeddings for 1D sequences.
164
165 Args:
166 embed_dim: The dimension of the embeddings.
167 max_seq_len: The maximum length to precompute the positional encodings.
168 """
169 def __init__(
170 self,
171 embedding_dim: int = 512,
172 num_pos: int = 1024) -> None:
173 super(LearnedAbsolutePositionEmbedding1D, self).__init__()
174 self.embeddings = nn.Embedding(num_pos, embedding_dim)
175 self.num_pos = num_pos
176
177 def forward(self, seq_embeds: torch.Tensor) -> torch.Tensor:
178 """
179 Args:
180 seq_embeds: The sequence embeddings in order. Allowed size:
181 1. [T, D], where T is the length of the sequence, and D is the
182 frame embedding dimension.
183 2. [B, T, D], where B is the batch size and T and D are the
184 same as above.
185
186 Returns a tensor of with the same dimensions as the input: i.e.,
187 [1, T, D] or [T, D].
188 """
189 shape_len = len(seq_embeds.shape)
190 assert 2 <= shape_len <= 3
191 len_seq = seq_embeds.size(-2)
192 assert len_seq <= self.num_pos
193 # [T, D]
194 pos_embeds = self.embeddings(torch.arange(len_seq).to(seq_embeds.device))
195 # Adapt pre-computed positional embeddings to the input.
196 if shape_len == 3:
197 pos_embeds = pos_embeds.view(
198 (1, pos_embeds.size(0), pos_embeds.size(1)))
199 return pos_embeds
200
201
202
203 class MySequential(nn.Sequential):
204 def forward(self, *inputs):
205 for module in self._modules.values():
206 if type(inputs) == tuple:
207 inputs = module(*inputs)
208 else:
209 inputs = module(inputs)
210 return inputs
211
212
213 class PreNorm(nn.Module):
214 def __init__(self, norm, fn, drop_path=None):
215 super().__init__()
216 self.norm = norm
217 self.fn = fn
218 self.drop_path = drop_path
219
220 def forward(self, x, *args, **kwargs):
221 shortcut = x
222 if self.norm != None:
223 x, size = self.fn(self.norm(x), *args, **kwargs)
224 else:
225 x, size = self.fn(x, *args, **kwargs)
226
227 if self.drop_path:
228 x = self.drop_path(x)
229
230 x = shortcut + x
231
232 return x, size
233
234
235 class Mlp(nn.Module):
236 def __init__(
237 self,
238 in_features,
239 hidden_features=None,
240 out_features=None,
241 act_layer=nn.GELU,
242 ):
243 super().__init__()
244 out_features = out_features or in_features
245 hidden_features = hidden_features or in_features
246 self.net = nn.Sequential(OrderedDict([
247 ("fc1", nn.Linear(in_features, hidden_features)),
248 ("act", act_layer()),
249 ("fc2", nn.Linear(hidden_features, out_features))
250 ]))
251
252 def forward(self, x, size):
253 return self.net(x), size
254
255
256 class DepthWiseConv2d(nn.Module):
257 def __init__(
258 self,
259 dim_in,
260 kernel_size,
261 padding,
262 stride,
263 bias=True,
264 ):
265 super().__init__()
266 self.dw = nn.Conv2d(
267 dim_in, dim_in,
268 kernel_size=kernel_size,
269 padding=padding,
270 groups=dim_in,
271 stride=stride,
272 bias=bias
273 )
274
275 def forward(self, x, size):
276 B, N, C = x.shape
277 H, W = size
278 assert N == H * W
279
280 x = self.dw(x.transpose(1, 2).view(B, C, H, W))
281 size = (x.size(-2), x.size(-1))
282 x = x.flatten(2).transpose(1, 2)
283 return x, size
284
285
286 class ConvEmbed(nn.Module):
287 """ Image to Patch Embedding
288 """
289
290 def __init__(
291 self,
292 patch_size=7,
293 in_chans=3,
294 embed_dim=64,
295 stride=4,
296 padding=2,
297 norm_layer=None,
298 pre_norm=True
299 ):
300 super().__init__()
301 self.patch_size = patch_size
302
303 self.proj = nn.Conv2d(
304 in_chans, embed_dim,
305 kernel_size=patch_size,
306 stride=stride,
307 padding=padding
308 )
309
310 dim_norm = in_chans if pre_norm else embed_dim
311 self.norm = norm_layer(dim_norm) if norm_layer else None
312
313 self.pre_norm = pre_norm
314
315 def forward(self, x, size):
316 H, W = size
317 if len(x.size()) == 3:
318 if self.norm and self.pre_norm:
319 x = self.norm(x)
320 x = rearrange(
321 x, 'b (h w) c -> b c h w',
322 h=H, w=W
323 )
324
325 x = self.proj(x)
326
327 _, _, H, W = x.shape
328 x = rearrange(x, 'b c h w -> b (h w) c')
329 if self.norm and not self.pre_norm:
330 x = self.norm(x)
331
332 return x, (H, W)
333
334
335 class ChannelAttention(nn.Module):
336
337 def __init__(self, dim, groups=8, qkv_bias=True):
338 super().__init__()
339
340 self.groups = groups
341 self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
342 self.proj = nn.Linear(dim, dim)
343
344 def forward(self, x, size):
345 B, N, C = x.shape
346
347 qkv = self.qkv(x).reshape(B, N, 3, self.groups, C // self.groups).permute(2, 0, 3, 1, 4)
348 q, k, v = qkv[0], qkv[1], qkv[2]
349
350 q = q * (float(N) ** -0.5)
351 attention = q.transpose(-1, -2) @ k
352 attention = attention.softmax(dim=-1)
353 x = (attention @ v.transpose(-1, -2)).transpose(-1, -2)
354 x = x.transpose(1, 2).reshape(B, N, C)
355 x = self.proj(x)
356 return x, size
357
358
359 class ChannelBlock(nn.Module):
360
361 def __init__(self, dim, groups, mlp_ratio=4., qkv_bias=True,
362 drop_path_rate=0., act_layer=nn.GELU, norm_layer=nn.LayerNorm,
363 conv_at_attn=True, conv_at_ffn=True):
364 super().__init__()
365
366 drop_path = DropPath(drop_path_rate) if drop_path_rate > 0. else nn.Identity()
367
368 self.conv1 = PreNorm(None, DepthWiseConv2d(dim, 3, 1, 1)) if conv_at_attn else None
369 self.channel_attn = PreNorm(
370 norm_layer(dim),
371 ChannelAttention(dim, groups=groups, qkv_bias=qkv_bias),
372 drop_path
373 )
374 self.conv2 = PreNorm(None, DepthWiseConv2d(dim, 3, 1, 1)) if conv_at_ffn else None
375 self.ffn = PreNorm(
376 norm_layer(dim),
377 Mlp(in_features=dim, hidden_features=int(dim*mlp_ratio), act_layer=act_layer),
378 drop_path
379 )
380
381 def forward(self, x, size):
382 if self.conv1:
383 x, size = self.conv1(x, size)
384 x, size = self.channel_attn(x, size)
385
386 if self.conv2:
387 x, size = self.conv2(x, size)
388 x, size = self.ffn(x, size)
389
390 return x, size
391
392
393 def window_partition(x, window_size: int):
394 B, H, W, C = x.shape
395 x = x.view(B, H // window_size, window_size, W // window_size, window_size, C)
396 windows = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(-1, window_size, window_size, C)
397 return windows
398
399
400 def window_reverse(windows, batch_size: int, window_size: int, H: int, W: int):
401 B = batch_size
402 # this will cause onnx conversion failed for dynamic axis, because treated as constant
403 # int(windows.shape[0] / (H * W / window_size / window_size))
404 x = windows.view(B, H // window_size, W // window_size, window_size, window_size, -1)
405 x = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(B, H, W, -1)
406 return x
407
408
409 class WindowAttention(nn.Module):
410 def __init__(self, dim, num_heads, window_size, qkv_bias=True):
411
412 super().__init__()
413 self.dim = dim
414 self.window_size = window_size
415 self.num_heads = num_heads
416 head_dim = dim // num_heads
417 self.scale = float(head_dim) ** -0.5
418
419 self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
420 self.proj = nn.Linear(dim, dim)
421
422 self.softmax = nn.Softmax(dim=-1)
423
424 def forward(self, x, size):
425
426 H, W = size
427 B, L, C = x.shape
428 assert L == H * W, "input feature has wrong size"
429
430 x = x.view(B, H, W, C)
431
432 pad_l = pad_t = 0
433 pad_r = (self.window_size - W % self.window_size) % self.window_size
434 pad_b = (self.window_size - H % self.window_size) % self.window_size
435 x = F.pad(x, (0, 0, pad_l, pad_r, pad_t, pad_b))
436 _, Hp, Wp, _ = x.shape
437
438 x = window_partition(x, self.window_size)
439 x = x.view(-1, self.window_size * self.window_size, C)
440
441 # W-MSA/SW-MSA
442 # attn_windows = self.attn(x_windows)
443
444 B_, N, C = x.shape
445 qkv = self.qkv(x).reshape(B_, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)
446 q, k, v = qkv[0], qkv[1], qkv[2]
447
448 q = q * self.scale
449 attn = (q @ k.transpose(-2, -1))
450 attn = self.softmax(attn)
451
452 x = (attn @ v).transpose(1, 2).reshape(B_, N, C)
453 x = self.proj(x)
454
455 # merge windows
456 x = x.view(
457 -1, self.window_size, self.window_size, C
458 )
459 x = window_reverse(x, B, self.window_size, Hp, Wp)
460
461 if pad_r > 0 or pad_b > 0:
462 x = x[:, :H, :W, :].contiguous()
463
464 x = x.view(B, H * W, C)
465
466 return x, size
467
468
469 class SpatialBlock(nn.Module):
470
471 def __init__(self, dim, num_heads, window_size,
472 mlp_ratio=4., qkv_bias=True, drop_path_rate=0., act_layer=nn.GELU,
473 norm_layer=nn.LayerNorm, conv_at_attn=True, conv_at_ffn=True):
474 super().__init__()
475
476 drop_path = DropPath(drop_path_rate) if drop_path_rate > 0. else nn.Identity()
477
478 self.conv1 = PreNorm(None, DepthWiseConv2d(dim, 3, 1, 1)) if conv_at_attn else None
479 self.window_attn = PreNorm(
480 norm_layer(dim),
481 WindowAttention(dim, num_heads, window_size, qkv_bias=qkv_bias),
482 drop_path
483 )
484 self.conv2 = PreNorm(None, DepthWiseConv2d(dim, 3, 1, 1)) if conv_at_ffn else None
485 self.ffn = PreNorm(
486 norm_layer(dim),
487 Mlp(in_features=dim, hidden_features=int(dim*mlp_ratio), act_layer=act_layer),
488 drop_path
489 )
490
491 def forward(self, x, size):
492 if self.conv1:
493 x, size = self.conv1(x, size)
494 x, size = self.window_attn(x, size)
495
496 if self.conv2:
497 x, size = self.conv2(x, size)
498 x, size = self.ffn(x, size)
499 return x, size
500
501
502 class DaViT(nn.Module):
503 """ DaViT: Dual-Attention Transformer
504
505 Args:
506 in_chans (int): Number of input image channels. Default: 3.
507 num_classes (int): Number of classes for classification head. Default: 1000.
508 patch_size (tuple(int)): Patch size of convolution in different stages. Default: (7, 2, 2, 2).
509 patch_stride (tuple(int)): Patch stride of convolution in different stages. Default: (4, 2, 2, 2).
510 patch_padding (tuple(int)): Patch padding of convolution in different stages. Default: (3, 0, 0, 0).
511 patch_prenorm (tuple(bool)): If True, perform norm before convlution layer. Default: (True, False, False, False).
512 embed_dims (tuple(int)): Patch embedding dimension in different stages. Default: (64, 128, 192, 256).
513 num_heads (tuple(int)): Number of spatial attention heads in different stages. Default: (4, 8, 12, 16).
514 num_groups (tuple(int)): Number of channel groups in different stages. Default: (4, 8, 12, 16).
515 window_size (int): Window size. Default: 7.
516 mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. Default: 4.
517 qkv_bias (bool): If True, add a learnable bias to query, key, value. Default: True.
518 drop_path_rate (float): Stochastic depth rate. Default: 0.1.
519 norm_layer (nn.Module): Normalization layer. Default: nn.LayerNorm.
520 enable_checkpoint (bool): If True, enable checkpointing. Default: False.
521 conv_at_attn (bool): If True, performe depthwise convolution before attention layer. Default: True.
522 conv_at_ffn (bool): If True, performe depthwise convolution before ffn layer. Default: True.
523 """
524
525 def __init__(
526 self,
527 in_chans=3,
528 num_classes=1000,
529 depths=(1, 1, 3, 1),
530 patch_size=(7, 2, 2, 2),
531 patch_stride=(4, 2, 2, 2),
532 patch_padding=(3, 0, 0, 0),
533 patch_prenorm=(False, False, False, False),
534 embed_dims=(64, 128, 192, 256),
535 num_heads=(3, 6, 12, 24),
536 num_groups=(3, 6, 12, 24),
537 window_size=7,
538 mlp_ratio=4.,
539 qkv_bias=True,
540 drop_path_rate=0.1,
541 norm_layer=nn.LayerNorm,
542 enable_checkpoint=False,
543 conv_at_attn=True,
544 conv_at_ffn=True,
545 ):
546 super().__init__()
547
548 self.num_classes = num_classes
549 self.embed_dims = embed_dims
550 self.num_heads = num_heads
551 self.num_groups = num_groups
552 self.num_stages = len(self.embed_dims)
553 self.enable_checkpoint = enable_checkpoint
554 assert self.num_stages == len(self.num_heads) == len(self.num_groups)
555
556 num_stages = len(embed_dims)
557 dpr = [x.item() for x in torch.linspace(0, drop_path_rate, sum(depths)*2)]
558
559 depth_offset = 0
560 convs = []
561 blocks = []
562 for i in range(num_stages):
563 conv_embed = ConvEmbed(
564 patch_size=patch_size[i],
565 stride=patch_stride[i],
566 padding=patch_padding[i],
567 in_chans=in_chans if i == 0 else self.embed_dims[i - 1],
568 embed_dim=self.embed_dims[i],
569 norm_layer=norm_layer,
570 pre_norm=patch_prenorm[i]
571 )
572 convs.append(conv_embed)
573
574 block = MySequential(
575 *[
576 MySequential(OrderedDict([
577 (
578 'spatial_block', SpatialBlock(
579 embed_dims[i],
580 num_heads[i],
581 window_size,
582 drop_path_rate=dpr[depth_offset+j*2],
583 qkv_bias=qkv_bias,
584 mlp_ratio=mlp_ratio,
585 conv_at_attn=conv_at_attn,
586 conv_at_ffn=conv_at_ffn,
587 )
588 ),
589 (
590 'channel_block', ChannelBlock(
591 embed_dims[i],
592 num_groups[i],
593 drop_path_rate=dpr[depth_offset+j*2+1],
594 qkv_bias=qkv_bias,
595 mlp_ratio=mlp_ratio,
596 conv_at_attn=conv_at_attn,
597 conv_at_ffn=conv_at_ffn,
598 )
599 )
600 ])) for j in range(depths[i])
601 ]
602 )
603 blocks.append(block)
604 depth_offset += depths[i]*2
605
606 self.convs = nn.ModuleList(convs)
607 self.blocks = nn.ModuleList(blocks)
608
609 self.norms = norm_layer(self.embed_dims[-1])
610 self.avgpool = nn.AdaptiveAvgPool1d(1)
611 self.head = nn.Linear(self.embed_dims[-1], num_classes) if num_classes > 0 else nn.Identity()
612
613 @property
614 def dim_out(self):
615 return self.embed_dims[-1]
616
617 def forward_features_unpool(self, x):
618 """
619 forward until avg pooling
620 Args:
621 x (_type_): input image tensor
622 """
623 input_size = (x.size(2), x.size(3))
624 for conv, block in zip(self.convs, self.blocks):
625 x, input_size = conv(x, input_size)
626 if self.enable_checkpoint:
627 x, input_size = checkpoint.checkpoint(block, x, input_size)
628 else:
629 x, input_size = block(x, input_size)
630 return x
631
632 def forward_features(self, x):
633 x = self.forward_features_unpool(x)
634
635 # (batch_size, num_tokens, token_dim)
636 x = self.avgpool(x.transpose(1, 2))
637 # (batch_size, 1, num_tokens)
638 x = torch.flatten(x, 1)
639 x = self.norms(x)
640
641 return x
642
643 def forward(self, x):
644 x = self.forward_features(x)
645 x = self.head(x)
646 return x
647
648 @classmethod
649 def from_config(cls, config):
650 return cls(
651 depths=config.depths,
652 embed_dims=config.dim_embed,
653 num_heads=config.num_heads,
654 num_groups=config.num_groups,
655 patch_size=config.patch_size,
656 patch_stride=config.patch_stride,
657 patch_padding=config.patch_padding,
658 patch_prenorm=config.patch_prenorm,
659 drop_path_rate=config.drop_path_rate,
660 window_size=config.window_size,
661 )
662
663
664
665
666 if is_flash_attn_2_available():
667 from flash_attn import flash_attn_func, flash_attn_varlen_func
668 from flash_attn.bert_padding import index_first_axis, pad_input, unpad_input # noqa
669
670 # Copied from transformers.models.llama.modeling_llama._get_unpad_data
671 def _get_unpad_data(attention_mask):
672 seqlens_in_batch = attention_mask.sum(dim=-1, dtype=torch.int32)
673 indices = torch.nonzero(attention_mask.flatten(), as_tuple=False).flatten()
674 max_seqlen_in_batch = seqlens_in_batch.max().item()
675 cu_seqlens = F.pad(torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0))
676 return (
677 indices,
678 cu_seqlens,
679 max_seqlen_in_batch,
680 )
681
682
683 def shift_tokens_right(input_ids: torch.Tensor, pad_token_id: int, decoder_start_token_id: int):
684 """
685 Shift input ids one token to the right.
686 """
687 shifted_input_ids = input_ids.new_zeros(input_ids.shape)
688 shifted_input_ids[:, 1:] = input_ids[:, :-1].clone()
689 shifted_input_ids[:, 0] = decoder_start_token_id
690
691 if pad_token_id is None:
692 raise ValueError("self.model.config.pad_token_id has to be defined.")
693 # replace possible -100 values in labels by `pad_token_id`
694 shifted_input_ids.masked_fill_(shifted_input_ids == -100, pad_token_id)
695
696 return shifted_input_ids
697
698
699 class Florence2LearnedPositionalEmbedding(nn.Embedding):
700 """
701 This module learns positional embeddings up to a fixed maximum size.
702 """
703
704 def __init__(self, num_embeddings: int, embedding_dim: int):
705 # Florence2 is set up so that if padding_idx is specified then offset the embedding ids by 2
706 # and adjust num_embeddings appropriately. Other models don't have this hack
707 self.offset = 2
708 super().__init__(num_embeddings + self.offset, embedding_dim)
709
710 def forward(self, input_ids: torch.Tensor, past_key_values_length: int = 0):
711 """`input_ids' shape is expected to be [bsz x seqlen]."""
712
713 bsz, seq_len = input_ids.shape[:2]
714 positions = torch.arange(
715 past_key_values_length, past_key_values_length + seq_len, dtype=torch.long, device=self.weight.device
716 ).expand(bsz, -1)
717
718 return super().forward(positions + self.offset)
719
720
721 class Florence2ScaledWordEmbedding(nn.Embedding):
722 """
723 This module overrides nn.Embeddings' forward by multiplying with embeddings scale.
724 """
725
726 def __init__(self, num_embeddings: int, embedding_dim: int, padding_idx: int, embed_scale: Optional[float] = 1.0):
727 super().__init__(num_embeddings, embedding_dim, padding_idx)
728 self.embed_scale = embed_scale
729
730 def forward(self, input_ids: torch.Tensor):
731 return super().forward(input_ids) * self.embed_scale
732
733
734 class Florence2Attention(nn.Module):
735 """Multi-headed attention from 'Attention Is All You Need' paper"""
736
737 def __init__(
738 self,
739 embed_dim: int,
740 num_heads: int,
741 dropout: float = 0.0,
742 is_decoder: bool = False,
743 bias: bool = True,
744 is_causal: bool = False,
745 config: Optional[Florence2LanguageConfig] = None,
746 ):
747 super().__init__()
748 self.embed_dim = embed_dim
749 self.num_heads = num_heads
750 self.dropout = dropout
751 self.head_dim = embed_dim // num_heads
752 self.config = config
753
754 if (self.head_dim * num_heads) != self.embed_dim:
755 raise ValueError(
756 f"embed_dim must be divisible by num_heads (got `embed_dim`: {self.embed_dim}"
757 f" and `num_heads`: {num_heads})."
758 )
759 self.scaling = self.head_dim**-0.5
760 self.is_decoder = is_decoder
761 self.is_causal = is_causal
762
763 self.k_proj = nn.Linear(embed_dim, embed_dim, bias=bias)
764 self.v_proj = nn.Linear(embed_dim, embed_dim, bias=bias)
765 self.q_proj = nn.Linear(embed_dim, embed_dim, bias=bias)
766 self.out_proj = nn.Linear(embed_dim, embed_dim, bias=bias)
767
768 def _shape(self, tensor: torch.Tensor, seq_len: int, bsz: int):
769 return tensor.view(bsz, seq_len, self.num_heads, self.head_dim).transpose(1, 2).contiguous()
770
771 def forward(
772 self,
773 hidden_states: torch.Tensor,
774 key_value_states: Optional[torch.Tensor] = None,
775 past_key_value: Optional[Tuple[torch.Tensor]] = None,
776 attention_mask: Optional[torch.Tensor] = None,
777 layer_head_mask: Optional[torch.Tensor] = None,
778 output_attentions: bool = False,
779 ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]:
780 """Input shape: Batch x Time x Channel"""
781
782 # if key_value_states are provided this layer is used as a cross-attention layer
783 # for the decoder
784 is_cross_attention = key_value_states is not None
785
786 bsz, tgt_len, _ = hidden_states.size()
787
788 # get query proj
789 query_states = self.q_proj(hidden_states) * self.scaling
790 # get key, value proj
791 # `past_key_value[0].shape[2] == key_value_states.shape[1]`
792 # is checking that the `sequence_length` of the `past_key_value` is the same as
793 # the provided `key_value_states` to support prefix tuning
794 if (
795 is_cross_attention
796 and past_key_value is not None
797 and past_key_value[0].shape[2] == key_value_states.shape[1]
798 ):
799 # reuse k,v, cross_attentions
800 key_states = past_key_value[0]
801 value_states = past_key_value[1]
802 elif is_cross_attention:
803 # cross_attentions
804 key_states = self._shape(self.k_proj(key_value_states), -1, bsz)
805 value_states = self._shape(self.v_proj(key_value_states), -1, bsz)
806 elif past_key_value is not None:
807 # reuse k, v, self_attention
808 key_states = self._shape(self.k_proj(hidden_states), -1, bsz)
809 value_states = self._shape(self.v_proj(hidden_states), -1, bsz)
810 key_states = torch.cat([past_key_value[0], key_states], dim=2)
811 value_states = torch.cat([past_key_value[1], value_states], dim=2)
812 else:
813 # self_attention
814 key_states = self._shape(self.k_proj(hidden_states), -1, bsz)
815 value_states = self._shape(self.v_proj(hidden_states), -1, bsz)
816
817 if self.is_decoder:
818 # if cross_attention save Tuple(torch.Tensor, torch.Tensor) of all cross attention key/value_states.
819 # Further calls to cross_attention layer can then reuse all cross-attention
820 # key/value_states (first "if" case)
821 # if uni-directional self-attention (decoder) save Tuple(torch.Tensor, torch.Tensor) of
822 # all previous decoder key/value_states. Further calls to uni-directional self-attention
823 # can concat previous decoder key/value_states to current projected key/value_states (third "elif" case)
824 # if encoder bi-directional self-attention `past_key_value` is always `None`
825 past_key_value = (key_states, value_states)
826
827 proj_shape = (bsz * self.num_heads, -1, self.head_dim)
828 query_states = self._shape(query_states, tgt_len, bsz).view(*proj_shape)
829 key_states = key_states.reshape(*proj_shape)
830 value_states = value_states.reshape(*proj_shape)
831
832 src_len = key_states.size(1)
833 attn_weights = torch.bmm(query_states, key_states.transpose(1, 2))
834
835 if attn_weights.size() != (bsz * self.num_heads, tgt_len, src_len):
836 raise ValueError(
837 f"Attention weights should be of size {(bsz * self.num_heads, tgt_len, src_len)}, but is"
838 f" {attn_weights.size()}"
839 )
840
841 if attention_mask is not None:
842 if attention_mask.size() != (bsz, 1, tgt_len, src_len):
843 raise ValueError(
844 f"Attention mask should be of size {(bsz, 1, tgt_len, src_len)}, but is {attention_mask.size()}"
845 )
846 attn_weights = attn_weights.view(bsz, self.num_heads, tgt_len, src_len) + attention_mask
847 attn_weights = attn_weights.view(bsz * self.num_heads, tgt_len, src_len)
848
849 attn_weights = nn.functional.softmax(attn_weights, dim=-1)
850
851 if layer_head_mask is not None:
852 if layer_head_mask.size() != (self.num_heads,):
853 raise ValueError(
854 f"Head mask for a single layer should be of size {(self.num_heads,)}, but is"
855 f" {layer_head_mask.size()}"
856 )
857 attn_weights = layer_head_mask.view(1, -1, 1, 1) * attn_weights.view(bsz, self.num_heads, tgt_len, src_len)
858 attn_weights = attn_weights.view(bsz * self.num_heads, tgt_len, src_len)
859
860 if output_attentions:
861 # this operation is a bit awkward, but it's required to
862 # make sure that attn_weights keeps its gradient.
863 # In order to do so, attn_weights have to be reshaped
864 # twice and have to be reused in the following
865 attn_weights_reshaped = attn_weights.view(bsz, self.num_heads, tgt_len, src_len)
866 attn_weights = attn_weights_reshaped.view(bsz * self.num_heads, tgt_len, src_len)
867 else:
868 attn_weights_reshaped = None
869
870 attn_probs = nn.functional.dropout(attn_weights, p=self.dropout, training=self.training)
871
872 attn_output = torch.bmm(attn_probs, value_states)
873
874 if attn_output.size() != (bsz * self.num_heads, tgt_len, self.head_dim):
875 raise ValueError(
876 f"`attn_output` should be of size {(bsz * self.num_heads, tgt_len, self.head_dim)}, but is"
877 f" {attn_output.size()}"
878 )
879
880 attn_output = attn_output.view(bsz, self.num_heads, tgt_len, self.head_dim)
881 attn_output = attn_output.transpose(1, 2)
882
883 # Use the `embed_dim` from the config (stored in the class) rather than `hidden_state` because `attn_output` can be
884 # partitioned across GPUs when using tensor-parallelism.
885 attn_output = attn_output.reshape(bsz, tgt_len, self.embed_dim)
886
887 attn_output = self.out_proj(attn_output)
888
889 return attn_output, attn_weights_reshaped, past_key_value
890
891
892 class Florence2FlashAttention2(Florence2Attention):
893 """
894 Florence2 flash attention module. This module inherits from `Florence2Attention` as the weights of the module stays
895 untouched. The only required change would be on the forward pass where it needs to correctly call the public API of
896 flash attention and deal with padding tokens in case the input contains any of them.
897 """
898
899 # Copied from transformers.models.llama.modeling_llama.LlamaFlashAttention2.__init__
900 def __init__(self, *args, **kwargs):
901 super().__init__(*args, **kwargs)
902
903 # TODO: Should be removed once Flash Attention for RoCm is bumped to 2.1.
904 # flash_attn<2.1 generates top-left aligned causal mask, while what is needed here is bottom-right alignement, that was made default for flash_attn>=2.1. This attribute is used to handle this difference. Reference: https://github.com/Dao-AILab/flash-attention/releases/tag/v2.1.0.
905 # Beware that with flash_attn<2.1, using q_seqlen != k_seqlen (except for the case q_seqlen == 1) produces a wrong mask (top-left).
906 self._flash_attn_uses_top_left_mask = not is_flash_attn_greater_or_equal_2_10()
907
908 def _reshape(self, tensor: torch.Tensor, seq_len: int, bsz: int):
909 return tensor.view(bsz, seq_len, self.num_heads, self.head_dim)
910
911 def forward(
912 self,
913 hidden_states: torch.Tensor,
914 key_value_states: Optional[torch.Tensor] = None,
915 past_key_value: Optional[Tuple[torch.Tensor]] = None,
916 attention_mask: Optional[torch.Tensor] = None,
917 layer_head_mask: Optional[torch.Tensor] = None,
918 output_attentions: bool = False,
919 ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]:
920 # Florence2FlashAttention2 attention does not support output_attentions
921 if output_attentions:
922 raise ValueError("Florence2FlashAttention2 attention does not support output_attentions")
923
924 # if key_value_states are provided this layer is used as a cross-attention layer
925 # for the decoder
926 is_cross_attention = key_value_states is not None
927
928 bsz, q_len, _ = hidden_states.size()
929
930 # get query proj
931 query_states = self._reshape(self.q_proj(hidden_states), -1, bsz)
932 # get key, value proj
933 # `past_key_value[0].shape[2] == key_value_states.shape[1]`
934 # is checking that the `sequence_length` of the `past_key_value` is the same as
935 # the provided `key_value_states` to support prefix tuning
936 if (
937 is_cross_attention
938 and past_key_value is not None
939 and past_key_value[0].shape[2] == key_value_states.shape[1]
940 ):
941 # reuse k,v, cross_attentions
942 key_states = past_key_value[0].transpose(1, 2)
943 value_states = past_key_value[1].transpose(1, 2)
944 elif is_cross_attention:
945 # cross_attentions
946 key_states = self._reshape(self.k_proj(key_value_states), -1, bsz)
947 value_states = self._reshape(self.v_proj(key_value_states), -1, bsz)
948 elif past_key_value is not None:
949 # reuse k, v, self_attention
950 key_states = self._reshape(self.k_proj(hidden_states), -1, bsz)
951 value_states = self._reshape(self.v_proj(hidden_states), -1, bsz)
952 key_states = torch.cat([past_key_value[0].transpose(1, 2), key_states], dim=1)
953 value_states = torch.cat([past_key_value[1].transpose(1, 2), value_states], dim=1)
954 else:
955 # self_attention
956 key_states = self._reshape(self.k_proj(hidden_states), -1, bsz)
957 value_states = self._reshape(self.v_proj(hidden_states), -1, bsz)
958
959 if self.is_decoder:
960 # if cross_attention save Tuple(torch.Tensor, torch.Tensor) of all cross attention key/value_states.
961 # Further calls to cross_attention layer can then reuse all cross-attention
962 # key/value_states (first "if" case)
963 # if uni-directional self-attention (decoder) save Tuple(torch.Tensor, torch.Tensor) of
964 # all previous decoder key/value_states. Further calls to uni-directional self-attention
965 # can concat previous decoder key/value_states to current projected key/value_states (third "elif" case)
966 # if encoder bi-directional self-attention `past_key_value` is always `None`
967 past_key_value = (key_states.transpose(1, 2), value_states.transpose(1, 2))
968
969 kv_seq_len = key_states.shape[-2]
970 if past_key_value is not None:
971 kv_seq_len += past_key_value[0].shape[-2]
972
973 # In PEFT, usually we cast the layer norms in float32 for training stability reasons
974 # therefore the input hidden states gets silently casted in float32. Hence, we need
975 # cast them back in the correct dtype just to be sure everything works as expected.
976 # This might slowdown training & inference so it is recommended to not cast the LayerNorms
977 # in fp32. (LlamaRMSNorm handles it correctly)
978
979 input_dtype = query_states.dtype
980 if input_dtype == torch.float32:
981 if torch.is_autocast_enabled():
982 target_dtype = torch.get_autocast_gpu_dtype()
983 # Handle the case where the model is quantized
984 elif hasattr(self.config, "_pre_quantization_dtype"):
985 target_dtype = self.config._pre_quantization_dtype
986 else:
987 target_dtype = self.q_proj.weight.dtype
988
989 logger.warning_once(
990 f"The input hidden states seems to be silently casted in float32, this might be related to"
991 f" the fact you have upcasted embedding or layer norm layers in float32. We will cast back the input in"
992 f" {target_dtype}."
993 )
994
995 query_states = query_states.to(target_dtype)
996 key_states = key_states.to(target_dtype)
997 value_states = value_states.to(target_dtype)
998
999 attn_output = self._flash_attention_forward(
1000 query_states, key_states, value_states, attention_mask, q_len, dropout=self.dropout
1001 )
1002
1003 attn_output = attn_output.reshape(bsz, q_len, -1)
1004 attn_output = self.out_proj(attn_output)
1005
1006 if not output_attentions:
1007 attn_weights = None
1008
1009 return attn_output, attn_weights, past_key_value
1010
1011 # Copied from transformers.models.llama.modeling_llama.LlamaFlashAttention2._flash_attention_forward
1012 def _flash_attention_forward(
1013 self, query_states, key_states, value_states, attention_mask, query_length, dropout=0.0, softmax_scale=None
1014 ):
1015 """
1016 Calls the forward method of Flash Attention - if the input hidden states contain at least one padding token
1017 first unpad the input, then computes the attention scores and pad the final attention scores.
1018
1019 Args:
1020 query_states (`torch.Tensor`):
1021 Input query states to be passed to Flash Attention API
1022 key_states (`torch.Tensor`):
1023 Input key states to be passed to Flash Attention API
1024 value_states (`torch.Tensor`):
1025 Input value states to be passed to Flash Attention API
1026 attention_mask (`torch.Tensor`):
1027 The padding mask - corresponds to a tensor of size `(batch_size, seq_len)` where 0 stands for the
1028 position of padding tokens and 1 for the position of non-padding tokens.
1029 dropout (`float`):
1030 Attention dropout
1031 softmax_scale (`float`, *optional*):
1032 The scaling of QK^T before applying softmax. Default to 1 / sqrt(head_dim)
1033 """
1034 if not self._flash_attn_uses_top_left_mask:
1035 causal = self.is_causal
1036 else:
1037 # TODO: Remove the `query_length != 1` check once Flash Attention for RoCm is bumped to 2.1. For details, please see the comment in LlamaFlashAttention2 __init__.
1038 causal = self.is_causal and query_length != 1
1039
1040 # Contains at least one padding token in the sequence
1041 if attention_mask is not None:
1042 batch_size = query_states.shape[0]
1043 query_states, key_states, value_states, indices_q, cu_seq_lens, max_seq_lens = self._upad_input(
1044 query_states, key_states, value_states, attention_mask, query_length
1045 )
1046
1047 cu_seqlens_q, cu_seqlens_k = cu_seq_lens
1048 max_seqlen_in_batch_q, max_seqlen_in_batch_k = max_seq_lens
1049
1050 attn_output_unpad = flash_attn_varlen_func(
1051 query_states,
1052 key_states,
1053 value_states,
1054 cu_seqlens_q=cu_seqlens_q,
1055 cu_seqlens_k=cu_seqlens_k,
1056 max_seqlen_q=max_seqlen_in_batch_q,
1057 max_seqlen_k=max_seqlen_in_batch_k,
1058 dropout_p=dropout,
1059 softmax_scale=softmax_scale,
1060 causal=causal,
1061 )
1062
1063 attn_output = pad_input(attn_output_unpad, indices_q, batch_size, query_length)
1064 else:
1065 attn_output = flash_attn_func(
1066 query_states, key_states, value_states, dropout, softmax_scale=softmax_scale, causal=causal
1067 )
1068
1069 return attn_output
1070
1071 # Copied from transformers.models.llama.modeling_llama.LlamaFlashAttention2._upad_input
1072 def _upad_input(self, query_layer, key_layer, value_layer, attention_mask, query_length):
1073 indices_k, cu_seqlens_k, max_seqlen_in_batch_k = _get_unpad_data(attention_mask)
1074 batch_size, kv_seq_len, num_key_value_heads, head_dim = key_layer.shape
1075
1076 key_layer = index_first_axis(
1077 key_layer.reshape(batch_size * kv_seq_len, num_key_value_heads, head_dim), indices_k
1078 )
1079 value_layer = index_first_axis(
1080 value_layer.reshape(batch_size * kv_seq_len, num_key_value_heads, head_dim), indices_k
1081 )
1082 if query_length == kv_seq_len:
1083 query_layer = index_first_axis(
1084 query_layer.reshape(batch_size * kv_seq_len, self.num_heads, head_dim), indices_k
1085 )
1086 cu_seqlens_q = cu_seqlens_k
1087 max_seqlen_in_batch_q = max_seqlen_in_batch_k
1088 indices_q = indices_k
1089 elif query_length == 1:
1090 max_seqlen_in_batch_q = 1
1091 cu_seqlens_q = torch.arange(
1092 batch_size + 1, dtype=torch.int32, device=query_layer.device
1093 ) # There is a memcpy here, that is very bad.
1094 indices_q = cu_seqlens_q[:-1]
1095 query_layer = query_layer.squeeze(1)
1096 else:
1097 # The -q_len: slice assumes left padding.
1098 attention_mask = attention_mask[:, -query_length:]
1099 query_layer, indices_q, cu_seqlens_q, max_seqlen_in_batch_q = unpad_input(query_layer, attention_mask)
1100
1101 return (
1102 query_layer,
1103 key_layer,
1104 value_layer,
1105 indices_q,
1106 (cu_seqlens_q, cu_seqlens_k),
1107 (max_seqlen_in_batch_q, max_seqlen_in_batch_k),
1108 )
1109
1110
1111 class Florence2SdpaAttention(Florence2Attention):
1112 def forward(
1113 self,
1114 hidden_states: torch.Tensor,
1115 key_value_states: Optional[torch.Tensor] = None,
1116 past_key_value: Optional[Tuple[torch.Tensor]] = None,
1117 attention_mask: Optional[torch.Tensor] = None,
1118 layer_head_mask: Optional[torch.Tensor] = None,
1119 output_attentions: bool = False,
1120 ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]:
1121 """Input shape: Batch x Time x Channel"""
1122 if output_attentions or layer_head_mask is not None:
1123 # TODO: Improve this warning with e.g. `model.config._attn_implementation = "manual"` once this is implemented.
1124 logger.warning_once(
1125 "Florence2Model is using Florence2SdpaAttention, but `torch.nn.functional.scaled_dot_product_attention` does not support `output_attentions=True` or `layer_head_mask` not None. Falling back to the manual attention"
1126 ' implementation, but specifying the manual implementation will be required from Transformers version v5.0.0 onwards. This warning can be removed using the argument `attn_implementation="eager"` when loading the model.'
1127 )
1128 return super().forward(
1129 hidden_states,
1130 key_value_states=key_value_states,
1131 past_key_value=past_key_value,
1132 attention_mask=attention_mask,
1133 layer_head_mask=layer_head_mask,
1134 output_attentions=output_attentions,
1135 )
1136
1137 # if key_value_states are provided this layer is used as a cross-attention layer
1138 # for the decoder
1139 is_cross_attention = key_value_states is not None
1140
1141 bsz, tgt_len, _ = hidden_states.size()
1142
1143 # get query proj
1144 query_states = self.q_proj(hidden_states)
1145 # get key, value proj
1146 # `past_key_value[0].shape[2] == key_value_states.shape[1]`
1147 # is checking that the `sequence_length` of the `past_key_value` is the same as
1148 # the provided `key_value_states` to support prefix tuning
1149 if (
1150 is_cross_attention
1151 and past_key_value is not None
1152 and past_key_value[0].shape[2] == key_value_states.shape[1]
1153 ):
1154 # reuse k,v, cross_attentions
1155 key_states = past_key_value[0]
1156 value_states = past_key_value[1]
1157 elif is_cross_attention:
1158 # cross_attentions
1159 key_states = self._shape(self.k_proj(key_value_states), -1, bsz)
1160 value_states = self._shape(self.v_proj(key_value_states), -1, bsz)
1161 elif past_key_value is not None:
1162 # reuse k, v, self_attention
1163 key_states = self._shape(self.k_proj(hidden_states), -1, bsz)
1164 value_states = self._shape(self.v_proj(hidden_states), -1, bsz)
1165 key_states = torch.cat([past_key_value[0], key_states], dim=2)
1166 value_states = torch.cat([past_key_value[1], value_states], dim=2)
1167 else:
1168 # self_attention
1169 key_states = self._shape(self.k_proj(hidden_states), -1, bsz)
1170 value_states = self._shape(self.v_proj(hidden_states), -1, bsz)
1171
1172 if self.is_decoder:
1173 # if cross_attention save Tuple(torch.Tensor, torch.Tensor) of all cross attention key/value_states.
1174 # Further calls to cross_attention layer can then reuse all cross-attention
1175 # key/value_states (first "if" case)
1176 # if uni-directional self-attention (decoder) save Tuple(torch.Tensor, torch.Tensor) of
1177 # all previous decoder key/value_states. Further calls to uni-directional self-attention
1178 # can concat previous decoder key/value_states to current projected key/value_states (third "elif" case)
1179 # if encoder bi-directional self-attention `past_key_value` is always `None`
1180 past_key_value = (key_states, value_states)
1181
1182 query_states = self._shape(query_states, tgt_len, bsz)
1183
1184 # We dispatch to SDPA's Flash Attention or Efficient kernels via this `is_causal` if statement instead of an inline conditional assignment
1185 # in SDPA to support both torch.compile's dynamic shapes and full graph options. An inline conditional prevents dynamic shapes from compiling.
1186 # The tgt_len > 1 is necessary to match with AttentionMaskConverter.to_causal_4d that does not create a causal mask in case tgt_len == 1.
1187 is_causal = True if self.is_causal and attention_mask is None and tgt_len > 1 else False
1188
1189 # NOTE: SDPA with memory-efficient backend is currently (torch==2.1.2) bugged when using non-contiguous inputs and a custom attn_mask,
1190 # but we are fine here as `_shape` do call `.contiguous()`. Reference: https://github.com/pytorch/pytorch/issues/112577
1191 attn_output = torch.nn.functional.scaled_dot_product_attention(
1192 query_states,
1193 key_states,
1194 value_states,
1195 attn_mask=attention_mask,
1196 dropout_p=self.dropout if self.training else 0.0,
1197 is_causal=is_causal,
1198 )
1199
1200 if attn_output.size() != (bsz, self.num_heads, tgt_len, self.head_dim):
1201 raise ValueError(
1202 f"`attn_output` should be of size {(bsz, self.num_heads, tgt_len, self.head_dim)}, but is"
1203 f" {attn_output.size()}"
1204 )
1205
1206 attn_output = attn_output.transpose(1, 2)
1207
1208 # Use the `embed_dim` from the config (stored in the class) rather than `hidden_state` because `attn_output` can be
1209 # partitioned across GPUs when using tensor-parallelism.
1210 attn_output = attn_output.reshape(bsz, tgt_len, self.embed_dim)
1211
1212 attn_output = self.out_proj(attn_output)
1213
1214 return attn_output, None, past_key_value
1215
1216
1217 FLORENCE2_ATTENTION_CLASSES = {
1218 "eager": Florence2Attention,
1219 "sdpa": Florence2SdpaAttention,
1220 "flash_attention_2": Florence2FlashAttention2,
1221 }
1222
1223
1224 class Florence2EncoderLayer(nn.Module):
1225 def __init__(self, config: Florence2LanguageConfig):
1226 super().__init__()
1227 self.embed_dim = config.d_model
1228
1229 self.self_attn = FLORENCE2_ATTENTION_CLASSES[config._attn_implementation](
1230 embed_dim=self.embed_dim,
1231 num_heads=config.encoder_attention_heads,
1232 dropout=config.attention_dropout,
1233 config=config,
1234 )
1235 self.self_attn_layer_norm = nn.LayerNorm(self.embed_dim)
1236 self.dropout = config.dropout
1237 self.activation_fn = ACT2FN[config.activation_function]
1238 self.activation_dropout = config.activation_dropout
1239 self.fc1 = nn.Linear(self.embed_dim, config.encoder_ffn_dim)
1240 self.fc2 = nn.Linear(config.encoder_ffn_dim, self.embed_dim)
1241 self.final_layer_norm = nn.LayerNorm(self.embed_dim)
1242
1243 def forward(
1244 self,
1245 hidden_states: torch.FloatTensor,
1246 attention_mask: torch.FloatTensor,
1247 layer_head_mask: torch.FloatTensor,
1248 output_attentions: Optional[bool] = False,
1249 ) -> Tuple[torch.FloatTensor, Optional[torch.FloatTensor]]:
1250 """
1251 Args:
1252 hidden_states (`torch.FloatTensor`): input to the layer of shape `(batch, seq_len, embed_dim)`
1253 attention_mask (`torch.FloatTensor`): attention mask of size
1254 `(batch, 1, tgt_len, src_len)` where padding elements are indicated by very large negative values.
1255 layer_head_mask (`torch.FloatTensor`): mask for attention heads in a given layer of size
1256 `(encoder_attention_heads,)`.
1257 output_attentions (`bool`, *optional*):
1258 Whether or not to return the attentions tensors of all attention layers. See `attentions` under
1259 returned tensors for more detail.
1260 """
1261 residual = hidden_states
1262 hidden_states, attn_weights, _ = self.self_attn(
1263 hidden_states=hidden_states,
1264 attention_mask=attention_mask,
1265 layer_head_mask=layer_head_mask,
1266 output_attentions=output_attentions,
1267 )
1268 hidden_states = nn.functional.dropout(hidden_states, p=self.dropout, training=self.training)
1269 hidden_states = residual + hidden_states
1270 hidden_states = self.self_attn_layer_norm(hidden_states)
1271
1272 residual = hidden_states
1273 hidden_states = self.activation_fn(self.fc1(hidden_states))
1274 hidden_states = nn.functional.dropout(hidden_states, p=self.activation_dropout, training=self.training)
1275 hidden_states = self.fc2(hidden_states)
1276 hidden_states = nn.functional.dropout(hidden_states, p=self.dropout, training=self.training)
1277 hidden_states = residual + hidden_states
1278 hidden_states = self.final_layer_norm(hidden_states)
1279
1280 if hidden_states.dtype == torch.float16 and (
1281 torch.isinf(hidden_states).any() or torch.isnan(hidden_states).any()
1282 ):
1283 clamp_value = torch.finfo(hidden_states.dtype).max - 1000
1284 hidden_states = torch.clamp(hidden_states, min=-clamp_value, max=clamp_value)
1285
1286 outputs = (hidden_states,)
1287
1288 if output_attentions:
1289 outputs += (attn_weights,)
1290
1291 return outputs
1292
1293
1294 class Florence2DecoderLayer(nn.Module):
1295 def __init__(self, config: Florence2LanguageConfig):
1296 super().__init__()
1297 self.embed_dim = config.d_model
1298
1299 self.self_attn = FLORENCE2_ATTENTION_CLASSES[config._attn_implementation](
1300 embed_dim=self.embed_dim,
1301 num_heads=config.decoder_attention_heads,
1302 dropout=config.attention_dropout,
1303 is_decoder=True,
1304 is_causal=True,
1305 config=config,
1306 )
1307 self.dropout = config.dropout
1308 self.activation_fn = ACT2FN[config.activation_function]
1309 self.activation_dropout = config.activation_dropout
1310
1311 self.self_attn_layer_norm = nn.LayerNorm(self.embed_dim)
1312 self.encoder_attn = FLORENCE2_ATTENTION_CLASSES[config._attn_implementation](
1313 self.embed_dim,
1314 config.decoder_attention_heads,
1315 dropout=config.attention_dropout,
1316 is_decoder=True,
1317 config=config,
1318 )
1319 self.encoder_attn_layer_norm = nn.LayerNorm(self.embed_dim)
1320 self.fc1 = nn.Linear(self.embed_dim, config.decoder_ffn_dim)
1321 self.fc2 = nn.Linear(config.decoder_ffn_dim, self.embed_dim)
1322 self.final_layer_norm = nn.LayerNorm(self.embed_dim)
1323
1324 def forward(
1325 self,
1326 hidden_states: torch.Tensor,
1327 attention_mask: Optional[torch.Tensor] = None,
1328 encoder_hidden_states: Optional[torch.Tensor] = None,
1329 encoder_attention_mask: Optional[torch.Tensor] = None,
1330 layer_head_mask: Optional[torch.Tensor] = None,
1331 cross_attn_layer_head_mask: Optional[torch.Tensor] = None,
1332 past_key_value: Optional[Tuple[torch.Tensor]] = None,
1333 output_attentions: Optional[bool] = False,
1334 use_cache: Optional[bool] = True,
1335 ) -> Tuple[torch.FloatTensor, Optional[Tuple[torch.FloatTensor, torch.FloatTensor]]]:
1336 """
1337 Args:
1338 hidden_states (`torch.FloatTensor`): input to the layer of shape `(batch, seq_len, embed_dim)`
1339 attention_mask (`torch.FloatTensor`): attention mask of size
1340 `(batch, 1, tgt_len, src_len)` where padding elements are indicated by very large negative values.
1341 encoder_hidden_states (`torch.FloatTensor`):
1342 cross attention input to the layer of shape `(batch, seq_len, embed_dim)`
1343 encoder_attention_mask (`torch.FloatTensor`): encoder attention mask of size
1344 `(batch, 1, tgt_len, src_len)` where padding elements are indicated by very large negative values.
1345 layer_head_mask (`torch.FloatTensor`): mask for attention heads in a given layer of size
1346 `(encoder_attention_heads,)`.
1347 cross_attn_layer_head_mask (`torch.FloatTensor`): mask for cross-attention heads in a given layer of
1348 size `(decoder_attention_heads,)`.
1349 past_key_value (`Tuple(torch.FloatTensor)`): cached past key and value projection states
1350 output_attentions (`bool`, *optional*):
1351 Whether or not to return the attentions tensors of all attention layers. See `attentions` under
1352 returned tensors for more detail.
1353 """
1354 residual = hidden_states
1355
1356 # Self Attention
1357 # decoder uni-directional self-attention cached key/values tuple is at positions 1,2
1358 self_attn_past_key_value = past_key_value[:2] if past_key_value is not None else None
1359 # add present self-attn cache to positions 1,2 of present_key_value tuple
1360 hidden_states, self_attn_weights, present_key_value = self.self_attn(
1361 hidden_states=hidden_states,
1362 past_key_value=self_attn_past_key_value,
1363 attention_mask=attention_mask,
1364 layer_head_mask=layer_head_mask,
1365 output_attentions=output_attentions,
1366 )
1367 hidden_states = nn.functional.dropout(hidden_states, p=self.dropout, training=self.training)
1368 hidden_states = residual + hidden_states
1369 hidden_states = self.self_attn_layer_norm(hidden_states)
1370
1371 # Cross-Attention Block
1372 cross_attn_present_key_value = None
1373 cross_attn_weights = None
1374 if encoder_hidden_states is not None:
1375 residual = hidden_states
1376
1377 # cross_attn cached key/values tuple is at positions 3,4 of present_key_value tuple
1378 cross_attn_past_key_value = past_key_value[-2:] if past_key_value is not None else None
1379 hidden_states, cross_attn_weights, cross_attn_present_key_value = self.encoder_attn(
1380 hidden_states=hidden_states,
1381 key_value_states=encoder_hidden_states,
1382 attention_mask=encoder_attention_mask,
1383 layer_head_mask=cross_attn_layer_head_mask,
1384 past_key_value=cross_attn_past_key_value,
1385 output_attentions=output_attentions,
1386 )
1387 hidden_states = nn.functional.dropout(hidden_states, p=self.dropout, training=self.training)
1388 hidden_states = residual + hidden_states
1389 hidden_states = self.encoder_attn_layer_norm(hidden_states)
1390
1391 # add cross-attn to positions 3,4 of present_key_value tuple
1392 present_key_value = present_key_value + cross_attn_present_key_value
1393
1394 # Fully Connected
1395 residual = hidden_states
1396 hidden_states = self.activation_fn(self.fc1(hidden_states))
1397 hidden_states = nn.functional.dropout(hidden_states, p=self.activation_dropout, training=self.training)
1398 hidden_states = self.fc2(hidden_states)
1399 hidden_states = nn.functional.dropout(hidden_states, p=self.dropout, training=self.training)
1400 hidden_states = residual + hidden_states
1401 hidden_states = self.final_layer_norm(hidden_states)
1402
1403 outputs = (hidden_states,)
1404
1405 if output_attentions:
1406 outputs += (self_attn_weights, cross_attn_weights)
1407
1408 if use_cache:
1409 outputs += (present_key_value,)
1410
1411 return outputs
1412
1413
1414
1415 class Florence2LanguagePreTrainedModel(PreTrainedModel):
1416 config_class = Florence2LanguageConfig
1417 base_model_prefix = "model"
1418 supports_gradient_checkpointing = True
1419 _keys_to_ignore_on_load_unexpected = ["encoder.version", "decoder.version"]
1420 _no_split_modules = [r"Florence2EncoderLayer", r"Florence2DecoderLayer"]
1421 _skip_keys_device_placement = "past_key_values"
1422 _supports_flash_attn_2 = True
1423 _supports_sdpa = True
1424
1425 def _init_weights(self, module):
1426 std = self.config.init_std
1427 if isinstance(module, nn.Linear):
1428 module.weight.data.normal_(mean=0.0, std=std)
1429 if module.bias is not None:
1430 module.bias.data.zero_()
1431 elif isinstance(module, nn.Embedding):
1432 module.weight.data.normal_(mean=0.0, std=std)
1433 if module.padding_idx is not None:
1434 module.weight.data[module.padding_idx].zero_()
1435 elif isinstance(module, nn.Conv2d):
1436 nn.init.normal_(module.weight, std=0.02)
1437 for name, _ in module.named_parameters():
1438 if name == "bias":
1439 nn.init.constant_(module.bias, 0)
1440 elif isinstance(module, nn.LayerNorm):
1441 nn.init.constant_(module.weight, 1.0)
1442 nn.init.constant_(module.bias, 0)
1443 elif isinstance(module, nn.BatchNorm2d):
1444 nn.init.constant_(module.weight, 1.0)
1445 nn.init.constant_(module.bias, 0)
1446
1447 @property
1448 def dummy_inputs(self):
1449 pad_token = self.config.pad_token_id
1450 input_ids = torch.tensor([[0, 6, 10, 4, 2], [0, 8, 12, 2, pad_token]], device=self.device)
1451 dummy_inputs = {
1452 "attention_mask": input_ids.ne(pad_token),
1453 "input_ids": input_ids,
1454 }
1455 return dummy_inputs
1456
1457
1458 class Florence2Encoder(Florence2LanguagePreTrainedModel):
1459 """
1460 Transformer encoder consisting of *config.encoder_layers* self attention layers. Each layer is a
1461 [`Florence2EncoderLayer`].
1462
1463 Args:
1464 config: Florence2LanguageConfig
1465 embed_tokens (nn.Embedding): output embedding
1466 """
1467
1468 def __init__(self, config: Florence2LanguageConfig, embed_tokens: Optional[nn.Embedding] = None):
1469 super().__init__(config)
1470
1471 self.dropout = config.dropout
1472 self.layerdrop = config.encoder_layerdrop
1473
1474 embed_dim = config.d_model
1475 self.padding_idx = config.pad_token_id
1476 self.max_source_positions = config.max_position_embeddings
1477 embed_scale = math.sqrt(embed_dim) if config.scale_embedding else 1.0
1478
1479 self.embed_tokens = Florence2ScaledWordEmbedding(
1480 config.vocab_size, embed_dim, self.padding_idx, embed_scale=embed_scale
1481 )
1482
1483 if embed_tokens is not None:
1484 self.embed_tokens.weight = embed_tokens.weight
1485
1486 self.embed_positions = Florence2LearnedPositionalEmbedding(
1487 config.max_position_embeddings,
1488 embed_dim,
1489 )
1490 self.layers = nn.ModuleList([Florence2EncoderLayer(config) for _ in range(config.encoder_layers)])
1491 self._use_flash_attention_2 = config._attn_implementation == "flash_attention_2"
1492 self._use_sdpa = config._attn_implementation == "sdpa"
1493 self.layernorm_embedding = nn.LayerNorm(embed_dim)
1494
1495 self.gradient_checkpointing = False
1496 # Initialize weights and apply final processing
1497 self.post_init()
1498
1499 def get_input_embeddings(self):
1500 return self.embed_tokens
1501
1502 def set_input_embeddings(self, value):
1503 self.embed_tokens = value
1504
1505 def forward(
1506 self,
1507 input_ids: torch.LongTensor = None,
1508 attention_mask: Optional[torch.Tensor] = None,
1509 head_mask: Optional[torch.Tensor] = None,
1510 inputs_embeds: Optional[torch.FloatTensor] = None,
1511 output_attentions: Optional[bool] = None,
1512 output_hidden_states: Optional[bool] = None,
1513 return_dict: Optional[bool] = None,
1514 ) -> Union[Tuple, BaseModelOutput]:
1515 r"""
1516 Args:
1517 input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`):
1518 Indices of input sequence tokens in the vocabulary. Padding will be ignored by default should you
1519 provide it.
1520
1521 Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and
1522 [`PreTrainedTokenizer.__call__`] for details.
1523
1524 [What are input IDs?](../glossary#input-ids)
1525 attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):
1526 Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:
1527
1528 - 1 for tokens that are **not masked**,
1529 - 0 for tokens that are **masked**.
1530
1531 [What are attention masks?](../glossary#attention-mask)
1532 head_mask (`torch.Tensor` of shape `(encoder_layers, encoder_attention_heads)`, *optional*):
1533 Mask to nullify selected heads of the attention modules. Mask values selected in `[0, 1]`:
1534
1535 - 1 indicates the head is **not masked**,
1536 - 0 indicates the head is **masked**.
1537
1538 inputs_embeds (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`, *optional*):
1539 Optionally, instead of passing `input_ids` you can choose to directly pass an embedded representation.
1540 This is useful if you want more control over how to convert `input_ids` indices into associated vectors
1541 than the model's internal embedding lookup matrix.
1542 output_attentions (`bool`, *optional*):
1543 Whether or not to return the attentions tensors of all attention layers. See `attentions` under
1544 returned tensors for more detail.
1545 output_hidden_states (`bool`, *optional*):
1546 Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors
1547 for more detail.
1548 return_dict (`bool`, *optional*):
1549 Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.
1550 """
1551 output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
1552 output_hidden_states = (
1553 output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
1554 )
1555 return_dict = return_dict if return_dict is not None else self.config.use_return_dict
1556
1557 # retrieve input_ids and inputs_embeds
1558 if input_ids is not None and inputs_embeds is not None:
1559 raise ValueError("You cannot specify both input_ids and inputs_embeds at the same time")
1560 elif input_ids is not None:
1561 input = input_ids
1562 input_ids = input_ids.view(-1, input_ids.shape[-1])
1563 elif inputs_embeds is not None:
1564 input = inputs_embeds[:, :, -1]
1565 else:
1566 raise ValueError("You have to specify either input_ids or inputs_embeds")
1567
1568 if inputs_embeds is None:
1569 inputs_embeds = self.embed_tokens(input_ids)
1570
1571 embed_pos = self.embed_positions(input)
1572 embed_pos = embed_pos.to(inputs_embeds.device)
1573
1574 hidden_states = inputs_embeds + embed_pos
1575 hidden_states = self.layernorm_embedding(hidden_states)
1576 hidden_states = nn.functional.dropout(hidden_states, p=self.dropout, training=self.training)
1577
1578 # expand attention_mask
1579 if attention_mask is not None:
1580 if self._use_flash_attention_2:
1581 attention_mask = attention_mask if 0 in attention_mask else None
1582 elif self._use_sdpa and head_mask is None and not output_attentions:
1583 # output_attentions=True & head_mask can not be supported when using SDPA, fall back to
1584 # the manual implementation that requires a 4D causal mask in all cases.
1585 # [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]
1586 attention_mask = _prepare_4d_attention_mask_for_sdpa(attention_mask, inputs_embeds.dtype)
1587 else:
1588 # [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]
1589 attention_mask = _prepare_4d_attention_mask(attention_mask, inputs_embeds.dtype)
1590
1591 encoder_states = () if output_hidden_states else None
1592 all_attentions = () if output_attentions else None
1593
1594 # check if head_mask has a correct number of layers specified if desired
1595 if head_mask is not None:
1596 if head_mask.size()[0] != (len(self.layers)):
1597 raise ValueError(
1598 f"The head_mask should be specified for {len(self.layers)} layers, but it is for"
1599 f" {head_mask.size()[0]}."
1600 )
1601
1602 for idx, encoder_layer in enumerate(self.layers):
1603 if output_hidden_states:
1604 encoder_states = encoder_states + (hidden_states,)
1605 # add LayerDrop (see https://arxiv.org/abs/1909.11556 for description)
1606 to_drop = False
1607 if self.training:
1608 dropout_probability = torch.rand([])
1609 if dropout_probability < self.layerdrop: # skip the layer
1610 to_drop = True
1611
1612 if to_drop:
1613 layer_outputs = (None, None)
1614 else:
1615 if self.gradient_checkpointing and self.training:
1616 layer_outputs = self._gradient_checkpointing_func(
1617 encoder_layer.__call__,
1618 hidden_states,
1619 attention_mask,
1620 (head_mask[idx] if head_mask is not None else None),
1621 output_attentions,
1622 )
1623 else:
1624 layer_outputs = encoder_layer(
1625 hidden_states,
1626 attention_mask,
1627 layer_head_mask=(head_mask[idx] if head_mask is not None else None),
1628 output_attentions=output_attentions,
1629 )
1630
1631 hidden_states = layer_outputs[0]
1632
1633 if output_attentions:
1634 all_attentions = all_attentions + (layer_outputs[1],)
1635
1636 if output_hidden_states:
1637 encoder_states = encoder_states + (hidden_states,)
1638
1639 if not return_dict:
1640 return tuple(v for v in [hidden_states, encoder_states, all_attentions] if v is not None)
1641 return BaseModelOutput(
1642 last_hidden_state=hidden_states, hidden_states=encoder_states, attentions=all_attentions
1643 )
1644
1645
1646 class Florence2Decoder(Florence2LanguagePreTrainedModel):
1647 """
1648 Transformer decoder consisting of *config.decoder_layers* layers. Each layer is a [`Florence2DecoderLayer`]
1649
1650 Args:
1651 config: Florence2LanguageConfig
1652 embed_tokens (nn.Embedding): output embedding
1653 """
1654
1655 def __init__(self, config: Florence2LanguageConfig, embed_tokens: Optional[nn.Embedding] = None):
1656 super().__init__(config)
1657 self.dropout = config.dropout
1658 self.layerdrop = config.decoder_layerdrop
1659 self.padding_idx = config.pad_token_id
1660 self.max_target_positions = config.max_position_embeddings
1661 embed_scale = math.sqrt(config.d_model) if config.scale_embedding else 1.0
1662
1663 self.embed_tokens = Florence2ScaledWordEmbedding(
1664 config.vocab_size, config.d_model, self.padding_idx, embed_scale=embed_scale
1665 )
1666
1667 if embed_tokens is not None:
1668 self.embed_tokens.weight = embed_tokens.weight
1669
1670 self.embed_positions = Florence2LearnedPositionalEmbedding(
1671 config.max_position_embeddings,
1672 config.d_model,
1673 )
1674 self.layers = nn.ModuleList([Florence2DecoderLayer(config) for _ in range(config.decoder_layers)])
1675 self._use_flash_attention_2 = config._attn_implementation == "flash_attention_2"
1676 self._use_sdpa = config._attn_implementation == "sdpa"
1677
1678 self.layernorm_embedding = nn.LayerNorm(config.d_model)
1679
1680 self.gradient_checkpointing = False
1681 # Initialize weights and apply final processing
1682 self.post_init()
1683
1684 def get_input_embeddings(self):
1685 return self.embed_tokens
1686
1687 def set_input_embeddings(self, value):
1688 self.embed_tokens = value
1689
1690 def forward(
1691 self,
1692 input_ids: torch.LongTensor = None,
1693 attention_mask: Optional[torch.Tensor] = None,
1694 encoder_hidden_states: Optional[torch.FloatTensor] = None,
1695 encoder_attention_mask: Optional[torch.LongTensor] = None,
1696 head_mask: Optional[torch.Tensor] = None,
1697 cross_attn_head_mask: Optional[torch.Tensor] = None,
1698 past_key_values: Optional[List[torch.FloatTensor]] = None,
1699 inputs_embeds: Optional[torch.FloatTensor] = None,
1700 use_cache: Optional[bool] = None,
1701 output_attentions: Optional[bool] = None,
1702 output_hidden_states: Optional[bool] = None,
1703 return_dict: Optional[bool] = None,
1704 ) -> Union[Tuple, BaseModelOutputWithPastAndCrossAttentions]:
1705 r"""
1706 Args:
1707 input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`):
1708 Indices of input sequence tokens in the vocabulary. Padding will be ignored by default should you
1709 provide it.
1710
1711 Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and
1712 [`PreTrainedTokenizer.__call__`] for details.
1713
1714 [What are input IDs?](../glossary#input-ids)
1715 attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):
1716 Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:
1717
1718 - 1 for tokens that are **not masked**,
1719 - 0 for tokens that are **masked**.
1720
1721 [What are attention masks?](../glossary#attention-mask)
1722 encoder_hidden_states (`torch.FloatTensor` of shape `(batch_size, encoder_sequence_length, hidden_size)`, *optional*):
1723 Sequence of hidden-states at the output of the last layer of the encoder. Used in the cross-attention
1724 of the decoder.
1725 encoder_attention_mask (`torch.LongTensor` of shape `(batch_size, encoder_sequence_length)`, *optional*):
1726 Mask to avoid performing cross-attention on padding tokens indices of encoder input_ids. Mask values
1727 selected in `[0, 1]`:
1728
1729 - 1 for tokens that are **not masked**,
1730 - 0 for tokens that are **masked**.
1731
1732 [What are attention masks?](../glossary#attention-mask)
1733 head_mask (`torch.Tensor` of shape `(decoder_layers, decoder_attention_heads)`, *optional*):
1734 Mask to nullify selected heads of the attention modules. Mask values selected in `[0, 1]`:
1735
1736 - 1 indicates the head is **not masked**,
1737 - 0 indicates the head is **masked**.
1738
1739 cross_attn_head_mask (`torch.Tensor` of shape `(decoder_layers, decoder_attention_heads)`, *optional*):
1740 Mask to nullify selected heads of the cross-attention modules in the decoder to avoid performing
1741 cross-attention on hidden heads. Mask values selected in `[0, 1]`:
1742
1743 - 1 indicates the head is **not masked**,
1744 - 0 indicates the head is **masked**.
1745
1746 past_key_values (`tuple(tuple(torch.FloatTensor))`, *optional*, returned when `use_cache=True` is passed or when `config.use_cache=True`):
1747 Tuple of `tuple(torch.FloatTensor)` of length `config.n_layers`, with each tuple having 2 tensors of
1748 shape `(batch_size, num_heads, sequence_length, embed_size_per_head)`) and 2 additional tensors of
1749 shape `(batch_size, num_heads, encoder_sequence_length, embed_size_per_head)`.
1750
1751 Contains pre-computed hidden-states (key and values in the self-attention blocks and in the
1752 cross-attention blocks) that can be used (see `past_key_values` input) to speed up sequential decoding.
1753
1754 If `past_key_values` are used, the user can optionally input only the last `decoder_input_ids` (those
1755 that don't have their past key value states given to this model) of shape `(batch_size, 1)` instead of
1756 all `decoder_input_ids` of shape `(batch_size, sequence_length)`.
1757 inputs_embeds (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`, *optional*):
1758 Optionally, instead of passing `input_ids` you can choose to directly pass an embedded representation.
1759 This is useful if you want more control over how to convert `input_ids` indices into associated vectors
1760 than the model's internal embedding lookup matrix.
1761 output_attentions (`bool`, *optional*):
1762 Whether or not to return the attentions tensors of all attention layers. See `attentions` under
1763 returned tensors for more detail.
1764 output_hidden_states (`bool`, *optional*):
1765 Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors
1766 for more detail.
1767 return_dict (`bool`, *optional*):
1768 Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.
1769 """
1770 output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
1771 output_hidden_states = (
1772 output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
1773 )
1774 use_cache = use_cache if use_cache is not None else self.config.use_cache
1775 return_dict = return_dict if return_dict is not None else self.config.use_return_dict
1776
1777 # retrieve input_ids and inputs_embeds
1778 if input_ids is not None and inputs_embeds is not None:
1779 raise ValueError("You cannot specify both decoder_input_ids and decoder_inputs_embeds at the same time")
1780 elif input_ids is not None:
1781 input = input_ids
1782 input_shape = input.shape
1783 input_ids = input_ids.view(-1, input_shape[-1])
1784 elif inputs_embeds is not None:
1785 input_shape = inputs_embeds.size()[:-1]
1786 input = inputs_embeds[:, :, -1]
1787 else:
1788 raise ValueError("You have to specify either decoder_input_ids or decoder_inputs_embeds")
1789
1790 # past_key_values_length
1791 past_key_values_length = past_key_values[0][0].shape[2] if past_key_values is not None else 0
1792
1793 if inputs_embeds is None:
1794 inputs_embeds = self.embed_tokens(input)
1795
1796 if self._use_flash_attention_2:
1797 # 2d mask is passed through the layers
1798 attention_mask = attention_mask if (attention_mask is not None and 0 in attention_mask) else None
1799 elif self._use_sdpa and not output_attentions and cross_attn_head_mask is None:
1800 # output_attentions=True & cross_attn_head_mask can not be supported when using SDPA, and we fall back on
1801 # the manual implementation that requires a 4D causal mask in all cases.
1802 attention_mask = _prepare_4d_causal_attention_mask_for_sdpa(
1803 attention_mask,
1804 input_shape,
1805 inputs_embeds,
1806 past_key_values_length,
1807 )
1808 else:
1809 # 4d mask is passed through the layers
1810 attention_mask = _prepare_4d_causal_attention_mask(
1811 attention_mask, input_shape, inputs_embeds, past_key_values_length
1812 )
1813
1814 # expand encoder attention mask
1815 if encoder_hidden_states is not None and encoder_attention_mask is not None:
1816 if self._use_flash_attention_2:
1817 encoder_attention_mask = encoder_attention_mask if 0 in encoder_attention_mask else None
1818 elif self._use_sdpa and cross_attn_head_mask is None and not output_attentions:
1819 # output_attentions=True & cross_attn_head_mask can not be supported when using SDPA, and we fall back on
1820 # the manual implementation that requires a 4D causal mask in all cases.
1821 # [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]
1822 encoder_attention_mask = _prepare_4d_attention_mask_for_sdpa(
1823 encoder_attention_mask,
1824 inputs_embeds.dtype,
1825 tgt_len=input_shape[-1],
1826 )
1827 else:
1828 # [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]
1829 encoder_attention_mask = _prepare_4d_attention_mask(
1830 encoder_attention_mask, inputs_embeds.dtype, tgt_len=input_shape[-1]
1831 )
1832
1833 # embed positions
1834 positions = self.embed_positions(input, past_key_values_length)
1835 positions = positions.to(inputs_embeds.device)
1836
1837 hidden_states = inputs_embeds + positions
1838 hidden_states = self.layernorm_embedding(hidden_states)
1839
1840 hidden_states = nn.functional.dropout(hidden_states, p=self.dropout, training=self.training)
1841
1842 if self.gradient_checkpointing and self.training:
1843 if use_cache:
1844 logger.warning_once(
1845 "`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`..."
1846 )
1847 use_cache = False
1848
1849 # decoder layers
1850 all_hidden_states = () if output_hidden_states else None
1851 all_self_attns = () if output_attentions else None
1852 all_cross_attentions = () if (output_attentions and encoder_hidden_states is not None) else None
1853 next_decoder_cache = () if use_cache else None
1854
1855 # check if head_mask/cross_attn_head_mask has a correct number of layers specified if desired
1856 for attn_mask, mask_name in zip([head_mask, cross_attn_head_mask], ["head_mask", "cross_attn_head_mask"]):
1857 if attn_mask is not None:
1858 if attn_mask.size()[0] != (len(self.layers)):
1859 raise ValueError(
1860 f"The `{mask_name}` should be specified for {len(self.layers)} layers, but it is for"
1861 f" {head_mask.size()[0]}."
1862 )
1863
1864 for idx, decoder_layer in enumerate(self.layers):
1865 # add LayerDrop (see https://arxiv.org/abs/1909.11556 for description)
1866 if output_hidden_states:
1867 all_hidden_states += (hidden_states,)
1868 if self.training:
1869 dropout_probability = torch.rand([])
1870 if dropout_probability < self.layerdrop:
1871 continue
1872
1873 past_key_value = past_key_values[idx] if past_key_values is not None else None
1874
1875 if self.gradient_checkpointing and self.training:
1876 layer_outputs = self._gradient_checkpointing_func(
1877 decoder_layer.__call__,
1878 hidden_states,
1879 attention_mask,
1880 encoder_hidden_states,
1881 encoder_attention_mask,
1882 head_mask[idx] if head_mask is not None else None,
1883 cross_attn_head_mask[idx] if cross_attn_head_mask is not None else None,
1884 None,
1885 output_attentions,
1886 use_cache,
1887 )
1888 else:
1889 layer_outputs = decoder_layer(
1890 hidden_states,
1891 attention_mask=attention_mask,
1892 encoder_hidden_states=encoder_hidden_states,
1893 encoder_attention_mask=encoder_attention_mask,
1894 layer_head_mask=(head_mask[idx] if head_mask is not None else None),
1895 cross_attn_layer_head_mask=(
1896 cross_attn_head_mask[idx] if cross_attn_head_mask is not None else None
1897 ),
1898 past_key_value=past_key_value,
1899 output_attentions=output_attentions,
1900 use_cache=use_cache,
1901 )
1902 hidden_states = layer_outputs[0]
1903
1904 if use_cache:
1905 next_decoder_cache += (layer_outputs[3 if output_attentions else 1],)
1906
1907 if output_attentions:
1908 all_self_attns += (layer_outputs[1],)
1909
1910 if encoder_hidden_states is not None:
1911 all_cross_attentions += (layer_outputs[2],)
1912
1913 # add hidden states from the last decoder layer
1914 if output_hidden_states:
1915 all_hidden_states += (hidden_states,)
1916
1917 next_cache = next_decoder_cache if use_cache else None
1918 if not return_dict:
1919 return tuple(
1920 v
1921 for v in [hidden_states, next_cache, all_hidden_states, all_self_attns, all_cross_attentions]
1922 if v is not None
1923 )
1924 return BaseModelOutputWithPastAndCrossAttentions(
1925 last_hidden_state=hidden_states,
1926 past_key_values=next_cache,
1927 hidden_states=all_hidden_states,
1928 attentions=all_self_attns,
1929 cross_attentions=all_cross_attentions,
1930 )
1931
1932
1933 class Florence2LanguageModel(Florence2LanguagePreTrainedModel):
1934 _tied_weights_keys = ["encoder.embed_tokens.weight", "decoder.embed_tokens.weight"]
1935
1936 def __init__(self, config: Florence2LanguageConfig):
1937 super().__init__(config)
1938
1939 padding_idx, vocab_size = config.pad_token_id, config.vocab_size
1940 self.shared = nn.Embedding(vocab_size, config.d_model, padding_idx)
1941
1942 self.encoder = Florence2Encoder(config, self.shared)
1943 self.decoder = Florence2Decoder(config, self.shared)
1944
1945 # Initialize weights and apply final processing
1946 self.post_init()
1947
1948 def _tie_weights(self):
1949 if self.config.tie_word_embeddings:
1950 self._tie_or_clone_weights(self.encoder.embed_tokens, self.shared)
1951 self._tie_or_clone_weights(self.decoder.embed_tokens, self.shared)
1952
1953 def get_input_embeddings(self):
1954 return self.shared
1955
1956 def set_input_embeddings(self, value):
1957 self.shared = value
1958 self.encoder.embed_tokens = self.shared
1959 self.decoder.embed_tokens = self.shared
1960
1961 def get_encoder(self):
1962 return self.encoder
1963
1964 def get_decoder(self):
1965 return self.decoder
1966
1967 def forward(
1968 self,
1969 input_ids: torch.LongTensor = None,
1970 attention_mask: Optional[torch.Tensor] = None,
1971 decoder_input_ids: Optional[torch.LongTensor] = None,
1972 decoder_attention_mask: Optional[torch.LongTensor] = None,
1973 head_mask: Optional[torch.Tensor] = None,
1974 decoder_head_mask: Optional[torch.Tensor] = None,
1975 cross_attn_head_mask: Optional[torch.Tensor] = None,
1976 encoder_outputs: Optional[List[torch.FloatTensor]] = None,
1977 past_key_values: Optional[List[torch.FloatTensor]] = None,
1978 inputs_embeds: Optional[torch.FloatTensor] = None,
1979 decoder_inputs_embeds: Optional[torch.FloatTensor] = None,
1980 use_cache: Optional[bool] = None,
1981 output_attentions: Optional[bool] = None,
1982 output_hidden_states: Optional[bool] = None,
1983 return_dict: Optional[bool] = None,
1984 ) -> Union[Tuple, Seq2SeqModelOutput]:
1985 # different to other models, Florence2 automatically creates decoder_input_ids from
1986 # input_ids if no decoder_input_ids are provided
1987 if decoder_input_ids is None and decoder_inputs_embeds is None:
1988 if input_ids is None:
1989 raise ValueError(
1990 "If no `decoder_input_ids` or `decoder_inputs_embeds` are "
1991 "passed, `input_ids` cannot be `None`. Please pass either "
1992 "`input_ids` or `decoder_input_ids` or `decoder_inputs_embeds`."
1993 )
1994
1995 decoder_input_ids = shift_tokens_right(
1996 input_ids, self.config.pad_token_id, self.config.decoder_start_token_id
1997 )
1998
1999 output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
2000 output_hidden_states = (
2001 output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
2002 )
2003 use_cache = use_cache if use_cache is not None else self.config.use_cache
2004 return_dict = return_dict if return_dict is not None else self.config.use_return_dict
2005
2006 if encoder_outputs is None:
2007 encoder_outputs = self.encoder(
2008 input_ids=input_ids,
2009 attention_mask=attention_mask,
2010 head_mask=head_mask,
2011 inputs_embeds=inputs_embeds,
2012 output_attentions=output_attentions,
2013 output_hidden_states=output_hidden_states,
2014 return_dict=return_dict,
2015 )
2016 # If the user passed a tuple for encoder_outputs, we wrap it in a BaseModelOutput when return_dict=True
2017 elif return_dict and not isinstance(encoder_outputs, BaseModelOutput):
2018 encoder_outputs = BaseModelOutput(
2019 last_hidden_state=encoder_outputs[0],
2020 hidden_states=encoder_outputs[1] if len(encoder_outputs) > 1 else None,
2021 attentions=encoder_outputs[2] if len(encoder_outputs) > 2 else None,
2022 )
2023
2024 # decoder outputs consists of (dec_features, past_key_value, dec_hidden, dec_attn)
2025 decoder_outputs = self.decoder(
2026 input_ids=decoder_input_ids,
2027 attention_mask=decoder_attention_mask,
2028 encoder_hidden_states=encoder_outputs[0],
2029 encoder_attention_mask=attention_mask,
2030 head_mask=decoder_head_mask,
2031 cross_attn_head_mask=cross_attn_head_mask,
2032 past_key_values=past_key_values,
2033 inputs_embeds=decoder_inputs_embeds,
2034 use_cache=use_cache,
2035 output_attentions=output_attentions,
2036 output_hidden_states=output_hidden_states,
2037 return_dict=return_dict,
2038 )
2039
2040 if not return_dict:
2041 return decoder_outputs + encoder_outputs
2042
2043 return Seq2SeqModelOutput(
2044 last_hidden_state=decoder_outputs.last_hidden_state,
2045 past_key_values=decoder_outputs.past_key_values,
2046 decoder_hidden_states=decoder_outputs.hidden_states,
2047 decoder_attentions=decoder_outputs.attentions,
2048 cross_attentions=decoder_outputs.cross_attentions,
2049 encoder_last_hidden_state=encoder_outputs.last_hidden_state,
2050 encoder_hidden_states=encoder_outputs.hidden_states,
2051 encoder_attentions=encoder_outputs.attentions,
2052 )
2053
2054
2055 class Florence2LanguageForConditionalGeneration(Florence2LanguagePreTrainedModel, GenerationMixin):
2056 base_model_prefix = "model"
2057 _tied_weights_keys = ["encoder.embed_tokens.weight", "decoder.embed_tokens.weight", "lm_head.weight"]
2058 _keys_to_ignore_on_load_missing = ["final_logits_bias"]
2059
2060 def __init__(self, config: Florence2LanguageConfig):
2061 super().__init__(config)
2062 self.model = Florence2LanguageModel(config)
2063 self.register_buffer("final_logits_bias", torch.zeros((1, self.model.shared.num_embeddings)))
2064 self.lm_head = nn.Linear(config.d_model, self.model.shared.num_embeddings, bias=False)
2065
2066 # Initialize weights and apply final processing
2067 self.post_init()
2068
2069 def _tie_weights(self):
2070 if self.config.tie_word_embeddings:
2071 self._tie_or_clone_weights(self.model.encoder.embed_tokens, self.model.shared)
2072 self._tie_or_clone_weights(self.model.decoder.embed_tokens, self.model.shared)
2073 self._tie_or_clone_weights(self.lm_head, self.model.shared)
2074
2075 def get_encoder(self):
2076 return self.model.get_encoder()
2077
2078 def get_decoder(self):
2079 return self.model.get_decoder()
2080
2081 def resize_token_embeddings(self, new_num_tokens: int, pad_to_multiple_of: Optional[int] = None, **kwargs) -> nn.Embedding:
2082 new_embeddings = super().resize_token_embeddings(new_num_tokens, pad_to_multiple_of, **kwargs)
2083 self._resize_final_logits_bias(new_embeddings.weight.shape[0])
2084 return new_embeddings
2085
2086 def _resize_final_logits_bias(self, new_num_tokens: int) -> None:
2087 old_num_tokens = self.final_logits_bias.shape[-1]
2088 if new_num_tokens <= old_num_tokens:
2089 new_bias = self.final_logits_bias[:, :new_num_tokens]
2090 else:
2091 extra_bias = torch.zeros((1, new_num_tokens - old_num_tokens), device=self.final_logits_bias.device)
2092 new_bias = torch.cat([self.final_logits_bias, extra_bias], dim=1)
2093 self.register_buffer("final_logits_bias", new_bias)
2094
2095 def get_output_embeddings(self):
2096 return self.lm_head
2097
2098 def set_output_embeddings(self, new_embeddings):
2099 self.lm_head = new_embeddings
2100
2101 def forward(
2102 self,
2103 input_ids: torch.LongTensor = None,
2104 attention_mask: Optional[torch.Tensor] = None,
2105 decoder_input_ids: Optional[torch.LongTensor] = None,
2106 decoder_attention_mask: Optional[torch.LongTensor] = None,
2107 head_mask: Optional[torch.Tensor] = None,
2108 decoder_head_mask: Optional[torch.Tensor] = None,
2109 cross_attn_head_mask: Optional[torch.Tensor] = None,
2110 encoder_outputs: Optional[List[torch.FloatTensor]] = None,
2111 past_key_values: Optional[List[torch.FloatTensor]] = None,
2112 inputs_embeds: Optional[torch.FloatTensor] = None,
2113 decoder_inputs_embeds: Optional[torch.FloatTensor] = None,
2114 labels: Optional[torch.LongTensor] = None,
2115 use_cache: Optional[bool] = None,
2116 output_attentions: Optional[bool] = None,
2117 output_hidden_states: Optional[bool] = None,
2118 return_dict: Optional[bool] = None,
2119 ) -> Union[Tuple, Seq2SeqLMOutput]:
2120 r"""
2121 labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
2122 Labels for computing the masked language modeling loss. Indices should either be in `[0, ...,
2123 config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored
2124 (masked), the loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`.
2125
2126 Returns:
2127 """
2128 return_dict = return_dict if return_dict is not None else self.config.use_return_dict
2129
2130 if labels is not None:
2131 if use_cache:
2132 logger.warning("The `use_cache` argument is changed to `False` since `labels` is provided.")
2133 use_cache = False
2134 if decoder_input_ids is None and decoder_inputs_embeds is None:
2135 decoder_input_ids = shift_tokens_right(
2136 labels, self.config.pad_token_id, self.config.decoder_start_token_id
2137 )
2138
2139 outputs = self.model(
2140 input_ids,
2141 attention_mask=attention_mask,
2142 decoder_input_ids=decoder_input_ids,
2143 encoder_outputs=encoder_outputs,
2144 decoder_attention_mask=decoder_attention_mask,
2145 head_mask=head_mask,
2146 decoder_head_mask=decoder_head_mask,
2147 cross_attn_head_mask=cross_attn_head_mask,
2148 past_key_values=past_key_values,
2149 inputs_embeds=inputs_embeds,
2150 decoder_inputs_embeds=decoder_inputs_embeds,
2151 use_cache=use_cache,
2152 output_attentions=output_attentions,
2153 output_hidden_states=output_hidden_states,
2154 return_dict=return_dict,
2155 )
2156
2157 lm_logits = self.lm_head(outputs[0])
2158 lm_logits = lm_logits + self.final_logits_bias.to(lm_logits.device)
2159
2160 masked_lm_loss = None
2161 if labels is not None:
2162 labels = labels.to(lm_logits.device)
2163 loss_fct = CrossEntropyLoss()
2164 masked_lm_loss = loss_fct(lm_logits.view(-1, self.config.vocab_size), labels.view(-1))
2165
2166 if not return_dict:
2167 output = (lm_logits,) + outputs[1:]
2168 return ((masked_lm_loss,) + output) if masked_lm_loss is not None else output
2169
2170 return Seq2SeqLMOutput(
2171 loss=masked_lm_loss,
2172 logits=lm_logits,
2173 past_key_values=outputs.past_key_values,
2174 decoder_hidden_states=outputs.decoder_hidden_states,
2175 decoder_attentions=outputs.decoder_attentions,
2176 cross_attentions=outputs.cross_attentions,
2177 encoder_last_hidden_state=outputs.encoder_last_hidden_state,
2178 encoder_hidden_states=outputs.encoder_hidden_states,
2179 encoder_attentions=outputs.encoder_attentions,
2180 )
2181
2182 def prepare_inputs_for_generation(
2183 self,
2184 decoder_input_ids,
2185 past_key_values=None,
2186 attention_mask=None,
2187 decoder_attention_mask=None,
2188 head_mask=None,
2189 decoder_head_mask=None,
2190 cross_attn_head_mask=None,
2191 use_cache=None,
2192 encoder_outputs=None,
2193 **kwargs,
2194 ):
2195 # cut decoder_input_ids if past_key_values is used
2196 if past_key_values is not None:
2197 past_length = past_key_values[0][0].shape[2]
2198
2199 # Some generation methods already pass only the last input ID
2200 if decoder_input_ids.shape[1] > past_length:
2201 remove_prefix_length = past_length
2202 else:
2203 # Default to old behavior: keep only final ID
2204 remove_prefix_length = decoder_input_ids.shape[1] - 1
2205
2206 decoder_input_ids = decoder_input_ids[:, remove_prefix_length:]
2207
2208 return {
2209 "input_ids": None, # encoder_outputs is defined. input_ids not needed
2210 "encoder_outputs": encoder_outputs,
2211 "past_key_values": past_key_values,
2212 "decoder_input_ids": decoder_input_ids,
2213 "attention_mask": attention_mask,
2214 "decoder_attention_mask": decoder_attention_mask,
2215 "head_mask": head_mask,
2216 "decoder_head_mask": decoder_head_mask,
2217 "cross_attn_head_mask": cross_attn_head_mask,
2218 "use_cache": use_cache, # change this to avoid caching (presumably for debugging)
2219 }
2220
2221 def prepare_decoder_input_ids_from_labels(self, labels: torch.Tensor):
2222 return shift_tokens_right(labels, self.config.pad_token_id, self.config.decoder_start_token_id)
2223
2224 @staticmethod
2225 def _reorder_cache(past_key_values, beam_idx):
2226 reordered_past = ()
2227 for layer_past in past_key_values:
2228 # cached cross_attention states don't have to be reordered -> they are always the same
2229 reordered_past += (
2230 tuple(past_state.index_select(0, beam_idx.to(past_state.device)) for past_state in layer_past[:2])
2231 + layer_past[2:],
2232 )
2233 return reordered_past
2234
2235 @dataclass
2236 class Florence2Seq2SeqLMOutput(ModelOutput):
2237 """
2238 Base class for Florence-2 model's outputs that also contains : pre-computed hidden states that can speed up sequential
2239 decoding.
2240
2241 Args:
2242 loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `labels` is provided):
2243 Language modeling loss.
2244 logits (`torch.FloatTensor` of shape `(batch_size, sequence_length, config.vocab_size)`):
2245 Prediction scores of the language modeling head (scores for each vocabulary token before SoftMax).
2246 last_hidden_state (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`):
2247 Sequence of hidden-states at the output of the last layer of the decoder of the model.
2248
2249 If `past_key_values` is used only the last hidden-state of the sequences of shape `(batch_size, 1,
2250 hidden_size)` is output.
2251 past_key_values (`tuple(tuple(torch.FloatTensor))`, *optional*, returned when `use_cache=True` is passed or when `config.use_cache=True`):
2252 Tuple of `tuple(torch.FloatTensor)` of length `config.n_layers`, with each tuple having 2 tensors of shape
2253 `(batch_size, num_heads, sequence_length, embed_size_per_head)`) and 2 additional tensors of shape
2254 `(batch_size, num_heads, encoder_sequence_length, embed_size_per_head)`.
2255
2256 Contains pre-computed hidden-states (key and values in the self-attention blocks and in the cross-attention
2257 blocks) that can be used (see `past_key_values` input) to speed up sequential decoding.
2258 decoder_hidden_states (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`):
2259 Tuple of `torch.FloatTensor` (one for the output of the embeddings, if the model has an embedding layer, +
2260 one for the output of each layer) of shape `(batch_size, sequence_length, hidden_size)`.
2261
2262 Hidden-states of the decoder at the output of each layer plus the optional initial embedding outputs.
2263 decoder_attentions (`tuple(torch.FloatTensor)`, *optional*, returned when `output_attentions=True` is passed or when `config.output_attentions=True`):
2264 Tuple of `torch.FloatTensor` (one for each layer) of shape `(batch_size, num_heads, sequence_length,
2265 sequence_length)`.
2266
2267 Attentions weights of the decoder, after the attention softmax, used to compute the weighted average in the
2268 self-attention heads.
2269 cross_attentions (`tuple(torch.FloatTensor)`, *optional*, returned when `output_attentions=True` is passed or when `config.output_attentions=True`):
2270 Tuple of `torch.FloatTensor` (one for each layer) of shape `(batch_size, num_heads, sequence_length,
2271 sequence_length)`.
2272
2273 Attentions weights of the decoder's cross-attention layer, after the attention softmax, used to compute the
2274 weighted average in the cross-attention heads.
2275 encoder_last_hidden_state (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`, *optional*):
2276 Sequence of hidden-states at the output of the last layer of the encoder of the model.
2277 encoder_hidden_states (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`):
2278 Tuple of `torch.FloatTensor` (one for the output of the embeddings, if the model has an embedding layer, +
2279 one for the output of each layer) of shape `(batch_size, sequence_length, hidden_size)`.
2280
2281 Hidden-states of the encoder at the output of each layer plus the optional initial embedding outputs.
2282 encoder_attentions (`tuple(torch.FloatTensor)`, *optional*, returned when `output_attentions=True` is passed or when `config.output_attentions=True`):
2283 Tuple of `torch.FloatTensor` (one for each layer) of shape `(batch_size, num_heads, sequence_length,
2284 sequence_length)`.
2285
2286 Attentions weights of the encoder, after the attention softmax, used to compute the weighted average in the
2287 self-attention heads.
2288 image_hidden_states (`tuple(torch.FloatTensor)`, *optional*):
2289 Tuple of `torch.FloatTensor` (one for the output of the image embeddings, `(batch_size,
2290 num_image_tokens, hidden_size)`.
2291
2292 image_hidden_states of the model produced by the vision encoder
2293 """
2294 loss: Optional[torch.FloatTensor] = None
2295 logits: torch.FloatTensor = None
2296 last_hidden_state: torch.FloatTensor = None
2297 past_key_values: Optional[Tuple[Tuple[torch.FloatTensor]]] = None
2298 decoder_hidden_states: Optional[Tuple[torch.FloatTensor, ...]] = None
2299 decoder_attentions: Optional[Tuple[torch.FloatTensor, ...]] = None
2300 cross_attentions: Optional[Tuple[torch.FloatTensor, ...]] = None
2301 encoder_last_hidden_state: Optional[torch.FloatTensor] = None
2302 encoder_hidden_states: Optional[Tuple[torch.FloatTensor, ...]] = None
2303 encoder_attentions: Optional[Tuple[torch.FloatTensor, ...]] = None
2304 image_hidden_states: Optional[Tuple[torch.FloatTensor, ...]] = None
2305
2306
2307 FLORENCE2_START_DOCSTRING = r"""
2308 This model inherits from [`PreTrainedModel`]. Check the superclass documentation for the generic methods the
2309 library implements for all its model (such as downloading or saving, resizing the input embeddings, pruning heads
2310 etc.)
2311
2312 This model is also a PyTorch [torch.nn.Module](https://pytorch.org/docs/stable/nn.html#torch.nn.Module) subclass.
2313 Use it as a regular PyTorch Module and refer to the PyTorch documentation for all matter related to general usage
2314 and behavior.
2315
2316 Parameters:
2317 config ([`Florence2Config`] or [`Florence2VisionConfig`]):
2318 Model configuration class with all the parameters of the model. Initializing with a config file does not
2319 load the weights associated with the model, only the configuration. Check out the
2320 [`~PreTrainedModel.from_pretrained`] method to load the model weights.
2321 """
2322
2323
2324 @add_start_docstrings(
2325 "The bare Florence-2 Model outputting raw hidden-states without any specific head on top.",
2326 FLORENCE2_START_DOCSTRING,
2327 )
2328 class Florence2PreTrainedModel(PreTrainedModel):
2329 config_class = Florence2Config
2330 base_model_prefix = "model"
2331 supports_gradient_checkpointing = True
2332 _skip_keys_device_placement = "past_key_values"
2333
2334 @property
2335 def _supports_flash_attn_2(self):
2336 """
2337 Retrieve language_model's attribute to check whether the model supports
2338 Flash Attention 2 or not.
2339 """
2340 return self.language_model._supports_flash_attn_2
2341
2342 @property
2343 def _supports_sdpa(self):
2344 """
2345 Retrieve language_model's attribute to check whether the model supports
2346 SDPA or not.
2347 """
2348 return self.language_model._supports_sdpa
2349
2350
2351 FLORENCE2_INPUTS_DOCSTRING = r"""
2352 Args:
2353 input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`):
2354 Indices of input sequence tokens in the vocabulary. Padding will be ignored by default should you provide
2355 it.
2356
2357 Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and
2358 [`PreTrainedTokenizer.__call__`] for details.
2359
2360 [What are input IDs?](../glossary#input-ids)
2361 pixel_values (`torch.FloatTensor` of shape `(batch_size, num_channels, image_size, image_size)):
2362 The tensors corresponding to the input images. Pixel values can be obtained using
2363 [`AutoImageProcessor`]. See [`CLIPImageProcessor.__call__`] for details ([]`Florence2Processor`] uses
2364 [`CLIPImageProcessor`] for processing images).
2365 attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):
2366 Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:
2367
2368 - 1 for tokens that are **not masked**,
2369 - 0 for tokens that are **masked**.
2370
2371 [What are attention masks?](../glossary#attention-mask)
2372
2373 Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and
2374 [`PreTrainedTokenizer.__call__`] for details.
2375
2376 If `past_key_values` is used, optionally only the last `decoder_input_ids` have to be input (see
2377 `past_key_values`).
2378
2379 If you want to change padding behavior, you should read [`modeling_opt._prepare_decoder_attention_mask`]
2380 and modify to your needs. See diagram 1 in [the paper](https://arxiv.org/abs/1910.13461) for more
2381 information on the default strategy.
2382
2383 - 1 indicates the head is **not masked**,
2384 - 0 indicates the head is **masked**.
2385 position_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
2386 Indices of positions of each input sequence tokens in the position embeddings. Selected in the range `[0,
2387 config.n_positions - 1]`. [What are position IDs?](../glossary#position-ids)
2388 past_key_values (`tuple(tuple(torch.FloatTensor))`, *optional*, returned when `use_cache=True` is passed or when `config.use_cache=True`):
2389 Tuple of `tuple(torch.FloatTensor)` of length `config.n_layers`, with each tuple having 2 tensors of shape
2390 `(batch_size, num_heads, sequence_length, embed_size_per_head)`) and 2 additional tensors of shape
2391 `(batch_size, num_heads, encoder_sequence_length, embed_size_per_head)`.
2392
2393 Contains pre-computed hidden-states (key and values in the self-attention blocks and in the cross-attention
2394 blocks) that can be used (see `past_key_values` input) to speed up sequential decoding.
2395
2396 If `past_key_values` are used, the user can optionally input only the last `decoder_input_ids` (those that
2397 don't have their past key value states given to this model) of shape `(batch_size, 1)` instead of all
2398 `decoder_input_ids` of shape `(batch_size, sequence_length)`.
2399 inputs_embeds (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`, *optional*):
2400 Optionally, instead of passing `input_ids` you can choose to directly pass an embedded representation. This
2401 is useful if you want more control over how to convert `input_ids` indices into associated vectors than the
2402 model's internal embedding lookup matrix.
2403 use_cache (`bool`, *optional*):
2404 If set to `True`, `past_key_values` key value states are returned and can be used to speed up decoding (see
2405 `past_key_values`).
2406 output_attentions (`bool`, *optional*):
2407 Whether or not to return the attentions tensors of all attention layers. See `attentions` under returned
2408 tensors for more detail.
2409 output_hidden_states (`bool`, *optional*):
2410 Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors for
2411 more detail.
2412 return_dict (`bool`, *optional*):
2413 Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.
2414 """
2415
2416 @add_start_docstrings(
2417 """The FLORENCE2 vision model without any head""",
2418 FLORENCE2_START_DOCSTRING,
2419 )
2420 class Florence2VisionModel(Florence2PreTrainedModel):
2421 def __init__(self, config: Florence2VisionConfig):
2422 super().__init__(config)
2423 assert config.model_type == 'davit', 'only DaViT is supported for now'
2424 self.vision_tower = DaViT.from_config(config=config)
2425
2426 self.post_init()
2427
2428 def forward(self, pixel_values):
2429 if len(pixel_values.shape) == 4:
2430 x = self.vision_tower.forward_features_unpool(pixel_values)
2431 else:
2432 raise ValueError(f'invalid image shape {pixel_values.shape}')
2433 return x
2434
2435
2436 @add_start_docstrings(
2437 """The FLORENCE2 vision model with projection layer""",
2438 FLORENCE2_START_DOCSTRING,
2439 )
2440 class Florence2VisionModelWithProjection(Florence2PreTrainedModel):
2441 def __init__(self, config: Florence2VisionConfig):
2442 super().__init__(config)
2443 assert config.model_type == 'davit', 'only DaViT is supported for now'
2444 self.vision_tower = DaViT.from_config(config=config)
2445
2446 self._build_image_projection_layers(config)
2447
2448 self.post_init()
2449
2450 def _build_image_projection_layers(self, config):
2451 image_dim_out = config.dim_embed[-1]
2452 dim_projection = config.projection_dim
2453 self.image_projection = nn.Parameter(
2454 torch.empty(image_dim_out, dim_projection)
2455 )
2456 self.image_proj_norm = nn.LayerNorm(dim_projection)
2457 image_pos_embed_config = config.image_pos_embed
2458 if image_pos_embed_config['type'] == 'learned_abs_2d':
2459 self.image_pos_embed = LearnedAbsolutePositionEmbedding2D(
2460 embedding_dim=image_dim_out,
2461 num_pos=image_pos_embed_config['max_pos_embeddings']
2462 )
2463 else:
2464 raise NotImplementedError('Not implemented yet')
2465
2466 self.image_feature_source = config.image_feature_source
2467
2468 # temporal embedding
2469 visual_temporal_embedding_config = config.visual_temporal_embedding
2470 if visual_temporal_embedding_config['type'] == 'COSINE':
2471 self.visual_temporal_embed = PositionalEmbeddingCosine1D(
2472 embed_dim=image_dim_out,
2473 max_seq_len=visual_temporal_embedding_config['max_temporal_embeddings']
2474 )
2475 else:
2476 raise NotImplementedError('Not implemented yet')
2477
2478 def forward(self, pixel_values):
2479 if len(pixel_values.shape) == 4:
2480 batch_size, C, H, W = pixel_values.shape
2481 T = 1
2482 x = self.vision_tower.forward_features_unpool(pixel_values)
2483 else:
2484 raise ValueError(f'invalid image shape {pixel_values.shape}')
2485
2486 if self.image_pos_embed is not None:
2487 x = x.view(batch_size * T, -1, x.shape[-1])
2488 num_tokens = x.shape[-2]
2489 h, w = int(num_tokens ** 0.5), int(num_tokens ** 0.5)
2490 assert h * w == num_tokens, 'only support square feature maps for now'
2491 x = x.view(batch_size * T, h, w, x.shape[-1])
2492 pos_embed = self.image_pos_embed(x)
2493 x = x + pos_embed
2494 x = x.view(batch_size, T * h*w, x.shape[-1])
2495
2496 if self.visual_temporal_embed is not None:
2497 visual_temporal_embed = self.visual_temporal_embed(x.view(batch_size, T, -1, x.shape[-1])[:, :, 0])
2498 x = x.view(batch_size, T, -1, x.shape[-1]) + visual_temporal_embed.view(1, T, 1, x.shape[-1])
2499
2500 x_feat_dict = {}
2501
2502 spatial_avg_pool_x = x.view(batch_size, T, -1, x.shape[-1]).mean(dim=2)
2503 x_feat_dict['spatial_avg_pool'] = spatial_avg_pool_x
2504
2505 temporal_avg_pool_x = x.view(batch_size, T, -1, x.shape[-1]).mean(dim=1)
2506 x_feat_dict['temporal_avg_pool'] = temporal_avg_pool_x
2507
2508 x = x.view(batch_size, T, -1, x.shape[-1])[:, -1]
2509 x_feat_dict['last_frame'] = x
2510
2511 new_x = []
2512 for _image_feature_source in self.image_feature_source:
2513 if _image_feature_source not in x_feat_dict:
2514 raise ValueError('invalid image feature source: {}'.format(_image_feature_source))
2515 new_x.append(x_feat_dict[_image_feature_source])
2516
2517 x = torch.cat(new_x, dim=1)
2518
2519 x = x @ self.image_projection
2520 x = self.image_proj_norm(x)
2521
2522
2523 return x
2524
2525
2526
2527 @add_start_docstrings(
2528 """The FLORENCE2 model which consists of a vision backbone and a language model.""",
2529 FLORENCE2_START_DOCSTRING,
2530 )
2531 class Florence2ForConditionalGeneration(Florence2PreTrainedModel):
2532 _tied_weights_keys = ["language_model.encoder.embed_tokens.weight", "language_model.decoder.embed_tokens.weight", "language_model.lm_head.weight"]
2533
2534 def __init__(self, config: Florence2Config):
2535 super().__init__(config)
2536 assert config.vision_config.model_type == 'davit', 'only DaViT is supported for now'
2537 self.vision_tower = DaViT.from_config(config=config.vision_config)
2538 # remove unused layers
2539 del self.vision_tower.head
2540 del self.vision_tower.norms
2541
2542 self.vocab_size = config.vocab_size
2543 self._attn_implementation = config._attn_implementation
2544 self._build_image_projection_layers(config)
2545
2546 language_model = Florence2LanguageForConditionalGeneration(config=config.text_config)
2547
2548 self.language_model = language_model
2549
2550 self.pad_token_id = self.config.pad_token_id if self.config.pad_token_id is not None else -1
2551 self.post_init()
2552
2553 def _build_image_projection_layers(self, config):
2554 image_dim_out = config.vision_config.dim_embed[-1]
2555 dim_projection = config.vision_config.projection_dim
2556 self.image_projection = nn.Parameter(
2557 torch.empty(image_dim_out, dim_projection)
2558 )
2559 self.image_proj_norm = nn.LayerNorm(dim_projection)
2560 image_pos_embed_config = config.vision_config.image_pos_embed
2561 if image_pos_embed_config['type'] == 'learned_abs_2d':
2562 self.image_pos_embed = LearnedAbsolutePositionEmbedding2D(
2563 embedding_dim=image_dim_out,
2564 num_pos=image_pos_embed_config['max_pos_embeddings']
2565 )
2566 else:
2567 raise NotImplementedError('Not implemented yet')
2568
2569 self.image_feature_source = config.vision_config.image_feature_source
2570
2571 # temporal embedding
2572 visual_temporal_embedding_config = config.vision_config.visual_temporal_embedding
2573 if visual_temporal_embedding_config['type'] == 'COSINE':
2574 self.visual_temporal_embed = PositionalEmbeddingCosine1D(
2575 embed_dim=image_dim_out,
2576 max_seq_len=visual_temporal_embedding_config['max_temporal_embeddings']
2577 )
2578 else:
2579 raise NotImplementedError('Not implemented yet')
2580
2581 def get_encoder(self):
2582 return self.language_model.get_encoder()
2583
2584 def get_decoder(self):
2585 return self.language_model.get_decoder()
2586
2587 def get_input_embeddings(self):
2588 return self.language_model.get_input_embeddings()
2589
2590 def resize_token_embeddings(self, new_num_tokens: Optional[int] = None, pad_to_multiple_of=None, **kwargs) -> nn.Embedding:
2591 model_embeds = self.language_model.resize_token_embeddings(new_num_tokens, pad_to_multiple_of, **kwargs)
2592 # update vocab size
2593 self.config.text_config.vocab_size = model_embeds.num_embeddings
2594 self.config.vocab_size = model_embeds.num_embeddings
2595 self.vocab_size = model_embeds.num_embeddings
2596 return model_embeds
2597
2598 def _encode_image(self, pixel_values):
2599 if len(pixel_values.shape) == 4:
2600 batch_size, C, H, W = pixel_values.shape
2601 T = 1
2602 x = self.vision_tower.forward_features_unpool(pixel_values)
2603 else:
2604 raise ValueError(f'invalid image shape {pixel_values.shape}')
2605
2606 if self.image_pos_embed is not None:
2607 x = x.view(batch_size * T, -1, x.shape[-1])
2608 num_tokens = x.shape[-2]
2609 h, w = int(num_tokens ** 0.5), int(num_tokens ** 0.5)
2610 assert h * w == num_tokens, 'only support square feature maps for now'
2611 x = x.view(batch_size * T, h, w, x.shape[-1])
2612 pos_embed = self.image_pos_embed(x)
2613 x = x + pos_embed
2614 x = x.view(batch_size, T * h*w, x.shape[-1])
2615
2616 if self.visual_temporal_embed is not None:
2617 visual_temporal_embed = self.visual_temporal_embed(x.view(batch_size, T, -1, x.shape[-1])[:, :, 0])
2618 x = x.view(batch_size, T, -1, x.shape[-1]) + visual_temporal_embed.view(1, T, 1, x.shape[-1])
2619
2620 x_feat_dict = {}
2621
2622 spatial_avg_pool_x = x.view(batch_size, T, -1, x.shape[-1]).mean(dim=2)
2623 x_feat_dict['spatial_avg_pool'] = spatial_avg_pool_x
2624
2625 temporal_avg_pool_x = x.view(batch_size, T, -1, x.shape[-1]).mean(dim=1)
2626 x_feat_dict['temporal_avg_pool'] = temporal_avg_pool_x
2627
2628 x = x.view(batch_size, T, -1, x.shape[-1])[:, -1]
2629 x_feat_dict['last_frame'] = x
2630
2631 new_x = []
2632 for _image_feature_source in self.image_feature_source:
2633 if _image_feature_source not in x_feat_dict:
2634 raise ValueError('invalid image feature source: {}'.format(_image_feature_source))
2635 new_x.append(x_feat_dict[_image_feature_source])
2636
2637 x = torch.cat(new_x, dim=1)
2638
2639 x = x @ self.image_projection
2640 x = self.image_proj_norm(x)
2641
2642 return x
2643
2644 def _merge_input_ids_with_image_features(
2645 self, image_features, inputs_embeds
2646 ):
2647 batch_size, image_token_length = image_features.size()[:-1]
2648 device = image_features.device
2649 image_attention_mask = torch.ones(batch_size, image_token_length, device=device)
2650
2651 # task_prefix_embeds: [batch_size, padded_context_length, hidden_size]
2652 # task_prefix_attention_mask: [batch_size, context_length]
2653 if inputs_embeds is None:
2654 return image_features, image_attention_mask
2655
2656 task_prefix_embeds = inputs_embeds
2657 task_prefix_attention_mask = torch.ones(batch_size, task_prefix_embeds.size(1), device=device)
2658
2659 if len(task_prefix_attention_mask.shape) == 3:
2660 task_prefix_attention_mask = task_prefix_attention_mask[:, 0]
2661
2662 # concat [image embeds, task prefix embeds]
2663 inputs_embeds = torch.cat([image_features, task_prefix_embeds], dim=1)
2664 attention_mask = torch.cat([image_attention_mask, task_prefix_attention_mask], dim=1)
2665
2666 return inputs_embeds, attention_mask
2667
2668
2669 @add_start_docstrings_to_model_forward(FLORENCE2_INPUTS_DOCSTRING)
2670 @replace_return_docstrings(output_type=Florence2Seq2SeqLMOutput, config_class=_CONFIG_FOR_DOC)
2671 def forward(
2672 self,
2673 input_ids: torch.LongTensor = None,
2674 pixel_values: torch.FloatTensor = None,
2675 attention_mask: Optional[torch.Tensor] = None,
2676 decoder_input_ids: Optional[torch.LongTensor] = None,
2677 decoder_attention_mask: Optional[torch.LongTensor] = None,
2678 head_mask: Optional[torch.Tensor] = None,
2679 decoder_head_mask: Optional[torch.Tensor] = None,
2680 cross_attn_head_mask: Optional[torch.Tensor] = None,
2681 encoder_outputs: Optional[List[torch.FloatTensor]] = None,
2682 past_key_values: Optional[List[torch.FloatTensor]] = None,
2683 inputs_embeds: Optional[torch.FloatTensor] = None,
2684 decoder_inputs_embeds: Optional[torch.FloatTensor] = None,
2685 labels: Optional[torch.LongTensor] = None,
2686 use_cache: Optional[bool] = None,
2687 output_attentions: Optional[bool] = None,
2688 output_hidden_states: Optional[bool] = None,
2689 return_dict: Optional[bool] = None,
2690 ) -> Union[Tuple, Florence2Seq2SeqLMOutput]:
2691 r"""
2692 Args:
2693 labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
2694 Labels for computing the masked language modeling loss. Indices should either be in `[0, ...,
2695 config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored
2696 (masked), the loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`.
2697
2698 Returns:
2699
2700 Example:
2701
2702 ```python
2703 >>> from PIL import Image
2704 >>> import requests
2705 >>> from transformers import AutoProcessor, Florence2ForConditionalGeneration
2706
2707 >>> model = Florence2ForConditionalGeneration.from_pretrained("microsoft/Florence-2-large")
2708 >>> processor = AutoProcessor.from_pretrained("microsoft/Florence-2-large")
2709
2710 >>> prompt = "<CAPTION>"
2711 >>> url = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/transformers/tasks/car.jpg"
2712 >>> image = Image.open(requests.get(url, stream=True).raw)
2713
2714 >>> inputs = processor(text=prompt, images=image, return_tensors="pt")
2715
2716 >>> # Generate
2717 >>> generate_ids = model.generate(**inputs, max_length=100)
2718 >>> processor.batch_decode(generate_ids, skip_special_tokens=True, clean_up_tokenization_spaces=False)[0]
2719 "A green car parked in front of a yellow building."
2720 ```"""
2721 output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
2722 output_hidden_states = (
2723 output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
2724 )
2725 return_dict = return_dict if return_dict is not None else self.config.use_return_dict
2726
2727 image_features = None
2728 if inputs_embeds is None:
2729 # 1. Extra the input embeddings
2730 if input_ids is not None:
2731 inputs_embeds = self.get_input_embeddings()(input_ids)
2732 # 2. Merge text and images
2733 if pixel_values is not None:
2734 # (batch_size, num_image_tokens, hidden_size)
2735 image_features = self._encode_image(pixel_values)
2736 inputs_embeds, attention_mask = self._merge_input_ids_with_image_features(image_features, inputs_embeds)
2737
2738 if inputs_embeds is not None:
2739 attention_mask = attention_mask.to(inputs_embeds.dtype)
2740 outputs = self.language_model(
2741 attention_mask=attention_mask,
2742 labels=labels,
2743 inputs_embeds=inputs_embeds,
2744 decoder_input_ids=decoder_input_ids,
2745 encoder_outputs=encoder_outputs,
2746 decoder_attention_mask=decoder_attention_mask,
2747 head_mask=head_mask,
2748 decoder_head_mask=decoder_head_mask,
2749 cross_attn_head_mask=cross_attn_head_mask,
2750 past_key_values=past_key_values,
2751 decoder_inputs_embeds=decoder_inputs_embeds,
2752 use_cache=use_cache,
2753 output_attentions=output_attentions,
2754 output_hidden_states=output_hidden_states,
2755 return_dict=return_dict,
2756 )
2757
2758 logits = outputs.logits
2759 logits = logits.float()
2760 loss = outputs.loss
2761 if not return_dict:
2762 output = (logits,) + outputs[1:]
2763 return (loss,) + output if loss is not None else output
2764
2765 return Florence2Seq2SeqLMOutput(
2766 loss=loss,
2767 logits=logits,
2768 past_key_values=outputs.past_key_values,
2769 decoder_hidden_states=outputs.decoder_hidden_states,
2770 decoder_attentions=outputs.decoder_attentions,
2771 cross_attentions=outputs.cross_attentions,
2772 encoder_last_hidden_state=outputs.encoder_last_hidden_state,
2773 encoder_hidden_states=outputs.encoder_hidden_states,
2774 encoder_attentions=outputs.encoder_attentions,
2775 image_hidden_states=image_features
2776 )
2777
2778 def generate(
2779 self,
2780 input_ids,
2781 inputs_embeds=None,
2782 pixel_values=None,
2783 **kwargs
2784 ):
2785
2786 if inputs_embeds is None:
2787 # 1. Extra the input embeddings
2788 if input_ids is not None:
2789 inputs_embeds = self.get_input_embeddings()(input_ids)
2790 # 2. Merge text and images
2791 if pixel_values is not None:
2792 image_features = self._encode_image(pixel_values)
2793 inputs_embeds, attention_mask = self._merge_input_ids_with_image_features(image_features, inputs_embeds)
2794
2795 return self.language_model.generate(
2796 input_ids=None,
2797 inputs_embeds=inputs_embeds,
2798 **kwargs
2799 )
2800
2801 def prepare_inputs_for_generation(
2802 self,
2803 decoder_input_ids,
2804 past_key_values=None,
2805 attention_mask=None,
2806 pixel_values=None,
2807 decoder_attention_mask=None,
2808 head_mask=None,
2809 decoder_head_mask=None,
2810 cross_attn_head_mask=None,
2811 use_cache=None,
2812 encoder_outputs=None,
2813 **kwargs,
2814 ):
2815 # cut decoder_input_ids if past_key_values is used
2816 if past_key_values is not None:
2817 past_length = past_key_values[0][0].shape[2]
2818
2819 # Some generation methods already pass only the last input ID
2820 if decoder_input_ids.shape[1] > past_length:
2821 remove_prefix_length = past_length
2822 else:
2823 # Default to old behavior: keep only final ID
2824 remove_prefix_length = decoder_input_ids.shape[1] - 1
2825
2826 decoder_input_ids = decoder_input_ids[:, remove_prefix_length:]
2827
2828 return {
2829 "input_ids": None, # encoder_outputs is defined. input_ids not needed
2830 "encoder_outputs": encoder_outputs,
2831 "past_key_values": past_key_values,
2832 "decoder_input_ids": decoder_input_ids,
2833 "attention_mask": attention_mask,
2834 "pixel_values": pixel_values,
2835 "decoder_attention_mask": decoder_attention_mask,
2836 "head_mask": head_mask,
2837 "decoder_head_mask": decoder_head_mask,
2838 "cross_attn_head_mask": cross_attn_head_mask,
2839 "use_cache": use_cache, # change this to avoid caching (presumably for debugging)
2840 }
2841
2842 def prepare_decoder_input_ids_from_labels(self, labels: torch.Tensor):
2843 return self.language_model.shift_tokens_right(labels)
2844
2845 def _reorder_cache(self, *args, **kwargs):
2846 return self.language_model._reorder_cache(*args, **kwargs)