Prompt48 commited on
Commit
bfe04ab
·
verified ·
1 Parent(s): c47f05a

Upload edit\Qwen3-TTS-test\.venv\Lib\site-packages\transformers\models\groupvit\modeling_groupvit.py with huggingface_hub

Browse files
edit//Qwen3-TTS-test//.venv//Lib//site-packages//transformers//models//groupvit//modeling_groupvit.py ADDED
@@ -0,0 +1,1431 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # coding=utf-8
2
+ # Copyright 2022 NVIDIA and The HuggingFace 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
+ """PyTorch GroupViT model."""
16
+
17
+ import collections.abc
18
+ from dataclasses import dataclass
19
+ from typing import Any, Optional, Union
20
+
21
+ import numpy as np
22
+ import torch
23
+ from torch import nn
24
+
25
+ from ...activations import ACT2FN
26
+ from ...modeling_attn_mask_utils import _create_4d_causal_attention_mask, _prepare_4d_attention_mask
27
+ from ...modeling_layers import GradientCheckpointingLayer
28
+ from ...modeling_outputs import BaseModelOutput, BaseModelOutputWithPooling
29
+ from ...modeling_utils import PreTrainedModel
30
+ from ...utils import ModelOutput, auto_docstring, filter_out_non_signature_kwargs, logging, torch_int
31
+ from .configuration_groupvit import GroupViTConfig, GroupViTTextConfig, GroupViTVisionConfig
32
+
33
+
34
+ logger = logging.get_logger(__name__)
35
+
36
+
37
+ # contrastive loss function, adapted from
38
+ # https://sachinruk.github.io/blog/pytorch/pytorch%20lightning/loss%20function/gpu/2021/03/07/CLIP.html
39
+ def contrastive_loss(logits: torch.Tensor) -> torch.Tensor:
40
+ return nn.functional.cross_entropy(logits, torch.arange(len(logits), device=logits.device))
41
+
42
+
43
+ # Copied from transformers.models.clip.modeling_clip.clip_loss with clip->groupvit
44
+ def groupvit_loss(similarity: torch.Tensor) -> torch.Tensor:
45
+ caption_loss = contrastive_loss(similarity)
46
+ image_loss = contrastive_loss(similarity.t())
47
+ return (caption_loss + image_loss) / 2.0
48
+
49
+
50
+ def hard_softmax(logits: torch.Tensor, dim: int):
51
+ y_soft = logits.softmax(dim)
52
+ # Straight through.
53
+ index = y_soft.max(dim, keepdim=True)[1]
54
+ y_hard = torch.zeros_like(logits, memory_format=torch.legacy_contiguous_format).scatter_(dim, index, 1.0)
55
+ ret = y_hard - y_soft.detach() + y_soft
56
+
57
+ return ret
58
+
59
+
60
+ def gumbel_softmax(logits: torch.Tensor, tau: float = 1, hard: bool = False, dim: int = -1) -> torch.Tensor:
61
+ # more stable https://github.com/pytorch/pytorch/issues/41663
62
+ gumbel_dist = torch.distributions.gumbel.Gumbel(
63
+ torch.tensor(0.0, device=logits.device, dtype=logits.dtype),
64
+ torch.tensor(1.0, device=logits.device, dtype=logits.dtype),
65
+ )
66
+ gumbels = gumbel_dist.sample(logits.shape)
67
+
68
+ gumbels = (logits + gumbels) / tau # ~Gumbel(logits,tau)
69
+ y_soft = gumbels.softmax(dim)
70
+
71
+ if hard:
72
+ # Straight through.
73
+ index = y_soft.max(dim, keepdim=True)[1]
74
+ y_hard = torch.zeros_like(logits, memory_format=torch.legacy_contiguous_format).scatter_(dim, index, 1.0)
75
+ ret = y_hard - y_soft.detach() + y_soft
76
+ else:
77
+ # Reparameterization trick.
78
+ ret = y_soft
79
+ return ret
80
+
81
+
82
+ def resize_attention_map(attentions, height, width, align_corners=False):
83
+ """
84
+ Args:
85
+ attentions (`torch.Tensor`): attention map of shape [batch_size, groups, feat_height*feat_width]
86
+ height (`int`): height of the output attention map
87
+ width (`int`): width of the output attention map
88
+ align_corners (`bool`, *optional*): the `align_corner` argument for `nn.functional.interpolate`.
89
+
90
+ Returns:
91
+ `torch.Tensor`: resized attention map of shape [batch_size, groups, height, width]
92
+ """
93
+
94
+ scale = (height * width // attentions.shape[2]) ** 0.5
95
+ if height > width:
96
+ feat_width = int(np.round(width / scale))
97
+ feat_height = attentions.shape[2] // feat_width
98
+ else:
99
+ feat_height = int(np.round(height / scale))
100
+ feat_width = attentions.shape[2] // feat_height
101
+
102
+ batch_size = attentions.shape[0]
103
+ groups = attentions.shape[1] # number of group token
104
+ # [batch_size, groups, height*width, groups] -> [batch_size, groups, height, width]
105
+ attentions = attentions.reshape(batch_size, groups, feat_height, feat_width)
106
+ attentions = nn.functional.interpolate(
107
+ attentions, size=(height, width), mode="bilinear", align_corners=align_corners
108
+ )
109
+ return attentions
110
+
111
+
112
+ def get_grouping_from_attentions(attentions, hw_shape):
113
+ """
114
+ Args:
115
+ attentions (`tuple(torch.FloatTensor)`: tuple of attention maps returned by `GroupViTVisionTransformer`
116
+ hw_shape (`tuple(int)`): height and width of the output attention map
117
+ Returns:
118
+ `torch.Tensor`: the attention map of shape [batch_size, groups, height, width]
119
+ """
120
+
121
+ attn_maps = []
122
+ with torch.no_grad():
123
+ prev_attn_masks = None
124
+ for attn_masks in attentions:
125
+ # [batch_size, num_groups, height x width] -> [batch_size, height x width, num_groups]
126
+ attn_masks = attn_masks.permute(0, 2, 1).contiguous()
127
+ if prev_attn_masks is None:
128
+ prev_attn_masks = attn_masks
129
+ else:
130
+ prev_attn_masks = prev_attn_masks @ attn_masks
131
+ # [batch_size, heightxwidth, num_groups] -> [batch_size, num_groups, heightxwidth] -> [batch_size, num_groups, height, width]
132
+ cur_attn_map = resize_attention_map(prev_attn_masks.permute(0, 2, 1).contiguous(), *hw_shape)
133
+ attn_maps.append(cur_attn_map)
134
+
135
+ # [batch_size, num_groups, height, width]
136
+ final_grouping = attn_maps[-1]
137
+
138
+ return final_grouping
139
+
140
+
141
+ class GroupViTCrossAttentionLayer(nn.Module):
142
+ def __init__(self, config: GroupViTVisionConfig):
143
+ super().__init__()
144
+ self.attn = GroupViTAttention(config)
145
+ self.norm2 = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)
146
+ self.mlp = GroupViTMLP(config)
147
+ self.norm_post = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)
148
+
149
+ def forward(self, query, key):
150
+ x = query
151
+ x = x + self.attn(query, encoder_hidden_states=key)[0]
152
+ x = x + self.mlp(self.norm2(x))
153
+ x = self.norm_post(x)
154
+ return x
155
+
156
+
157
+ class GroupViTAssignAttention(nn.Module):
158
+ def __init__(self, config: GroupViTVisionConfig):
159
+ super().__init__()
160
+ self.scale = config.hidden_size**-0.5
161
+
162
+ self.q_proj = nn.Linear(config.hidden_size, config.hidden_size)
163
+ self.k_proj = nn.Linear(config.hidden_size, config.hidden_size)
164
+ self.v_proj = nn.Linear(config.hidden_size, config.hidden_size)
165
+ self.proj = nn.Linear(config.hidden_size, config.hidden_size)
166
+ self.assign_eps = config.assign_eps
167
+
168
+ def get_attn(self, attn, gumbel=True, hard=True):
169
+ if gumbel and self.training:
170
+ attn = gumbel_softmax(attn, dim=-2, hard=hard)
171
+ else:
172
+ if hard:
173
+ attn = hard_softmax(attn, dim=-2)
174
+ else:
175
+ attn = nn.functional.softmax(attn, dim=-2)
176
+
177
+ return attn
178
+
179
+ def forward(self, query, key):
180
+ value = key
181
+ # [batch_size, query_length, channels]
182
+ query = self.q_proj(query)
183
+
184
+ # [batch_size, key_length, channels]
185
+ key = self.k_proj(key)
186
+
187
+ # [batch_size, key_length, channels]
188
+ value = self.v_proj(value)
189
+
190
+ # [batch_size, query_length, key_length]
191
+ raw_attn = (query @ key.transpose(-2, -1)) * self.scale
192
+
193
+ attn = self.get_attn(raw_attn)
194
+ soft_attn = self.get_attn(raw_attn, gumbel=False, hard=False)
195
+
196
+ attn = attn / (attn.sum(dim=-1, keepdim=True) + self.assign_eps)
197
+
198
+ out = attn @ value
199
+
200
+ out = self.proj(out)
201
+
202
+ return out, soft_attn
203
+
204
+
205
+ class GroupViTTokenAssign(nn.Module):
206
+ def __init__(self, config: GroupViTVisionConfig, num_group_token, num_output_group):
207
+ super().__init__()
208
+ self.num_output_group = num_output_group
209
+ # norm on group_tokens
210
+ self.norm_tokens = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)
211
+ assign_mlp_ratio = (
212
+ config.assign_mlp_ratio
213
+ if isinstance(config.assign_mlp_ratio, collections.abc.Iterable)
214
+ else (config.assign_mlp_ratio, config.assign_mlp_ratio)
215
+ )
216
+ tokens_dim, channels_dim = [int(x * config.hidden_size) for x in assign_mlp_ratio]
217
+ self.mlp_inter = GroupViTMixerMLP(config, num_group_token, tokens_dim, num_output_group)
218
+ self.norm_post_tokens = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)
219
+ # norm on x
220
+ self.norm_x = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)
221
+ self.pre_assign_attn = GroupViTCrossAttentionLayer(config)
222
+
223
+ self.assign = GroupViTAssignAttention(config)
224
+ self.norm_new_x = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)
225
+ self.mlp_channels = GroupViTMLP(config, config.hidden_size, channels_dim, config.hidden_size)
226
+
227
+ def project_group_token(self, group_tokens):
228
+ """
229
+ Args:
230
+ group_tokens (torch.Tensor): group tokens, [batch_size, num_group_tokens, channels]
231
+
232
+ Returns:
233
+ projected_group_tokens (torch.Tensor): [batch_size, num_output_groups, channels]
234
+ """
235
+ # [B, num_output_groups, C] <- [B, num_group_tokens, C]
236
+ projected_group_tokens = self.mlp_inter(group_tokens)
237
+ projected_group_tokens = self.norm_post_tokens(projected_group_tokens)
238
+ return projected_group_tokens
239
+
240
+ def forward(self, image_tokens, group_tokens):
241
+ """
242
+ Args:
243
+ image_tokens (`torch.Tensor`): image tokens, of shape [batch_size, input_length, channels]
244
+ group_tokens (`torch.Tensor`): group tokens, [batch_size, num_group_tokens, channels]
245
+ """
246
+
247
+ group_tokens = self.norm_tokens(group_tokens)
248
+ image_tokens = self.norm_x(image_tokens)
249
+ # [batch_size, num_output_groups, channels]
250
+ projected_group_tokens = self.project_group_token(group_tokens)
251
+ projected_group_tokens = self.pre_assign_attn(projected_group_tokens, image_tokens)
252
+ new_image_tokens, attention = self.assign(projected_group_tokens, image_tokens)
253
+ new_image_tokens += projected_group_tokens
254
+
255
+ new_image_tokens = new_image_tokens + self.mlp_channels(self.norm_new_x(new_image_tokens))
256
+
257
+ return new_image_tokens, attention
258
+
259
+
260
+ @dataclass
261
+ @auto_docstring
262
+ class GroupViTModelOutput(ModelOutput):
263
+ r"""
264
+ loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `return_loss` is `True`):
265
+ Contrastive loss for image-text similarity.
266
+ logits_per_image (`torch.FloatTensor` of shape `(image_batch_size, text_batch_size)`):
267
+ The scaled dot product scores between `image_embeds` and `text_embeds`. This represents the image-text
268
+ similarity scores.
269
+ logits_per_text (`torch.FloatTensor` of shape `(text_batch_size, image_batch_size)`):
270
+ The scaled dot product scores between `text_embeds` and `image_embeds`. This represents the text-image
271
+ similarity scores.
272
+ segmentation_logits (`torch.FloatTensor` of shape `(batch_size, config.num_labels, logits_height, logits_width)`):
273
+ Classification scores for each pixel.
274
+
275
+ <Tip warning={true}>
276
+
277
+ The logits returned do not necessarily have the same size as the `pixel_values` passed as inputs. This is
278
+ to avoid doing two interpolations and lose some quality when a user needs to resize the logits to the
279
+ original image size as post-processing. You should always check your logits shape and resize as needed.
280
+
281
+ </Tip>
282
+ text_embeds (`torch.FloatTensor` of shape `(batch_size, output_dim`):
283
+ The text embeddings obtained by applying the projection layer to the pooled output of
284
+ [`GroupViTTextModel`].
285
+ image_embeds (`torch.FloatTensor` of shape `(batch_size, output_dim`):
286
+ The image embeddings obtained by applying the projection layer to the pooled output of
287
+ [`GroupViTVisionModel`].
288
+ text_model_output (`BaseModelOutputWithPooling`):
289
+ The output of the [`GroupViTTextModel`].
290
+ vision_model_output (`BaseModelOutputWithPooling`):
291
+ The output of the [`GroupViTVisionModel`].
292
+ """
293
+
294
+ loss: Optional[torch.FloatTensor] = None
295
+ logits_per_image: Optional[torch.FloatTensor] = None
296
+ logits_per_text: Optional[torch.FloatTensor] = None
297
+ segmentation_logits: Optional[torch.FloatTensor] = None
298
+ text_embeds: Optional[torch.FloatTensor] = None
299
+ image_embeds: Optional[torch.FloatTensor] = None
300
+ text_model_output: BaseModelOutputWithPooling = None
301
+ vision_model_output: BaseModelOutputWithPooling = None
302
+
303
+ def to_tuple(self) -> tuple[Any]:
304
+ return tuple(
305
+ self[k] if k not in ["text_model_output", "vision_model_output"] else getattr(self, k).to_tuple()
306
+ for k in self.keys()
307
+ )
308
+
309
+
310
+ class GroupViTPatchEmbeddings(nn.Module):
311
+ """
312
+ Image to Patch Embedding.
313
+ """
314
+
315
+ def __init__(
316
+ self,
317
+ image_size: int = 224,
318
+ patch_size: Union[int, tuple[int, int]] = 16,
319
+ num_channels: int = 3,
320
+ embed_dim: int = 768,
321
+ ):
322
+ super().__init__()
323
+ image_size = image_size if isinstance(image_size, collections.abc.Iterable) else (image_size, image_size)
324
+ patch_size = patch_size if isinstance(patch_size, collections.abc.Iterable) else (patch_size, patch_size)
325
+ num_patches = (image_size[1] // patch_size[1]) * (image_size[0] // patch_size[0])
326
+ self.image_size = image_size
327
+ self.patch_size = patch_size
328
+ self.num_patches = num_patches
329
+
330
+ self.projection = nn.Conv2d(num_channels, embed_dim, kernel_size=patch_size, stride=patch_size)
331
+
332
+ def forward(self, pixel_values: torch.Tensor, interpolate_pos_encoding: bool = False) -> torch.Tensor:
333
+ batch_size, num_channels, height, width = pixel_values.shape
334
+ if not interpolate_pos_encoding:
335
+ if height != self.image_size[0] or width != self.image_size[1]:
336
+ raise ValueError(
337
+ f"Input image size ({height}*{width}) doesn't match model"
338
+ f" ({self.image_size[0]}*{self.image_size[1]})."
339
+ )
340
+ x = self.projection(pixel_values).flatten(2).transpose(1, 2)
341
+ return x
342
+
343
+
344
+ class GroupViTVisionEmbeddings(nn.Module):
345
+ def __init__(self, config: GroupViTVisionConfig):
346
+ super().__init__()
347
+
348
+ self.patch_embeddings = GroupViTPatchEmbeddings(
349
+ image_size=config.image_size,
350
+ patch_size=config.patch_size,
351
+ num_channels=config.num_channels,
352
+ embed_dim=config.hidden_size,
353
+ )
354
+ num_patches = self.patch_embeddings.num_patches
355
+ self.position_embeddings = nn.Parameter(torch.zeros(1, num_patches, config.hidden_size))
356
+ self.dropout = nn.Dropout(config.dropout)
357
+ self.layernorm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)
358
+ self.patch_size = config.patch_size
359
+ self.config = config
360
+
361
+ def interpolate_pos_encoding(self, embeddings: torch.Tensor, height: int, width: int) -> torch.Tensor:
362
+ """
363
+ This method allows to interpolate the pre-trained position encodings, to be able to use the model on higher resolution
364
+ images. This method is also adapted to support torch.jit tracing and no class embeddings.
365
+
366
+ Adapted from:
367
+ - https://github.com/facebookresearch/dino/blob/de9ee3df6cf39fac952ab558447af1fa1365362a/vision_transformer.py#L174-L194, and
368
+ - https://github.com/facebookresearch/dinov2/blob/e1277af2ba9496fbadf7aec6eba56e8d882d1e35/dinov2/models/vision_transformer.py#L179-L211
369
+ """
370
+
371
+ num_patches = embeddings.shape[1]
372
+ num_positions = self.position_embeddings.shape[1]
373
+
374
+ # always interpolate when tracing to ensure the exported model works for dynamic input shapes
375
+ if not torch.jit.is_tracing() and num_patches == num_positions and height == width:
376
+ return self.position_embeddings
377
+
378
+ patch_pos_embed = self.position_embeddings
379
+
380
+ dim = embeddings.shape[-1]
381
+
382
+ new_height = height // self.patch_size
383
+ new_width = width // self.patch_size
384
+
385
+ sqrt_num_positions = torch_int(num_positions**0.5)
386
+ patch_pos_embed = patch_pos_embed.reshape(1, sqrt_num_positions, sqrt_num_positions, dim)
387
+ patch_pos_embed = patch_pos_embed.permute(0, 3, 1, 2)
388
+
389
+ patch_pos_embed = nn.functional.interpolate(
390
+ patch_pos_embed,
391
+ size=(new_height, new_width),
392
+ mode="bicubic",
393
+ align_corners=False,
394
+ )
395
+
396
+ patch_pos_embed = patch_pos_embed.permute(0, 2, 3, 1).view(1, -1, dim)
397
+ return patch_pos_embed
398
+
399
+ def forward(self, pixel_values: torch.Tensor, interpolate_pos_encoding: bool = False) -> torch.Tensor:
400
+ batch_size, num_channels, height, width = pixel_values.shape
401
+ embeddings = self.patch_embeddings(pixel_values, interpolate_pos_encoding=interpolate_pos_encoding)
402
+
403
+ embeddings = self.layernorm(embeddings)
404
+
405
+ batch_size, seq_len, _ = embeddings.size()
406
+
407
+ # add positional encoding to each token
408
+ if interpolate_pos_encoding:
409
+ embeddings = embeddings + self.interpolate_pos_encoding(embeddings, height, width)
410
+ else:
411
+ embeddings = embeddings + self.position_embeddings
412
+
413
+ embeddings = self.dropout(embeddings)
414
+
415
+ return embeddings
416
+
417
+
418
+ # Copied from transformers.models.clip.modeling_clip.CLIPTextEmbeddings with CLIP->GroupViT
419
+ class GroupViTTextEmbeddings(nn.Module):
420
+ def __init__(self, config: GroupViTTextConfig):
421
+ super().__init__()
422
+ embed_dim = config.hidden_size
423
+
424
+ self.token_embedding = nn.Embedding(config.vocab_size, embed_dim)
425
+ self.position_embedding = nn.Embedding(config.max_position_embeddings, embed_dim)
426
+
427
+ # position_ids (1, len position emb) is contiguous in memory and exported when serialized
428
+ self.register_buffer(
429
+ "position_ids", torch.arange(config.max_position_embeddings).expand((1, -1)), persistent=False
430
+ )
431
+
432
+ def forward(
433
+ self,
434
+ input_ids: Optional[torch.LongTensor] = None,
435
+ position_ids: Optional[torch.LongTensor] = None,
436
+ inputs_embeds: Optional[torch.FloatTensor] = None,
437
+ ) -> torch.Tensor:
438
+ seq_length = input_ids.shape[-1] if input_ids is not None else inputs_embeds.shape[-2]
439
+ max_position_embedding = self.position_embedding.weight.shape[0]
440
+
441
+ if seq_length > max_position_embedding:
442
+ raise ValueError(
443
+ f"Sequence length must be less than max_position_embeddings (got `sequence length`: "
444
+ f"{seq_length} and max_position_embeddings: {max_position_embedding}"
445
+ )
446
+
447
+ if position_ids is None:
448
+ position_ids = self.position_ids[:, :seq_length]
449
+
450
+ if inputs_embeds is None:
451
+ inputs_embeds = self.token_embedding(input_ids)
452
+
453
+ position_embeddings = self.position_embedding(position_ids)
454
+ embeddings = inputs_embeds + position_embeddings
455
+
456
+ return embeddings
457
+
458
+
459
+ class GroupViTStage(nn.Module):
460
+ """This corresponds to the `GroupingLayer` class in the GroupViT implementation."""
461
+
462
+ def __init__(
463
+ self,
464
+ config: GroupViTVisionConfig,
465
+ depth: int,
466
+ num_prev_group_token: int,
467
+ num_group_token: int,
468
+ num_output_group: int,
469
+ ):
470
+ super().__init__()
471
+ self.depth = depth
472
+ self.num_group_token = num_group_token
473
+ if num_group_token > 0:
474
+ self.group_token = nn.Parameter(torch.zeros(1, num_group_token, config.hidden_size))
475
+ else:
476
+ self.group_token = None
477
+ self.layers = nn.ModuleList([GroupViTEncoderLayer(config) for _ in range(depth)])
478
+
479
+ if num_group_token > 0:
480
+ self.downsample = GroupViTTokenAssign(
481
+ config=config,
482
+ num_group_token=num_group_token,
483
+ num_output_group=num_output_group,
484
+ )
485
+ else:
486
+ self.downsample = None
487
+
488
+ if num_prev_group_token > 0 and num_group_token > 0:
489
+ self.group_projector = nn.Sequential(
490
+ nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps),
491
+ GroupViTMixerMLP(config, num_prev_group_token, config.hidden_size // 2, num_group_token),
492
+ )
493
+ else:
494
+ self.group_projector = None
495
+
496
+ @property
497
+ def with_group_token(self):
498
+ return self.group_token is not None
499
+
500
+ def split_x(self, x):
501
+ if self.with_group_token:
502
+ return x[:, : -self.num_group_token], x[:, -self.num_group_token :]
503
+ else:
504
+ return x, None
505
+
506
+ def concat_x(self, x: torch.Tensor, group_token: Optional[torch.Tensor] = None) -> torch.Tensor:
507
+ if group_token is None:
508
+ return x
509
+ return torch.cat([x, group_token], dim=1)
510
+
511
+ def forward(
512
+ self,
513
+ hidden_states: torch.Tensor,
514
+ prev_group_token: Optional[torch.Tensor] = None,
515
+ output_attentions: Optional[bool] = False,
516
+ ) -> tuple[torch.FloatTensor]:
517
+ """
518
+ Args:
519
+ hidden_states (`torch.FloatTensor`): input to the layer of shape `(batch, seq_len, embed_dim)`
520
+ attention_mask (`torch.FloatTensor`): attention mask of size
521
+ `(batch, 1, tgt_len, src_len)` where padding elements are indicated by very large negative values.
522
+ `(config.encoder_attention_heads,)`.
523
+ output_attentions (`bool`, *optional*):
524
+ Whether or not to return the grouping tensors of Grouping block.
525
+ """
526
+ if self.with_group_token:
527
+ group_token = self.group_token.expand(hidden_states.size(0), -1, -1)
528
+ if self.group_projector is not None:
529
+ group_token = group_token + self.group_projector(prev_group_token)
530
+ else:
531
+ group_token = None
532
+
533
+ x = hidden_states
534
+
535
+ cat_x = self.concat_x(x, group_token)
536
+ for layer in self.layers:
537
+ layer_out = layer(cat_x, attention_mask=None, causal_attention_mask=None)
538
+ cat_x = layer_out[0]
539
+
540
+ x, group_token = self.split_x(cat_x)
541
+
542
+ attention = None
543
+ if self.downsample is not None:
544
+ x, attention = self.downsample(x, group_token)
545
+
546
+ outputs = (x, group_token)
547
+ if output_attentions:
548
+ outputs = outputs + (attention,)
549
+
550
+ return outputs
551
+
552
+
553
+ class GroupViTMLP(nn.Module):
554
+ def __init__(
555
+ self,
556
+ config: GroupViTVisionConfig,
557
+ hidden_size: Optional[int] = None,
558
+ intermediate_size: Optional[int] = None,
559
+ output_size: Optional[int] = None,
560
+ ):
561
+ super().__init__()
562
+ self.config = config
563
+ self.activation_fn = ACT2FN[config.hidden_act]
564
+ hidden_size = hidden_size if hidden_size is not None else config.hidden_size
565
+ intermediate_size = intermediate_size if intermediate_size is not None else config.intermediate_size
566
+ output_size = output_size if output_size is not None else hidden_size
567
+ self.fc1 = nn.Linear(hidden_size, intermediate_size)
568
+ self.fc2 = nn.Linear(intermediate_size, output_size)
569
+
570
+ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
571
+ hidden_states = self.fc1(hidden_states)
572
+ hidden_states = self.activation_fn(hidden_states)
573
+ hidden_states = self.fc2(hidden_states)
574
+ return hidden_states
575
+
576
+
577
+ class GroupViTMixerMLP(GroupViTMLP):
578
+ def forward(self, x):
579
+ x = super().forward(x.transpose(1, 2))
580
+ return x.transpose(1, 2)
581
+
582
+
583
+ class GroupViTAttention(nn.Module):
584
+ """Multi-headed attention from 'Attention Is All You Need' paper"""
585
+
586
+ def __init__(self, config):
587
+ super().__init__()
588
+ self.config = config
589
+ self.embed_dim = config.hidden_size
590
+ self.num_heads = config.num_attention_heads
591
+ self.head_dim = self.embed_dim // self.num_heads
592
+ if self.head_dim * self.num_heads != self.embed_dim:
593
+ raise ValueError(
594
+ f"embed_dim must be divisible by num_heads (got `embed_dim`: {self.embed_dim} and `num_heads`:"
595
+ f" {self.num_heads})."
596
+ )
597
+ self.scale = self.head_dim**-0.5
598
+ self.dropout = config.attention_dropout
599
+
600
+ self.k_proj = nn.Linear(self.embed_dim, self.embed_dim)
601
+ self.v_proj = nn.Linear(self.embed_dim, self.embed_dim)
602
+ self.q_proj = nn.Linear(self.embed_dim, self.embed_dim)
603
+ self.out_proj = nn.Linear(self.embed_dim, self.embed_dim)
604
+
605
+ def _shape(self, tensor: torch.Tensor, seq_len: int, bsz: int):
606
+ return tensor.view(bsz, seq_len, self.num_heads, self.head_dim).transpose(1, 2).contiguous()
607
+
608
+ def forward(
609
+ self,
610
+ hidden_states: torch.Tensor,
611
+ attention_mask: Optional[torch.Tensor] = None,
612
+ causal_attention_mask: Optional[torch.Tensor] = None,
613
+ encoder_hidden_states: Optional[torch.FloatTensor] = None,
614
+ output_attentions: Optional[bool] = False,
615
+ ) -> tuple[torch.Tensor, Optional[torch.Tensor], Optional[tuple[torch.Tensor]]]:
616
+ """Input shape: Batch x Time x Channel"""
617
+
618
+ bsz, tgt_len, embed_dim = hidden_states.size()
619
+ is_cross_attention = encoder_hidden_states is not None
620
+
621
+ # get query proj
622
+ query_states = self.q_proj(hidden_states) * self.scale
623
+ if is_cross_attention:
624
+ key_states = self._shape(self.k_proj(encoder_hidden_states), -1, bsz)
625
+ value_states = self._shape(self.v_proj(encoder_hidden_states), -1, bsz)
626
+ else:
627
+ key_states = self._shape(self.k_proj(hidden_states), -1, bsz)
628
+ value_states = self._shape(self.v_proj(hidden_states), -1, bsz)
629
+
630
+ proj_shape = (bsz * self.num_heads, -1, self.head_dim)
631
+ query_states = self._shape(query_states, tgt_len, bsz).view(*proj_shape)
632
+ key_states = key_states.view(*proj_shape)
633
+ value_states = value_states.view(*proj_shape)
634
+
635
+ src_len = key_states.size(1)
636
+ attn_weights = torch.bmm(query_states, key_states.transpose(1, 2))
637
+
638
+ if attn_weights.size() != (bsz * self.num_heads, tgt_len, src_len):
639
+ raise ValueError(
640
+ f"Attention weights should be of size {(bsz * self.num_heads, tgt_len, src_len)}, but is"
641
+ f" {attn_weights.size()}"
642
+ )
643
+
644
+ # apply the causal_attention_mask first
645
+ if causal_attention_mask is not None:
646
+ if causal_attention_mask.size() != (bsz, 1, tgt_len, src_len):
647
+ raise ValueError(
648
+ f"Attention mask should be of size {(bsz, 1, tgt_len, src_len)}, but is"
649
+ f" {causal_attention_mask.size()}"
650
+ )
651
+ attn_weights = attn_weights.view(bsz, self.num_heads, tgt_len, src_len) + causal_attention_mask
652
+ attn_weights = attn_weights.view(bsz * self.num_heads, tgt_len, src_len)
653
+
654
+ if attention_mask is not None:
655
+ if attention_mask.size() != (bsz, 1, tgt_len, src_len):
656
+ raise ValueError(
657
+ f"Attention mask should be of size {(bsz, 1, tgt_len, src_len)}, but is {attention_mask.size()}"
658
+ )
659
+ attn_weights = attn_weights.view(bsz, self.num_heads, tgt_len, src_len) + attention_mask
660
+ attn_weights = attn_weights.view(bsz * self.num_heads, tgt_len, src_len)
661
+
662
+ attn_weights = nn.functional.softmax(attn_weights, dim=-1)
663
+
664
+ if output_attentions:
665
+ # this operation is a bit awkward, but it's required to
666
+ # make sure that attn_weights keeps its gradient.
667
+ # In order to do so, attn_weights have to reshaped
668
+ # twice and have to be reused in the following
669
+ attn_weights_reshaped = attn_weights.view(bsz, self.num_heads, tgt_len, src_len)
670
+ attn_weights = attn_weights_reshaped.view(bsz * self.num_heads, tgt_len, src_len)
671
+ else:
672
+ attn_weights_reshaped = None
673
+
674
+ attn_probs = nn.functional.dropout(attn_weights, p=self.dropout, training=self.training)
675
+
676
+ attn_output = torch.bmm(attn_probs, value_states)
677
+
678
+ if attn_output.size() != (bsz * self.num_heads, tgt_len, self.head_dim):
679
+ raise ValueError(
680
+ f"`attn_output` should be of size {(bsz, self.num_heads, tgt_len, self.head_dim)}, but is"
681
+ f" {attn_output.size()}"
682
+ )
683
+
684
+ attn_output = attn_output.view(bsz, self.num_heads, tgt_len, self.head_dim)
685
+ attn_output = attn_output.transpose(1, 2)
686
+ attn_output = attn_output.reshape(bsz, tgt_len, embed_dim)
687
+
688
+ attn_output = self.out_proj(attn_output)
689
+
690
+ return attn_output, attn_weights_reshaped
691
+
692
+
693
+ # Copied from transformers.models.altclip.modeling_altclip.AltCLIPEncoderLayer with AltCLIP->GroupViT
694
+ class GroupViTEncoderLayer(GradientCheckpointingLayer):
695
+ def __init__(self, config: GroupViTConfig):
696
+ super().__init__()
697
+ self.embed_dim = config.hidden_size
698
+ self.self_attn = GroupViTAttention(config)
699
+ self.layer_norm1 = nn.LayerNorm(self.embed_dim, eps=config.layer_norm_eps)
700
+ self.mlp = GroupViTMLP(config)
701
+ self.layer_norm2 = nn.LayerNorm(self.embed_dim, eps=config.layer_norm_eps)
702
+
703
+ def forward(
704
+ self,
705
+ hidden_states: torch.Tensor,
706
+ attention_mask: torch.Tensor,
707
+ causal_attention_mask: torch.Tensor,
708
+ output_attentions: Optional[bool] = False,
709
+ ) -> tuple[torch.FloatTensor]:
710
+ """
711
+ Args:
712
+ hidden_states (`torch.FloatTensor`): input to the layer of shape `(batch, seq_len, embed_dim)`
713
+ attention_mask (`torch.FloatTensor`): attention mask of size
714
+ `(batch, 1, tgt_len, src_len)` where padding elements are indicated by very large negative values.
715
+ `(config.encoder_attention_heads,)`.
716
+ output_attentions (`bool`, *optional*):
717
+ Whether or not to return the attentions tensors of all attention layers. See `attentions` under
718
+ returned tensors for more detail.
719
+ """
720
+ residual = hidden_states
721
+
722
+ hidden_states = self.layer_norm1(hidden_states)
723
+ hidden_states, attn_weights = self.self_attn(
724
+ hidden_states=hidden_states,
725
+ attention_mask=attention_mask,
726
+ causal_attention_mask=causal_attention_mask,
727
+ output_attentions=output_attentions,
728
+ )
729
+ hidden_states = residual + hidden_states
730
+
731
+ residual = hidden_states
732
+ hidden_states = self.layer_norm2(hidden_states)
733
+ hidden_states = self.mlp(hidden_states)
734
+ hidden_states = residual + hidden_states
735
+
736
+ outputs = (hidden_states,)
737
+
738
+ if output_attentions:
739
+ outputs += (attn_weights,)
740
+
741
+ return outputs
742
+
743
+
744
+ @auto_docstring
745
+ class GroupViTPreTrainedModel(PreTrainedModel):
746
+ config: GroupViTConfig
747
+ base_model_prefix = "groupvit"
748
+ supports_gradient_checkpointing = True
749
+
750
+ def _init_weights(self, module):
751
+ """Initialize the weights"""
752
+
753
+ init_range = self.config.initializer_range
754
+ if isinstance(module, (nn.Linear, nn.Conv2d)):
755
+ # Slightly different from the TF version which uses truncated_normal for initialization
756
+ # cf https://github.com/pytorch/pytorch/pull/5617
757
+ module.weight.data.normal_(mean=0.0, std=init_range)
758
+ if module.bias is not None:
759
+ module.bias.data.zero_()
760
+ elif isinstance(module, nn.LayerNorm):
761
+ module.bias.data.zero_()
762
+ module.weight.data.fill_(1.0)
763
+
764
+ factor = self.config.initializer_factor
765
+ if isinstance(module, GroupViTTextEmbeddings):
766
+ module.token_embedding.weight.data.normal_(mean=0.0, std=factor * 0.02)
767
+ module.position_embedding.weight.data.normal_(mean=0.0, std=factor * 0.02)
768
+ elif isinstance(module, GroupViTAttention):
769
+ factor = self.config.initializer_factor
770
+ in_proj_std = (module.embed_dim**-0.5) * ((2 * module.config.num_hidden_layers) ** -0.5) * factor
771
+ out_proj_std = (module.embed_dim**-0.5) * factor
772
+ nn.init.normal_(module.q_proj.weight, std=in_proj_std)
773
+ nn.init.normal_(module.k_proj.weight, std=in_proj_std)
774
+ nn.init.normal_(module.v_proj.weight, std=in_proj_std)
775
+ nn.init.normal_(module.out_proj.weight, std=out_proj_std)
776
+ elif isinstance(module, GroupViTMLP):
777
+ factor = self.config.initializer_factor
778
+ in_proj_std = (module.config.hidden_size**-0.5) * ((2 * module.config.num_hidden_layers) ** -0.5) * factor
779
+ fc_std = (2 * module.config.hidden_size) ** -0.5 * factor
780
+ nn.init.normal_(module.fc1.weight, std=fc_std)
781
+ nn.init.normal_(module.fc2.weight, std=in_proj_std)
782
+
783
+
784
+ class GroupViTVisionEncoder(nn.Module):
785
+ def __init__(self, config: GroupViTVisionConfig) -> None:
786
+ super().__init__()
787
+ self.config = config
788
+ self.stages = nn.ModuleList(
789
+ [
790
+ GroupViTStage(
791
+ config=config,
792
+ depth=config.depths[i],
793
+ num_group_token=config.num_group_tokens[i],
794
+ num_output_group=config.num_output_groups[i],
795
+ num_prev_group_token=config.num_output_groups[i - 1] if i > 0 else 0,
796
+ )
797
+ for i in range(len(config.depths))
798
+ ]
799
+ )
800
+ self.gradient_checkpointing = False
801
+
802
+ def forward(
803
+ self,
804
+ hidden_states: torch.Tensor,
805
+ output_hidden_states: Optional[bool] = None,
806
+ output_attentions: Optional[bool] = None,
807
+ return_dict: Optional[bool] = None,
808
+ ) -> Union[tuple, BaseModelOutput]:
809
+ output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
810
+ output_hidden_states = (
811
+ output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
812
+ )
813
+ return_dict = return_dict if return_dict is not None else self.config.use_return_dict
814
+
815
+ all_hidden_states = () if output_hidden_states else None
816
+ all_groupings = () if output_attentions else None
817
+
818
+ group_tokens = None
819
+
820
+ for i, stage in enumerate(self.stages):
821
+ if output_hidden_states:
822
+ all_hidden_states = all_hidden_states + (hidden_states,)
823
+
824
+ layer_outputs = stage(hidden_states, group_tokens, output_attentions)
825
+
826
+ hidden_states = layer_outputs[0]
827
+ group_tokens = layer_outputs[1]
828
+
829
+ if output_attentions and layer_outputs[2] is not None:
830
+ all_groupings = all_groupings + (layer_outputs[2],)
831
+
832
+ if output_hidden_states:
833
+ all_hidden_states = all_hidden_states + (hidden_states,)
834
+
835
+ if not return_dict:
836
+ return tuple(v for v in [hidden_states, all_hidden_states, all_groupings] if v is not None)
837
+ return BaseModelOutput(
838
+ last_hidden_state=hidden_states, hidden_states=all_hidden_states, attentions=all_groupings
839
+ )
840
+
841
+
842
+ class GroupViTTextEncoder(nn.Module):
843
+ """
844
+ Transformer encoder consisting of `config.num_hidden_layers` self-attention layers. Each layer is a
845
+ [`GroupViTEncoderLayer`].
846
+
847
+ Args:
848
+ config: GroupViTTextConfig
849
+ """
850
+
851
+ def __init__(self, config: GroupViTTextConfig):
852
+ super().__init__()
853
+ self.config = config
854
+ self.layers = nn.ModuleList([GroupViTEncoderLayer(config) for _ in range(config.num_hidden_layers)])
855
+ self.gradient_checkpointing = False
856
+
857
+ def forward(
858
+ self,
859
+ inputs_embeds,
860
+ attention_mask: Optional[torch.Tensor] = None,
861
+ causal_attention_mask: Optional[torch.Tensor] = None,
862
+ output_attentions: Optional[bool] = None,
863
+ output_hidden_states: Optional[bool] = None,
864
+ return_dict: Optional[bool] = None,
865
+ ) -> Union[tuple, BaseModelOutput]:
866
+ r"""
867
+ Args:
868
+ inputs_embeds (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`):
869
+ Optionally, instead of passing `input_ids` you can choose to directly pass an embedded representation.
870
+ This is useful if you want more control over how to convert `input_ids` indices into associated vectors
871
+ than the model's internal embedding lookup matrix.
872
+ attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):
873
+ Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:
874
+
875
+ - 1 for tokens that are **not masked**,
876
+ - 0 for tokens that are **masked**.
877
+
878
+ [What are attention masks?](../glossary#attention-mask)
879
+ causal_attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):
880
+ Causal mask for the text model. Mask values selected in `[0, 1]`:
881
+
882
+ - 1 for tokens that are **not masked**,
883
+ - 0 for tokens that are **masked**.
884
+
885
+ [What are attention masks?](../glossary#attention-mask)
886
+ output_attentions (`bool`, *optional*):
887
+ Whether or not to return the attentions tensors of all attention layers. See `attentions` under
888
+ returned tensors for more detail.
889
+ output_hidden_states (`bool`, *optional*):
890
+ Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors
891
+ for more detail.
892
+ return_dict (`bool`, *optional*):
893
+ Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.
894
+ """
895
+ output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
896
+ output_hidden_states = (
897
+ output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
898
+ )
899
+ return_dict = return_dict if return_dict is not None else self.config.use_return_dict
900
+
901
+ encoder_states = () if output_hidden_states else None
902
+ all_attentions = () if output_attentions else None
903
+
904
+ hidden_states = inputs_embeds
905
+ for idx, encoder_layer in enumerate(self.layers):
906
+ if output_hidden_states:
907
+ encoder_states = encoder_states + (hidden_states,)
908
+ layer_outputs = encoder_layer(
909
+ hidden_states,
910
+ attention_mask,
911
+ causal_attention_mask,
912
+ output_attentions=output_attentions,
913
+ )
914
+
915
+ hidden_states = layer_outputs[0]
916
+
917
+ if output_attentions:
918
+ all_attentions = all_attentions + (layer_outputs[1],)
919
+
920
+ if output_hidden_states:
921
+ encoder_states = encoder_states + (hidden_states,)
922
+
923
+ if not return_dict:
924
+ return tuple(v for v in [hidden_states, encoder_states, all_attentions] if v is not None)
925
+ return BaseModelOutput(
926
+ last_hidden_state=hidden_states, hidden_states=encoder_states, attentions=all_attentions
927
+ )
928
+
929
+
930
+ class GroupViTTextTransformer(nn.Module):
931
+ def __init__(self, config: GroupViTTextConfig):
932
+ super().__init__()
933
+ self.config = config
934
+ embed_dim = config.hidden_size
935
+ self.embeddings = GroupViTTextEmbeddings(config)
936
+ self.encoder = GroupViTTextEncoder(config)
937
+ self.final_layer_norm = nn.LayerNorm(embed_dim, eps=config.layer_norm_eps)
938
+
939
+ # For `pooled_output` computation
940
+ self.eos_token_id = config.eos_token_id
941
+
942
+ @auto_docstring
943
+ def forward(
944
+ self,
945
+ input_ids: Optional[torch.Tensor] = None,
946
+ attention_mask: Optional[torch.Tensor] = None,
947
+ position_ids: Optional[torch.Tensor] = None,
948
+ output_attentions: Optional[bool] = None,
949
+ output_hidden_states: Optional[bool] = None,
950
+ return_dict: Optional[bool] = None,
951
+ ) -> Union[tuple, BaseModelOutputWithPooling]:
952
+ output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
953
+ output_hidden_states = (
954
+ output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
955
+ )
956
+ return_dict = return_dict if return_dict is not None else self.config.use_return_dict
957
+
958
+ if input_ids is None:
959
+ raise ValueError("You have to specify input_ids")
960
+
961
+ input_shape = input_ids.size()
962
+ input_ids = input_ids.view(-1, input_shape[-1])
963
+
964
+ hidden_states = self.embeddings(input_ids=input_ids, position_ids=position_ids)
965
+
966
+ # CLIP's text model uses causal mask, prepare it here.
967
+ # https://github.com/openai/CLIP/blob/cfcffb90e69f37bf2ff1e988237a0fbe41f33c04/clip/model.py#L324
968
+ causal_attention_mask = _create_4d_causal_attention_mask(
969
+ input_shape, hidden_states.dtype, device=hidden_states.device
970
+ )
971
+
972
+ # expand attention_mask
973
+ if attention_mask is not None:
974
+ # [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]
975
+ attention_mask = _prepare_4d_attention_mask(attention_mask, hidden_states.dtype)
976
+
977
+ encoder_outputs = self.encoder(
978
+ inputs_embeds=hidden_states,
979
+ attention_mask=attention_mask,
980
+ causal_attention_mask=causal_attention_mask,
981
+ output_attentions=output_attentions,
982
+ output_hidden_states=output_hidden_states,
983
+ return_dict=return_dict,
984
+ )
985
+
986
+ last_hidden_state = encoder_outputs[0]
987
+ last_hidden_state = self.final_layer_norm(last_hidden_state)
988
+
989
+ if self.eos_token_id == 2:
990
+ # The `eos_token_id` was incorrect before PR #24773: Let's keep what have been done here.
991
+ # A CLIP model with such `eos_token_id` in the config can't work correctly with extra new tokens added
992
+ # ------------------------------------------------------------
993
+ # text_embeds.shape = [batch_size, sequence_length, transformer.width]
994
+ # take features from the eot embedding (eot_token is the highest number in each sequence)
995
+ # casting to torch.int for onnx compatibility: argmax doesn't support int64 inputs with opset 14
996
+ pooled_output = last_hidden_state[
997
+ torch.arange(last_hidden_state.shape[0], device=last_hidden_state.device),
998
+ input_ids.to(dtype=torch.int, device=last_hidden_state.device).argmax(dim=-1),
999
+ ]
1000
+ else:
1001
+ # The config gets updated `eos_token_id` from PR #24773 (so the use of extra new tokens is possible)
1002
+ pooled_output = last_hidden_state[
1003
+ torch.arange(last_hidden_state.shape[0], device=last_hidden_state.device),
1004
+ # We need to get the first position of `eos_token_id` value (`pad_token_ids` might equal to `eos_token_id`)
1005
+ # Note: we assume each sequence (along batch dim.) contains an `eos_token_id` (e.g. prepared by the tokenizer)
1006
+ (input_ids.to(dtype=torch.int, device=last_hidden_state.device) == self.eos_token_id)
1007
+ .int()
1008
+ .argmax(dim=-1),
1009
+ ]
1010
+
1011
+ if not return_dict:
1012
+ return (last_hidden_state, pooled_output) + encoder_outputs[1:]
1013
+
1014
+ return BaseModelOutputWithPooling(
1015
+ last_hidden_state=last_hidden_state,
1016
+ pooler_output=pooled_output,
1017
+ hidden_states=encoder_outputs.hidden_states,
1018
+ attentions=encoder_outputs.attentions,
1019
+ )
1020
+
1021
+
1022
+ class GroupViTTextModel(GroupViTPreTrainedModel):
1023
+ config: GroupViTTextConfig
1024
+
1025
+ def __init__(self, config: GroupViTTextConfig):
1026
+ super().__init__(config)
1027
+ self.text_model = GroupViTTextTransformer(config)
1028
+ # Initialize weights and apply final processing
1029
+ self.post_init()
1030
+
1031
+ def get_input_embeddings(self) -> nn.Module:
1032
+ return self.text_model.embeddings.token_embedding
1033
+
1034
+ def set_input_embeddings(self, value):
1035
+ self.text_model.embeddings.token_embedding = value
1036
+
1037
+ @auto_docstring
1038
+ def forward(
1039
+ self,
1040
+ input_ids: Optional[torch.Tensor] = None,
1041
+ attention_mask: Optional[torch.Tensor] = None,
1042
+ position_ids: Optional[torch.Tensor] = None,
1043
+ output_attentions: Optional[bool] = None,
1044
+ output_hidden_states: Optional[bool] = None,
1045
+ return_dict: Optional[bool] = None,
1046
+ ) -> Union[tuple, BaseModelOutputWithPooling]:
1047
+ r"""
1048
+ Examples:
1049
+
1050
+ ```python
1051
+ >>> from transformers import CLIPTokenizer, GroupViTTextModel
1052
+
1053
+ >>> tokenizer = CLIPTokenizer.from_pretrained("nvidia/groupvit-gcc-yfcc")
1054
+ >>> model = GroupViTTextModel.from_pretrained("nvidia/groupvit-gcc-yfcc")
1055
+
1056
+ >>> inputs = tokenizer(["a photo of a cat", "a photo of a dog"], padding=True, return_tensors="pt")
1057
+
1058
+ >>> outputs = model(**inputs)
1059
+ >>> last_hidden_state = outputs.last_hidden_state
1060
+ >>> pooled_output = outputs.pooler_output # pooled (EOS token) states
1061
+ ```"""
1062
+ return self.text_model(
1063
+ input_ids=input_ids,
1064
+ attention_mask=attention_mask,
1065
+ position_ids=position_ids,
1066
+ output_attentions=output_attentions,
1067
+ output_hidden_states=output_hidden_states,
1068
+ return_dict=return_dict,
1069
+ )
1070
+
1071
+
1072
+ class GroupViTVisionTransformer(nn.Module):
1073
+ def __init__(self, config: GroupViTVisionConfig):
1074
+ super().__init__()
1075
+ self.config = config
1076
+ embed_dim = config.hidden_size
1077
+
1078
+ self.embeddings = GroupViTVisionEmbeddings(config)
1079
+ self.encoder = GroupViTVisionEncoder(config)
1080
+ self.layernorm = nn.LayerNorm(embed_dim, eps=config.layer_norm_eps)
1081
+
1082
+ @auto_docstring
1083
+ def forward(
1084
+ self,
1085
+ pixel_values: Optional[torch.FloatTensor] = None,
1086
+ output_hidden_states: Optional[bool] = None,
1087
+ output_attentions: Optional[bool] = None,
1088
+ return_dict: Optional[bool] = None,
1089
+ ) -> Union[tuple, BaseModelOutputWithPooling]:
1090
+ output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
1091
+ output_hidden_states = (
1092
+ output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
1093
+ )
1094
+ return_dict = return_dict if return_dict is not None else self.config.use_return_dict
1095
+
1096
+ if pixel_values is None:
1097
+ raise ValueError("You have to specify pixel_values")
1098
+
1099
+ hidden_states = self.embeddings(pixel_values)
1100
+
1101
+ encoder_outputs = self.encoder(
1102
+ hidden_states=hidden_states,
1103
+ output_hidden_states=output_hidden_states,
1104
+ output_attentions=output_attentions,
1105
+ return_dict=return_dict,
1106
+ )
1107
+
1108
+ last_hidden_state = encoder_outputs[0]
1109
+
1110
+ # normalize the last hidden state
1111
+ last_hidden_state = self.layernorm(last_hidden_state)
1112
+ pooled_output = last_hidden_state.mean(dim=1)
1113
+
1114
+ if not return_dict:
1115
+ return (last_hidden_state, pooled_output) + encoder_outputs[1:]
1116
+
1117
+ return BaseModelOutputWithPooling(
1118
+ last_hidden_state=last_hidden_state,
1119
+ pooler_output=pooled_output,
1120
+ hidden_states=encoder_outputs.hidden_states,
1121
+ attentions=encoder_outputs.attentions,
1122
+ )
1123
+
1124
+
1125
+ class GroupViTVisionModel(GroupViTPreTrainedModel):
1126
+ config: GroupViTVisionConfig
1127
+ main_input_name = "pixel_values"
1128
+
1129
+ def __init__(self, config: GroupViTVisionConfig):
1130
+ super().__init__(config)
1131
+ self.vision_model = GroupViTVisionTransformer(config)
1132
+ # Initialize weights and apply final processing
1133
+ self.post_init()
1134
+
1135
+ def get_input_embeddings(self) -> GroupViTPatchEmbeddings:
1136
+ return self.vision_model.embeddings.patch_embeddings
1137
+
1138
+ @auto_docstring
1139
+ def forward(
1140
+ self,
1141
+ pixel_values: Optional[torch.FloatTensor] = None,
1142
+ output_attentions: Optional[bool] = None,
1143
+ output_hidden_states: Optional[bool] = None,
1144
+ return_dict: Optional[bool] = None,
1145
+ ) -> Union[tuple, BaseModelOutputWithPooling]:
1146
+ r"""
1147
+ Examples:
1148
+
1149
+ ```python
1150
+ >>> from PIL import Image
1151
+ >>> import requests
1152
+ >>> from transformers import AutoProcessor, GroupViTVisionModel
1153
+
1154
+ >>> processor = AutoProcessor.from_pretrained("nvidia/groupvit-gcc-yfcc")
1155
+ >>> model = GroupViTVisionModel.from_pretrained("nvidia/groupvit-gcc-yfcc")
1156
+
1157
+ >>> url = "http://images.cocodataset.org/val2017/000000039769.jpg"
1158
+ >>> image = Image.open(requests.get(url, stream=True).raw)
1159
+
1160
+ >>> inputs = processor(images=image, return_tensors="pt")
1161
+
1162
+ >>> outputs = model(**inputs)
1163
+ >>> last_hidden_state = outputs.last_hidden_state
1164
+ >>> pooled_output = outputs.pooler_output # pooled CLS states
1165
+ ```"""
1166
+ return self.vision_model(
1167
+ pixel_values=pixel_values,
1168
+ output_attentions=output_attentions,
1169
+ output_hidden_states=output_hidden_states,
1170
+ return_dict=return_dict,
1171
+ )
1172
+
1173
+
1174
+ @auto_docstring
1175
+ class GroupViTModel(GroupViTPreTrainedModel):
1176
+ config: GroupViTConfig
1177
+
1178
+ def __init__(self, config: GroupViTConfig):
1179
+ super().__init__(config)
1180
+
1181
+ if not isinstance(config.text_config, GroupViTTextConfig):
1182
+ raise TypeError(
1183
+ "config.text_config is expected to be of type GroupViTTextConfig but is of type"
1184
+ f" {type(config.text_config)}."
1185
+ )
1186
+
1187
+ if not isinstance(config.vision_config, GroupViTVisionConfig):
1188
+ raise TypeError(
1189
+ "config.vision_config is expected to be of type GroupViTVisionConfig but is of type"
1190
+ f" {type(config.vision_config)}."
1191
+ )
1192
+
1193
+ text_config = config.text_config
1194
+ vision_config = config.vision_config
1195
+
1196
+ self.projection_dim = config.projection_dim
1197
+ self.projection_intermediate_dim = config.projection_intermediate_dim
1198
+ self.text_embed_dim = text_config.hidden_size
1199
+ self.vision_embed_dim = vision_config.hidden_size
1200
+
1201
+ self.text_model = GroupViTTextTransformer(text_config)
1202
+ self.vision_model = GroupViTVisionTransformer(vision_config)
1203
+
1204
+ self.visual_projection = nn.Sequential(
1205
+ nn.Linear(self.vision_embed_dim, self.projection_intermediate_dim, bias=True),
1206
+ nn.BatchNorm1d(self.projection_intermediate_dim),
1207
+ nn.ReLU(inplace=True),
1208
+ nn.Linear(self.projection_intermediate_dim, self.projection_dim, bias=True),
1209
+ )
1210
+ self.text_projection = nn.Sequential(
1211
+ nn.Linear(self.text_embed_dim, self.projection_intermediate_dim, bias=True),
1212
+ nn.BatchNorm1d(self.projection_intermediate_dim),
1213
+ nn.ReLU(inplace=True),
1214
+ nn.Linear(self.projection_intermediate_dim, self.projection_dim, bias=True),
1215
+ )
1216
+ self.logit_scale = nn.Parameter(torch.tensor(self.config.logit_scale_init_value))
1217
+
1218
+ # Initialize weights and apply final processing
1219
+ self.post_init()
1220
+
1221
+ @filter_out_non_signature_kwargs()
1222
+ @auto_docstring
1223
+ def get_text_features(
1224
+ self,
1225
+ input_ids: torch.Tensor,
1226
+ attention_mask: Optional[torch.Tensor] = None,
1227
+ position_ids: Optional[torch.Tensor] = None,
1228
+ ) -> torch.FloatTensor:
1229
+ r"""
1230
+ Returns:
1231
+ text_features (`torch.FloatTensor` of shape `(batch_size, output_dim`): The text embeddings obtained by
1232
+ applying the projection layer to the pooled output of [`GroupViTTextModel`].
1233
+
1234
+ Examples:
1235
+
1236
+ ```python
1237
+ >>> import torch
1238
+ >>> from transformers import CLIPTokenizer, GroupViTModel
1239
+
1240
+ >>> model = GroupViTModel.from_pretrained("nvidia/groupvit-gcc-yfcc")
1241
+ >>> tokenizer = CLIPTokenizer.from_pretrained("nvidia/groupvit-gcc-yfcc")
1242
+
1243
+ >>> inputs = tokenizer(["a photo of a cat", "a photo of a dog"], padding=True, return_tensors="pt")
1244
+ >>> with torch.inference_mode():
1245
+ ... text_features = model.get_text_features(**inputs)
1246
+ ```"""
1247
+ text_outputs: BaseModelOutputWithPooling = self.text_model(
1248
+ input_ids=input_ids,
1249
+ attention_mask=attention_mask,
1250
+ position_ids=position_ids,
1251
+ )
1252
+ text_features = self.text_projection(text_outputs.pooler_output)
1253
+ return text_features
1254
+
1255
+ @filter_out_non_signature_kwargs()
1256
+ @auto_docstring
1257
+ def get_image_features(self, pixel_values: torch.Tensor) -> torch.FloatTensor:
1258
+ r"""
1259
+ Returns:
1260
+ image_features (`torch.FloatTensor` of shape `(batch_size, output_dim`): The image embeddings obtained by
1261
+ applying the projection layer to the pooled output of [`GroupViTVisionModel`].
1262
+
1263
+ Examples:
1264
+
1265
+ ```python
1266
+ >>> import torch
1267
+ >>> from transformers import AutoProcessor, GroupViTModel
1268
+ >>> from transformers.image_utils import load_image
1269
+
1270
+ >>> model = GroupViTModel.from_pretrained("nvidia/groupvit-gcc-yfcc")
1271
+ >>> processor = AutoProcessor.from_pretrained("nvidia/groupvit-gcc-yfcc")
1272
+
1273
+ >>> url = "http://images.cocodataset.org/val2017/000000039769.jpg"
1274
+ >>> image = load_image(url)
1275
+
1276
+ >>> inputs = processor(images=image, return_tensors="pt")
1277
+
1278
+ >>> with torch.inference_mode():
1279
+ ... image_features = model.get_image_features(**inputs)
1280
+ ```"""
1281
+ vision_outputs: BaseModelOutputWithPooling = self.vision_model(pixel_values)
1282
+ image_features = self.visual_projection(vision_outputs.pooler_output)
1283
+ return image_features
1284
+
1285
+ @auto_docstring
1286
+ def forward(
1287
+ self,
1288
+ input_ids: Optional[torch.LongTensor] = None,
1289
+ pixel_values: Optional[torch.FloatTensor] = None,
1290
+ attention_mask: Optional[torch.Tensor] = None,
1291
+ position_ids: Optional[torch.LongTensor] = None,
1292
+ return_loss: Optional[bool] = None,
1293
+ output_attentions: Optional[bool] = None,
1294
+ output_hidden_states: Optional[bool] = None,
1295
+ output_segmentation: Optional[bool] = None,
1296
+ return_dict: Optional[bool] = None,
1297
+ ) -> Union[tuple, GroupViTModelOutput]:
1298
+ r"""
1299
+ return_loss (`bool`, *optional*):
1300
+ Whether or not to return the contrastive loss.
1301
+ output_segmentation (`bool`, *optional*):
1302
+ Whether or not to return the segmentation logits.
1303
+
1304
+ Examples:
1305
+
1306
+ ```python
1307
+ >>> from PIL import Image
1308
+ >>> import requests
1309
+ >>> from transformers import AutoProcessor, GroupViTModel
1310
+
1311
+ >>> model = GroupViTModel.from_pretrained("nvidia/groupvit-gcc-yfcc")
1312
+ >>> processor = AutoProcessor.from_pretrained("nvidia/groupvit-gcc-yfcc")
1313
+
1314
+ >>> url = "http://images.cocodataset.org/val2017/000000039769.jpg"
1315
+ >>> image = Image.open(requests.get(url, stream=True).raw)
1316
+
1317
+ >>> inputs = processor(
1318
+ ... text=["a photo of a cat", "a photo of a dog"], images=image, return_tensors="pt", padding=True
1319
+ ... )
1320
+
1321
+ >>> outputs = model(**inputs)
1322
+ >>> logits_per_image = outputs.logits_per_image # this is the image-text similarity score
1323
+ >>> probs = logits_per_image.softmax(dim=1) # we can take the softmax to get the label probabilities
1324
+ ```"""
1325
+ # Use GROUPVIT model's config for some fields (if specified) instead of those of vision & text components.
1326
+ output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
1327
+ output_segmentation = (
1328
+ output_segmentation if output_segmentation is not None else self.config.output_segmentation
1329
+ )
1330
+ if output_segmentation:
1331
+ output_attentions = True
1332
+ output_hidden_states = (
1333
+ output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
1334
+ )
1335
+ return_dict = return_dict if return_dict is not None else self.config.use_return_dict
1336
+
1337
+ vision_outputs = self.vision_model(
1338
+ pixel_values=pixel_values,
1339
+ output_attentions=output_attentions,
1340
+ output_hidden_states=output_hidden_states,
1341
+ return_dict=return_dict,
1342
+ )
1343
+
1344
+ text_outputs = self.text_model(
1345
+ input_ids=input_ids,
1346
+ attention_mask=attention_mask,
1347
+ position_ids=position_ids,
1348
+ output_attentions=output_attentions,
1349
+ output_hidden_states=output_hidden_states,
1350
+ return_dict=return_dict,
1351
+ )
1352
+
1353
+ image_embeds = vision_outputs[1]
1354
+ image_embeds = self.visual_projection(image_embeds)
1355
+
1356
+ text_embeds = text_outputs[1]
1357
+ text_embeds = self.text_projection(text_embeds)
1358
+
1359
+ # normalized features
1360
+ image_embeds = image_embeds / image_embeds.norm(dim=-1, keepdim=True)
1361
+ text_embeds = text_embeds / text_embeds.norm(dim=-1, keepdim=True)
1362
+
1363
+ # cosine similarity as logits
1364
+ logit_scale = self.logit_scale.exp()
1365
+ logits_per_text = torch.matmul(text_embeds, image_embeds.t()) * logit_scale
1366
+ logits_per_image = logits_per_text.t()
1367
+
1368
+ seg_logits = None
1369
+ if output_segmentation:
1370
+ # grouped features
1371
+ # [batch_size_image, num_group, hidden_size]
1372
+ image_group_embeds = vision_outputs[0]
1373
+ # [batch_size_image*num_group, hidden_size]
1374
+ image_group_embeds = self.visual_projection(image_group_embeds.reshape(-1, image_group_embeds.shape[-1]))
1375
+ if output_hidden_states:
1376
+ attentions = vision_outputs[3]
1377
+ else:
1378
+ attentions = vision_outputs[2]
1379
+ # [batch_size_image, num_group, height, width]
1380
+ grouping = get_grouping_from_attentions(attentions, pixel_values.shape[2:])
1381
+
1382
+ # normalized features
1383
+ image_group_embeds = image_group_embeds / image_group_embeds.norm(dim=-1, keepdim=True)
1384
+ # [batch_size_image x num_group, batch_size_text]
1385
+ logits_per_image_group = torch.matmul(image_group_embeds, text_embeds.t()) * logit_scale
1386
+ # [batch_size_image, batch_size_text, num_group]
1387
+ logits_per_image_group = logits_per_image_group.reshape(
1388
+ image_embeds.shape[0], -1, text_embeds.shape[0]
1389
+ ).permute(0, 2, 1)
1390
+
1391
+ # [batch_size_image, batch_size_text, height x width]
1392
+ flatten_grouping = grouping.reshape(grouping.shape[0], grouping.shape[1], -1)
1393
+
1394
+ # [batch_size_image, batch_size_text, height, width]
1395
+ seg_logits = torch.matmul(logits_per_image_group, flatten_grouping) * logit_scale
1396
+ seg_logits = seg_logits.reshape(
1397
+ seg_logits.shape[0], seg_logits.shape[1], grouping.shape[2], grouping.shape[3]
1398
+ )
1399
+
1400
+ loss = None
1401
+ if return_loss:
1402
+ loss = groupvit_loss(logits_per_text)
1403
+
1404
+ if not return_dict:
1405
+ if seg_logits is not None:
1406
+ output = (
1407
+ logits_per_image,
1408
+ logits_per_text,
1409
+ seg_logits,
1410
+ text_embeds,
1411
+ image_embeds,
1412
+ text_outputs,
1413
+ vision_outputs,
1414
+ )
1415
+ else:
1416
+ output = (logits_per_image, logits_per_text, text_embeds, image_embeds, text_outputs, vision_outputs)
1417
+ return ((loss,) + output) if loss is not None else output
1418
+
1419
+ return GroupViTModelOutput(
1420
+ loss=loss,
1421
+ logits_per_image=logits_per_image,
1422
+ logits_per_text=logits_per_text,
1423
+ segmentation_logits=seg_logits,
1424
+ text_embeds=text_embeds,
1425
+ image_embeds=image_embeds,
1426
+ text_model_output=text_outputs,
1427
+ vision_model_output=vision_outputs,
1428
+ )
1429
+
1430
+
1431
+ __all__ = ["GroupViTModel", "GroupViTPreTrainedModel", "GroupViTTextModel", "GroupViTVisionModel"]