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

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

Browse files
edit//Qwen3-TTS-test//.venv//Lib//site-packages//transformers//models//groupvit//modeling_tf_groupvit.py ADDED
@@ -0,0 +1,2141 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
+ """TF 2.0 GroupViT model."""
16
+
17
+ from __future__ import annotations
18
+
19
+ import collections.abc
20
+ import math
21
+ from dataclasses import dataclass
22
+ from typing import Any
23
+
24
+ import numpy as np
25
+ import tensorflow as tf
26
+
27
+ from ...activations_tf import get_tf_activation
28
+ from ...modeling_tf_outputs import TFBaseModelOutput, TFBaseModelOutputWithPooling
29
+ from ...modeling_tf_utils import (
30
+ TFModelInputType,
31
+ TFPreTrainedModel,
32
+ get_initializer,
33
+ keras,
34
+ keras_serializable,
35
+ unpack_inputs,
36
+ )
37
+ from ...tf_utils import check_embeddings_within_bounds, shape_list, stable_softmax
38
+ from ...utils import (
39
+ ModelOutput,
40
+ add_start_docstrings,
41
+ add_start_docstrings_to_model_forward,
42
+ is_tensorflow_probability_available,
43
+ logging,
44
+ replace_return_docstrings,
45
+ )
46
+ from .configuration_groupvit import GroupViTConfig, GroupViTTextConfig, GroupViTVisionConfig
47
+
48
+
49
+ logger = logging.get_logger(__name__)
50
+
51
+ # soft dependency
52
+ if is_tensorflow_probability_available():
53
+ try:
54
+ import tensorflow_probability as tfp
55
+
56
+ # On the first call, check whether a compatible version of TensorFlow is installed
57
+ # TensorFlow Probability depends on a recent stable release of TensorFlow
58
+ _ = tfp.distributions.Normal(loc=0.0, scale=1.0)
59
+ except ImportError:
60
+ logger.error(
61
+ "GroupViT models are not usable since `tensorflow_probability` can't be loaded. "
62
+ "It seems you have `tensorflow_probability` installed with the wrong tensorflow version."
63
+ "Please try to reinstall it following the instructions here: https://github.com/tensorflow/probability."
64
+ )
65
+ else:
66
+ try:
67
+ import tensorflow_probability as tfp
68
+
69
+ # On the first call, check whether a compatible version of TensorFlow is installed
70
+ # TensorFlow Probability depends on a recent stable release of TensorFlow
71
+ _ = tfp.distributions.Normal(loc=0.0, scale=1.0)
72
+ except ImportError:
73
+ pass
74
+
75
+ _CHECKPOINT_FOR_DOC = "nvidia/groupvit-gcc-yfcc"
76
+
77
+
78
+ LARGE_NEGATIVE = -1e8
79
+
80
+
81
+ # Copied from transformers.models.bart.modeling_tf_bart._expand_mask
82
+ def _expand_mask(mask: tf.Tensor, tgt_len: int | None = None):
83
+ """
84
+ Expands attention_mask from `[bsz, seq_len]` to `[bsz, 1, tgt_seq_len, src_seq_len]`.
85
+ """
86
+ src_len = shape_list(mask)[1]
87
+ tgt_len = tgt_len if tgt_len is not None else src_len
88
+ one_cst = tf.constant(1.0)
89
+ mask = tf.cast(mask, dtype=one_cst.dtype)
90
+ expanded_mask = tf.tile(mask[:, None, None, :], (1, 1, tgt_len, 1))
91
+
92
+ return (one_cst - expanded_mask) * LARGE_NEGATIVE
93
+
94
+
95
+ # contrastive loss function, adapted from
96
+ # https://sachinruk.github.io/blog/pytorch/pytorch%20lightning/loss%20function/gpu/2021/03/07/CLIP.html
97
+ def contrastive_loss(logits: tf.Tensor) -> tf.Tensor:
98
+ return tf.math.reduce_mean(
99
+ keras.metrics.sparse_categorical_crossentropy(
100
+ y_true=tf.range(shape_list(logits)[0]), y_pred=logits, from_logits=True
101
+ )
102
+ )
103
+
104
+
105
+ # Copied from transformers.models.clip.modeling_tf_clip.clip_loss with clip->groupvit
106
+ def groupvit_loss(similarity: tf.Tensor) -> tf.Tensor:
107
+ caption_loss = contrastive_loss(similarity)
108
+ image_loss = contrastive_loss(tf.transpose(similarity))
109
+ return (caption_loss + image_loss) / 2.0
110
+
111
+
112
+ def hard_softmax(logits: tf.Tensor, dim: int) -> tf.Tensor:
113
+ y_soft = stable_softmax(logits, dim)
114
+ # Straight through.
115
+ index = tf.argmax(y_soft, dim)
116
+ y_hard = tf.one_hot(
117
+ index,
118
+ depth=shape_list(logits)[dim],
119
+ # TensorFlow expects axis to be -1 or between [0, 3). But received: -2
120
+ # This is why the following code snippet is used.
121
+ axis=range(len(shape_list(logits)))[dim],
122
+ dtype=y_soft.dtype,
123
+ )
124
+ ret = y_hard - tf.stop_gradient(y_soft) + y_soft
125
+
126
+ return ret
127
+
128
+
129
+ def gumbel_softmax(logits: tf.Tensor, tau: float = 1, hard: bool = False, dim: int = -1) -> tf.Tensor:
130
+ gumbel_dist = tfp.distributions.Gumbel(0.0, 1.0)
131
+ gumbels = gumbel_dist.sample(tf.shape(logits), dtype=logits.dtype)
132
+
133
+ gumbels = (logits + gumbels) / tau # ~Gumbel(logits,tau)
134
+ y_soft = stable_softmax(gumbels, dim)
135
+
136
+ if hard:
137
+ # Straight through.
138
+ index = tf.argmax(y_soft, dim)
139
+ y_hard = tf.one_hot(
140
+ index,
141
+ depth=shape_list(logits)[dim],
142
+ # TensorFlow expects axis to be -1 or between [0, 3). But received: -2
143
+ # This is why the following code snippet is used.
144
+ axis=range(len(shape_list(logits)))[dim],
145
+ dtype=y_soft.dtype,
146
+ )
147
+ ret = y_hard - tf.stop_gradient(y_soft) + y_soft
148
+ else:
149
+ # Reparametrization trick.
150
+ ret = y_soft
151
+ return ret
152
+
153
+
154
+ def resize_attention_map(attentions: tf.Tensor, height: int, width: int, align_corners: bool = False) -> tf.Tensor:
155
+ """
156
+ Args:
157
+ attentions (`tf.Tensor`): attention map of shape [batch_size, groups, feat_height*feat_width]
158
+ height (`int`): height of the output attention map
159
+ width (`int`): width of the output attention map
160
+ align_corners (`bool`, *optional*): the `align_corner` argument for `nn.functional.interpolate`.
161
+
162
+ Returns:
163
+ `tf.Tensor`: resized attention map of shape [batch_size, groups, height, width]
164
+ """
165
+
166
+ scale = (height * width // attentions.shape[2]) ** 0.5
167
+ if height > width:
168
+ feat_width = int(np.round(width / scale))
169
+ feat_height = shape_list(attentions)[2] // feat_width
170
+ else:
171
+ feat_height = int(np.round(height / scale))
172
+ feat_width = shape_list(attentions)[2] // feat_height
173
+
174
+ batch_size = shape_list(attentions)[0]
175
+ groups = shape_list(attentions)[1] # number of group token
176
+ # [batch_size, groups, height x width, groups] -> [batch_size, groups, height, width]
177
+ attentions = tf.reshape(attentions, (batch_size, groups, feat_height, feat_width))
178
+ attentions = tf.transpose(attentions, perm=(0, 2, 3, 1))
179
+ if align_corners:
180
+ attentions = tf.compat.v1.image.resize(
181
+ attentions,
182
+ size=(height, width),
183
+ method="bilinear",
184
+ align_corners=align_corners,
185
+ )
186
+ else:
187
+ attentions = tf.image.resize(attentions, size=(height, width), method="bilinear")
188
+ attentions = tf.transpose(attentions, perm=(0, 3, 1, 2))
189
+ return attentions
190
+
191
+
192
+ def get_grouping_from_attentions(attentions: tuple[tf.Tensor], hw_shape: tuple[int]) -> tf.Tensor:
193
+ """
194
+ Args:
195
+ attentions (`tuple(tf.Tensor)`: tuple of attention maps returned by `TFGroupViTVisionTransformer`
196
+ hw_shape (`tuple(int)`): height and width of the output attention map
197
+ Returns:
198
+ `tf.Tensor`: the attention map of shape [batch_size, groups, height, width]
199
+ """
200
+
201
+ attn_maps = []
202
+ prev_attn_masks = None
203
+ for attn_masks in attentions:
204
+ # [batch_size, num_groups, height x width] -> [batch_size, height x width, num_groups]
205
+ attn_masks = tf.transpose(attn_masks, perm=(0, 2, 1))
206
+ if prev_attn_masks is None:
207
+ prev_attn_masks = attn_masks
208
+ else:
209
+ prev_attn_masks = tf.matmul(prev_attn_masks, attn_masks)
210
+ # [batch_size, height x width, num_groups] -> [batch_size, num_groups, height x width] -> [batch_size, num_groups, height, width]
211
+ cur_attn_map = resize_attention_map(tf.transpose(prev_attn_masks, perm=(0, 2, 1)), *hw_shape)
212
+ attn_maps.append(cur_attn_map)
213
+
214
+ # [batch_size, num_groups, height, width]
215
+ final_grouping = attn_maps[-1]
216
+
217
+ return tf.stop_gradient(final_grouping)
218
+
219
+
220
+ @dataclass
221
+ class TFGroupViTModelOutput(ModelOutput):
222
+ """
223
+ Args:
224
+ loss (`tf.Tensor` of shape `(1,)`, *optional*, returned when `return_loss` is `True`):
225
+ Contrastive loss for image-text similarity.
226
+ logits_per_image (`tf.Tensor` of shape `(image_batch_size, text_batch_size)`):
227
+ The scaled dot product scores between `image_embeds` and `text_embeds`. This represents the image-text
228
+ similarity scores.
229
+ logits_per_text (`tf.Tensor` of shape `(text_batch_size, image_batch_size)`):
230
+ The scaled dot product scores between `text_embeds` and `image_embeds`. This represents the text-image
231
+ similarity scores.
232
+ segmentation_logits (`tf.Tensor` of shape `(batch_size, config.num_labels, logits_height, logits_width)`):
233
+ Classification scores for each pixel.
234
+
235
+ <Tip warning={true}>
236
+
237
+ The logits returned do not necessarily have the same size as the `pixel_values` passed as inputs. This is
238
+ to avoid doing two interpolations and lose some quality when a user needs to resize the logits to the
239
+ original image size as post-processing. You should always check your logits shape and resize as needed.
240
+
241
+ </Tip>
242
+
243
+ text_embeds (`tf.Tensor` of shape `(batch_size, output_dim`):
244
+ The text embeddings obtained by applying the projection layer to the pooled output of
245
+ [`TFGroupViTTextModel`].
246
+ image_embeds (`tf.Tensor` of shape `(batch_size, output_dim`):
247
+ The image embeddings obtained by applying the projection layer to the pooled output of
248
+ [`TFGroupViTVisionModel`].
249
+ text_model_output (`TFBaseModelOutputWithPooling`):
250
+ The output of the [`TFGroupViTTextModel`].
251
+ vision_model_output (`TFBaseModelOutputWithPooling`):
252
+ The output of the [`TFGroupViTVisionModel`].
253
+ """
254
+
255
+ loss: tf.Tensor | None = None
256
+ logits_per_image: tf.Tensor | None = None
257
+ logits_per_text: tf.Tensor | None = None
258
+ segmentation_logits: tf.Tensor | None = None
259
+ text_embeds: tf.Tensor | None = None
260
+ image_embeds: tf.Tensor | None = None
261
+ text_model_output: TFBaseModelOutputWithPooling = None
262
+ vision_model_output: TFBaseModelOutputWithPooling = None
263
+
264
+ def to_tuple(self) -> tuple[Any]:
265
+ return tuple(
266
+ self[k] if k not in ["text_model_output", "vision_model_output"] else getattr(self, k).to_tuple()
267
+ for k in self.keys()
268
+ )
269
+
270
+
271
+ class TFGroupViTCrossAttentionLayer(keras.layers.Layer):
272
+ def __init__(self, config: GroupViTVisionConfig, **kwargs):
273
+ super().__init__(**kwargs)
274
+ self.attn = TFGroupViTAttention(config, name="attn")
275
+ self.norm2 = keras.layers.LayerNormalization(epsilon=config.layer_norm_eps, name="norm2")
276
+ self.mlp = TFGroupViTMLP(config, name="mlp")
277
+ self.norm_post = keras.layers.LayerNormalization(epsilon=config.layer_norm_eps, name="norm_post")
278
+ self.config = config
279
+
280
+ def call(self, query: tf.Tensor, key: tf.Tensor, training: bool = False) -> tf.Tensor:
281
+ x = query
282
+ x = x + self.attn(query, encoder_hidden_states=key)[0]
283
+ x = x + self.mlp(self.norm2(x))
284
+ x = self.norm_post(x)
285
+ return x
286
+
287
+ def build(self, input_shape=None):
288
+ if self.built:
289
+ return
290
+ self.built = True
291
+ if getattr(self, "attn", None) is not None:
292
+ with tf.name_scope(self.attn.name):
293
+ self.attn.build(None)
294
+ if getattr(self, "norm2", None) is not None:
295
+ with tf.name_scope(self.norm2.name):
296
+ self.norm2.build([None, None, self.config.hidden_size])
297
+ if getattr(self, "mlp", None) is not None:
298
+ with tf.name_scope(self.mlp.name):
299
+ self.mlp.build(None)
300
+ if getattr(self, "norm_post", None) is not None:
301
+ with tf.name_scope(self.norm_post.name):
302
+ self.norm_post.build([None, None, self.config.hidden_size])
303
+
304
+
305
+ class TFGroupViTAssignAttention(keras.layers.Layer):
306
+ def __init__(self, config: GroupViTVisionConfig, **kwargs):
307
+ super().__init__(**kwargs)
308
+ self.scale = config.hidden_size**-0.5
309
+
310
+ self.q_proj = keras.layers.Dense(config.hidden_size, name="q_proj")
311
+ self.k_proj = keras.layers.Dense(config.hidden_size, name="k_proj")
312
+ self.v_proj = keras.layers.Dense(config.hidden_size, name="v_proj")
313
+ self.proj = keras.layers.Dense(config.hidden_size, name="proj")
314
+ self.assign_eps = config.assign_eps
315
+ self.config = config
316
+
317
+ def get_attn(self, attn: tf.Tensor, gumbel: bool = True, hard: bool = True, training: bool = False) -> tf.Tensor:
318
+ if gumbel and training:
319
+ attn = gumbel_softmax(attn, dim=-2, hard=hard)
320
+ else:
321
+ if hard:
322
+ attn = hard_softmax(attn, dim=-2)
323
+ else:
324
+ attn = stable_softmax(attn, axis=-2)
325
+
326
+ return attn
327
+
328
+ def call(self, query: tf.Tensor, key: tf.Tensor, training: bool = False):
329
+ value = key
330
+ # [batch_size, query_length, channels]
331
+ query = self.q_proj(query)
332
+
333
+ # [batch_size, key_length, channels]
334
+ key = self.k_proj(key)
335
+
336
+ # [batch_size, key_length, channels]
337
+ value = self.v_proj(value)
338
+
339
+ # [batch_size, query_length, key_length]
340
+ raw_attn = tf.matmul(query, key, transpose_b=True) * self.scale
341
+
342
+ attn = self.get_attn(raw_attn, training=training)
343
+ soft_attn = self.get_attn(raw_attn, training=training, gumbel=False, hard=False)
344
+
345
+ attn = attn / (tf.math.reduce_sum(attn, axis=-1, keepdims=True) + self.assign_eps)
346
+
347
+ out = tf.matmul(attn, value)
348
+
349
+ out = self.proj(out)
350
+
351
+ return out, soft_attn
352
+
353
+ def build(self, input_shape=None):
354
+ if self.built:
355
+ return
356
+ self.built = True
357
+ if getattr(self, "q_proj", None) is not None:
358
+ with tf.name_scope(self.q_proj.name):
359
+ self.q_proj.build([None, None, self.config.hidden_size])
360
+ if getattr(self, "k_proj", None) is not None:
361
+ with tf.name_scope(self.k_proj.name):
362
+ self.k_proj.build([None, None, self.config.hidden_size])
363
+ if getattr(self, "v_proj", None) is not None:
364
+ with tf.name_scope(self.v_proj.name):
365
+ self.v_proj.build([None, None, self.config.hidden_size])
366
+ if getattr(self, "proj", None) is not None:
367
+ with tf.name_scope(self.proj.name):
368
+ self.proj.build([None, None, self.config.hidden_size])
369
+
370
+
371
+ class TFGroupViTTokenAssign(keras.layers.Layer):
372
+ def __init__(self, config: GroupViTVisionConfig, num_group_token: int, num_output_group: int, **kwargs):
373
+ super().__init__(**kwargs)
374
+ self.num_output_group = num_output_group
375
+ # norm on group_tokens
376
+ self.norm_tokens = keras.layers.LayerNormalization(epsilon=config.layer_norm_eps, name="norm_tokens")
377
+ assign_mlp_ratio = (
378
+ config.assign_mlp_ratio
379
+ if isinstance(config.assign_mlp_ratio, collections.abc.Iterable)
380
+ else (config.assign_mlp_ratio, config.assign_mlp_ratio)
381
+ )
382
+ tokens_dim, channels_dim = [int(x * config.hidden_size) for x in assign_mlp_ratio]
383
+ self.mlp_inter = TFGroupViTMixerMLP(config, num_group_token, tokens_dim, num_output_group, name="mlp_inter")
384
+ self.norm_post_tokens = keras.layers.LayerNormalization(epsilon=config.layer_norm_eps, name="norm_post_tokens")
385
+ # norm on x
386
+ self.norm_x = keras.layers.LayerNormalization(epsilon=config.layer_norm_eps, name="norm_x")
387
+ self.pre_assign_attn = TFGroupViTCrossAttentionLayer(config, name="pre_assign_attn")
388
+
389
+ self.assign = TFGroupViTAssignAttention(config, name="assign")
390
+ self.norm_new_x = keras.layers.LayerNormalization(epsilon=config.layer_norm_eps, name="norm_new_x")
391
+ self.mlp_channels = TFGroupViTMLP(
392
+ config, config.hidden_size, channels_dim, config.hidden_size, name="mlp_channels"
393
+ )
394
+ self.config = config
395
+
396
+ def project_group_token(self, group_tokens: tf.Tensor) -> tf.Tensor:
397
+ """
398
+ Args:
399
+ group_tokens (tf.Tensor): group tokens, [batch_size, num_group_tokens, channels]
400
+
401
+ Returns:
402
+ projected_group_tokens (tf.Tensor): [batch_size, num_output_groups, channels]
403
+ """
404
+ # [B, num_output_groups, C] <- [B, num_group_tokens, C]
405
+ projected_group_tokens = self.mlp_inter(group_tokens)
406
+ projected_group_tokens = self.norm_post_tokens(projected_group_tokens)
407
+ return projected_group_tokens
408
+
409
+ def call(self, image_tokens: tf.Tensor, group_tokens: tf.Tensor, training: bool = False):
410
+ """
411
+ Args:
412
+ image_tokens (`tf.Tensor`): image tokens, of shape [batch_size, input_length, channels]
413
+ group_tokens (`tf.Tensor`): group tokens, [batch_size, num_group_tokens, channels]
414
+ """
415
+
416
+ group_tokens = self.norm_tokens(group_tokens)
417
+ image_tokens = self.norm_x(image_tokens)
418
+ # [batch_size, num_output_groups, channels]
419
+ projected_group_tokens = self.project_group_token(group_tokens)
420
+ projected_group_tokens = self.pre_assign_attn(projected_group_tokens, image_tokens)
421
+ new_image_tokens, attention = self.assign(projected_group_tokens, image_tokens)
422
+ new_image_tokens += projected_group_tokens
423
+
424
+ new_image_tokens = new_image_tokens + self.mlp_channels(self.norm_new_x(new_image_tokens))
425
+
426
+ return new_image_tokens, attention
427
+
428
+ def build(self, input_shape=None):
429
+ if self.built:
430
+ return
431
+ self.built = True
432
+ if getattr(self, "norm_tokens", None) is not None:
433
+ with tf.name_scope(self.norm_tokens.name):
434
+ self.norm_tokens.build([None, None, self.config.hidden_size])
435
+ if getattr(self, "mlp_inter", None) is not None:
436
+ with tf.name_scope(self.mlp_inter.name):
437
+ self.mlp_inter.build(None)
438
+ if getattr(self, "norm_post_tokens", None) is not None:
439
+ with tf.name_scope(self.norm_post_tokens.name):
440
+ self.norm_post_tokens.build([None, None, self.config.hidden_size])
441
+ if getattr(self, "norm_x", None) is not None:
442
+ with tf.name_scope(self.norm_x.name):
443
+ self.norm_x.build([None, None, self.config.hidden_size])
444
+ if getattr(self, "pre_assign_attn", None) is not None:
445
+ with tf.name_scope(self.pre_assign_attn.name):
446
+ self.pre_assign_attn.build(None)
447
+ if getattr(self, "assign", None) is not None:
448
+ with tf.name_scope(self.assign.name):
449
+ self.assign.build(None)
450
+ if getattr(self, "norm_new_x", None) is not None:
451
+ with tf.name_scope(self.norm_new_x.name):
452
+ self.norm_new_x.build([None, None, self.config.hidden_size])
453
+ if getattr(self, "mlp_channels", None) is not None:
454
+ with tf.name_scope(self.mlp_channels.name):
455
+ self.mlp_channels.build(None)
456
+
457
+
458
+ # Adapted from transformers.models.vit.modeling_tf_vit.TFViTPatchEmbeddings with ViT->GroupViT
459
+ class TFGroupViTPatchEmbeddings(keras.layers.Layer):
460
+ """
461
+ This class turns `pixel_values` of shape `(batch_size, num_channels, height, width)` into the initial
462
+ `hidden_states` (patch embeddings) of shape `(batch_size, seq_length, hidden_size)` to be consumed by a
463
+ Transformer.
464
+ """
465
+
466
+ def __init__(self, config: GroupViTConfig, **kwargs):
467
+ super().__init__(**kwargs)
468
+ image_size, patch_size = config.image_size, config.patch_size
469
+ num_channels = config.num_channels
470
+ # hidden_size is a member as it will be required in the call method
471
+ self.hidden_size = config.hidden_size
472
+
473
+ image_size = image_size if isinstance(image_size, collections.abc.Iterable) else (image_size, image_size)
474
+ patch_size = patch_size if isinstance(patch_size, collections.abc.Iterable) else (patch_size, patch_size)
475
+ num_patches = (image_size[1] // patch_size[1]) * (image_size[0] // patch_size[0])
476
+ self.image_size = image_size
477
+ self.patch_size = patch_size
478
+ self.num_patches = num_patches
479
+ self.num_channels = num_channels
480
+ self.config = config
481
+
482
+ self.projection = keras.layers.Conv2D(
483
+ filters=self.hidden_size,
484
+ kernel_size=patch_size,
485
+ strides=patch_size,
486
+ padding="valid",
487
+ data_format="channels_last",
488
+ use_bias=True,
489
+ kernel_initializer=get_initializer(self.config.initializer_range),
490
+ bias_initializer="zeros",
491
+ name="projection",
492
+ )
493
+
494
+ def call(
495
+ self, pixel_values: tf.Tensor, interpolate_pos_encoding: bool = False, training: bool = False
496
+ ) -> tf.Tensor:
497
+ batch_size, num_channels, height, width = shape_list(pixel_values)
498
+ if tf.executing_eagerly() and num_channels != self.num_channels:
499
+ raise ValueError(
500
+ "Make sure that the channel dimension of the pixel values match with the one set in the configuration."
501
+ )
502
+ if (
503
+ not interpolate_pos_encoding
504
+ and tf.executing_eagerly()
505
+ and (height != self.image_size[0] or width != self.image_size[1])
506
+ ):
507
+ raise ValueError(
508
+ f"Input image size ({height}*{width}) doesn't match model ({self.image_size[0]}*{self.image_size[1]})."
509
+ )
510
+
511
+ # When running on CPU, `keras.layers.Conv2D` doesn't support `NCHW` format.
512
+ # So change the input format from `NCHW` to `NHWC`.
513
+ # shape = (batch_size, in_height, in_width, in_channels=num_channels)
514
+ pixel_values = tf.transpose(pixel_values, perm=(0, 2, 3, 1))
515
+
516
+ projection = self.projection(pixel_values)
517
+
518
+ # Change the 2D spatial dimensions to a single temporal dimension.
519
+ # shape = (batch_size, num_patches, out_channels=embed_dim)
520
+ num_patches = (width // self.patch_size[1]) * (height // self.patch_size[0])
521
+ # In the TFGroupViTVisionEmbeddings the embeddings from this layer will be layer normalized
522
+ # LayerNormalization layer needs to have static last dimension (otherwise the test_keras_save_load fails with symbolic tensors)
523
+ # This is why we have used the hidden_size in the reshape method
524
+ embeddings = tf.reshape(tensor=projection, shape=(batch_size, num_patches, self.hidden_size))
525
+
526
+ return embeddings
527
+
528
+ def build(self, input_shape=None):
529
+ if self.built:
530
+ return
531
+ self.built = True
532
+ if getattr(self, "projection", None) is not None:
533
+ with tf.name_scope(self.projection.name):
534
+ self.projection.build([None, None, None, self.num_channels])
535
+
536
+
537
+ # Adapted from transformers.vit.modeling_tf_vit.TFViTEmbeddings
538
+ class TFGroupViTVisionEmbeddings(keras.layers.Layer):
539
+ """
540
+ Construct the position and patch embeddings.
541
+
542
+ """
543
+
544
+ def __init__(self, config: GroupViTVisionConfig, **kwargs):
545
+ super().__init__(**kwargs)
546
+
547
+ self.patch_embeddings = TFGroupViTPatchEmbeddings(config, name="patch_embeddings")
548
+ self.dropout = keras.layers.Dropout(rate=config.dropout, name="dropout")
549
+ self.layernorm = keras.layers.LayerNormalization(epsilon=config.layer_norm_eps, name="layernorm")
550
+ self.config = config
551
+
552
+ def build(self, input_shape=None):
553
+ num_patches = self.patch_embeddings.num_patches
554
+ self.position_embeddings = self.add_weight(
555
+ shape=(1, num_patches, self.config.hidden_size),
556
+ initializer="zeros",
557
+ trainable=True,
558
+ name="position_embeddings",
559
+ )
560
+
561
+ if self.built:
562
+ return
563
+ self.built = True
564
+ if getattr(self, "patch_embeddings", None) is not None:
565
+ with tf.name_scope(self.patch_embeddings.name):
566
+ self.patch_embeddings.build(None)
567
+ if getattr(self, "dropout", None) is not None:
568
+ with tf.name_scope(self.dropout.name):
569
+ self.dropout.build(None)
570
+ if getattr(self, "layernorm", None) is not None:
571
+ with tf.name_scope(self.layernorm.name):
572
+ self.layernorm.build([None, None, self.config.hidden_size])
573
+
574
+ def interpolate_pos_encoding(self, embeddings, height, width) -> tf.Tensor:
575
+ """
576
+ This method allows to interpolate the pre-trained position encodings, to be able to use the model on higher
577
+ resolution images.
578
+
579
+ Source:
580
+ https://github.com/facebookresearch/dino/blob/de9ee3df6cf39fac952ab558447af1fa1365362a/vision_transformer.py#L174
581
+ """
582
+
583
+ batch_size, num_patches, dim = shape_list(embeddings)
584
+ num_positions = shape_list(self.position_embeddings)[1]
585
+
586
+ if num_patches == num_positions and height == width:
587
+ return self.position_embeddings
588
+ patch_pos_embed = self.position_embeddings
589
+ h0 = height // self.config.patch_size
590
+ w0 = width // self.config.patch_size
591
+ patch_pos_embed = tf.image.resize(
592
+ images=tf.reshape(
593
+ patch_pos_embed, shape=(1, int(math.sqrt(num_positions)), int(math.sqrt(num_positions)), dim)
594
+ ),
595
+ size=(h0, w0),
596
+ method="bicubic",
597
+ )
598
+ patch_pos_embed = tf.reshape(tensor=patch_pos_embed, shape=(1, -1, dim))
599
+ return patch_pos_embed
600
+
601
+ def call(
602
+ self, pixel_values: tf.Tensor, interpolate_pos_encoding: bool = False, training: bool = False
603
+ ) -> tf.Tensor:
604
+ _, _, height, width = shape_list(pixel_values)
605
+ embeddings = self.patch_embeddings(pixel_values, interpolate_pos_encoding=interpolate_pos_encoding)
606
+ embeddings = self.layernorm(embeddings)
607
+
608
+ # add positional encoding to each token
609
+ if interpolate_pos_encoding:
610
+ embeddings = embeddings + self.interpolate_pos_encoding(embeddings, height, width)
611
+ else:
612
+ embeddings = embeddings + self.position_embeddings
613
+
614
+ embeddings = self.dropout(embeddings)
615
+
616
+ return embeddings
617
+
618
+
619
+ # Copied from transformers.models.clip.modeling_tf_clip.TFCLIPTextEmbeddings with CLIP->GroupViT
620
+ class TFGroupViTTextEmbeddings(keras.layers.Layer):
621
+ def __init__(self, config: GroupViTTextConfig, **kwargs):
622
+ super().__init__(**kwargs)
623
+
624
+ self.embed_dim = config.hidden_size
625
+
626
+ self.config = config
627
+
628
+ def build(self, input_shape: tf.TensorShape = None):
629
+ with tf.name_scope("token_embedding"):
630
+ self.weight = self.add_weight(
631
+ shape=(self.config.vocab_size, self.embed_dim),
632
+ initializer=get_initializer(self.config.initializer_factor * self.config.initializer_range),
633
+ trainable=True,
634
+ name="weight",
635
+ )
636
+
637
+ with tf.name_scope("position_embedding"):
638
+ self.position_embedding = self.add_weight(
639
+ shape=(self.config.max_position_embeddings, self.embed_dim),
640
+ initializer=get_initializer(self.config.initializer_factor * self.config.initializer_range),
641
+ trainable=True,
642
+ name="embeddings",
643
+ )
644
+
645
+ super().build(input_shape)
646
+
647
+ def call(
648
+ self,
649
+ input_ids: tf.Tensor | None = None,
650
+ position_ids: tf.Tensor | None = None,
651
+ inputs_embeds: tf.Tensor | None = None,
652
+ ) -> tf.Tensor:
653
+ """
654
+ Applies embedding based on inputs tensor.
655
+
656
+ Returns:
657
+ final_embeddings (`tf.Tensor`): output embedding tensor.
658
+ """
659
+ if input_ids is None and inputs_embeds is None:
660
+ raise ValueError("You have to specify either input_ids or inputs_embeds")
661
+
662
+ if inputs_embeds is None:
663
+ check_embeddings_within_bounds(input_ids, self.config.vocab_size)
664
+ inputs_embeds = tf.gather(params=self.weight, indices=input_ids)
665
+
666
+ input_shape = shape_list(inputs_embeds)[:-1]
667
+
668
+ if position_ids is None:
669
+ position_ids = tf.expand_dims(tf.range(start=0, limit=input_shape[-1]), axis=0)
670
+
671
+ position_embeds = tf.gather(params=self.position_embedding, indices=position_ids)
672
+ position_embeds = tf.tile(input=position_embeds, multiples=(input_shape[0], 1, 1))
673
+ final_embeddings = inputs_embeds + position_embeds
674
+
675
+ return final_embeddings
676
+
677
+
678
+ class TFGroupViTStage(keras.layers.Layer):
679
+ """This corresponds to the `GroupingLayer` class in the GroupViT implementation."""
680
+
681
+ def __init__(
682
+ self,
683
+ config: GroupViTVisionConfig,
684
+ depth: int,
685
+ num_prev_group_token: int,
686
+ num_group_token: int,
687
+ num_output_group: int,
688
+ **kwargs,
689
+ ):
690
+ super().__init__(**kwargs)
691
+ self.config = config
692
+ self.depth = depth
693
+ self.num_group_token = num_group_token
694
+ self.layers = [TFGroupViTEncoderLayer(config, name=f"layers_._{i}") for i in range(depth)]
695
+
696
+ if num_group_token > 0:
697
+ self.downsample = TFGroupViTTokenAssign(
698
+ config=config,
699
+ num_group_token=num_group_token,
700
+ num_output_group=num_output_group,
701
+ name="downsample",
702
+ )
703
+ else:
704
+ self.downsample = None
705
+
706
+ if num_prev_group_token > 0 and num_group_token > 0:
707
+ self.group_projector = [
708
+ keras.layers.LayerNormalization(epsilon=config.layer_norm_eps, name="group_projector.0"),
709
+ TFGroupViTMixerMLP(
710
+ config, num_prev_group_token, config.hidden_size // 2, num_group_token, name="group_projector.1"
711
+ ),
712
+ ]
713
+ else:
714
+ self.group_projector = None
715
+
716
+ def build(self, input_shape=None):
717
+ if self.num_group_token > 0:
718
+ self.group_token = self.add_weight(
719
+ shape=(1, self.num_group_token, self.config.hidden_size),
720
+ initializer="zeros",
721
+ trainable=True,
722
+ name="group_token",
723
+ )
724
+ else:
725
+ self.group_token = None
726
+
727
+ if self.built:
728
+ return
729
+ self.built = True
730
+ if getattr(self, "downsample", None) is not None:
731
+ with tf.name_scope(self.downsample.name):
732
+ self.downsample.build(None)
733
+ if getattr(self, "layers", None) is not None:
734
+ for layer in self.layers:
735
+ with tf.name_scope(layer.name):
736
+ layer.build(None)
737
+ if getattr(self, "group_projector", None) is not None:
738
+ with tf.name_scope(self.group_projector[0].name):
739
+ self.group_projector[0].build([None, None, self.config.hidden_size])
740
+ with tf.name_scope(self.group_projector[1].name):
741
+ self.group_projector[1].build(None)
742
+
743
+ @property
744
+ def with_group_token(self):
745
+ return self.group_token is not None
746
+
747
+ def split_x(self, x: tf.Tensor) -> tf.Tensor:
748
+ if self.with_group_token:
749
+ return x[:, : -self.num_group_token], x[:, -self.num_group_token :]
750
+ else:
751
+ return x, None
752
+
753
+ def concat_x(self, x: tf.Tensor, group_token: tf.Tensor | None = None) -> tf.Tensor:
754
+ if group_token is None:
755
+ return x
756
+ return tf.concat([x, group_token], axis=1)
757
+
758
+ def call(
759
+ self,
760
+ hidden_states: tf.Tensor,
761
+ prev_group_token: tf.Tensor | None = None,
762
+ output_attentions: bool = False,
763
+ training: bool = False,
764
+ ) -> tuple[tf.Tensor]:
765
+ """
766
+ Args:
767
+ hidden_states (`tf.Tensor`): input to the layer of shape `(batch, seq_len, embed_dim)`
768
+ attention_mask (`tf.Tensor`): attention mask of size
769
+ `(batch, 1, tgt_len, src_len)` where padding elements are indicated by very large negative values.
770
+ `(config.encoder_attention_heads,)`.
771
+ output_attentions (`bool`, *optional*):
772
+ Whether or not to return the grouping tensors of Grouping block.
773
+ """
774
+ if self.with_group_token:
775
+ group_token = tf.tile(self.group_token, multiples=(shape_list(hidden_states)[0], 1, 1))
776
+ if self.group_projector is not None:
777
+ for layer in self.group_projector:
778
+ prev_group_token = layer(prev_group_token)
779
+ group_token = group_token + prev_group_token
780
+ else:
781
+ group_token = None
782
+
783
+ x = hidden_states
784
+
785
+ cat_x = self.concat_x(x, group_token)
786
+ for layer in self.layers:
787
+ layer_out = layer(
788
+ cat_x,
789
+ attention_mask=None,
790
+ causal_attention_mask=None,
791
+ output_attentions=None,
792
+ )
793
+ cat_x = layer_out[0]
794
+
795
+ x, group_token = self.split_x(cat_x)
796
+
797
+ attention = None
798
+ if self.downsample is not None:
799
+ x, attention = self.downsample(x, group_token)
800
+
801
+ outputs = (x, group_token)
802
+ if output_attentions:
803
+ outputs = outputs + (attention,)
804
+
805
+ return outputs
806
+
807
+
808
+ class TFGroupViTMLP(keras.layers.Layer):
809
+ def __init__(
810
+ self,
811
+ config: GroupViTVisionConfig,
812
+ hidden_size: int | None = None,
813
+ intermediate_size: int | None = None,
814
+ output_size: int | None = None,
815
+ **kwargs,
816
+ ):
817
+ super().__init__(**kwargs)
818
+ self.config = config
819
+ self.activation_fn = get_tf_activation(config.hidden_act)
820
+ hidden_size = hidden_size if hidden_size is not None else config.hidden_size
821
+ intermediate_size = intermediate_size if intermediate_size is not None else config.intermediate_size
822
+ output_size = output_size if output_size is not None else hidden_size
823
+ self.fc1 = keras.layers.Dense(intermediate_size, name="fc1")
824
+ self.fc2 = keras.layers.Dense(output_size, name="fc2")
825
+ self.intermediate_size = intermediate_size
826
+ self.hidden_size = hidden_size
827
+
828
+ def call(self, hidden_states: tf.Tensor, training: bool = False) -> tf.Tensor:
829
+ hidden_states = self.fc1(hidden_states)
830
+ hidden_states = self.activation_fn(hidden_states)
831
+ hidden_states = self.fc2(hidden_states)
832
+ return hidden_states
833
+
834
+ def build(self, input_shape=None):
835
+ if self.built:
836
+ return
837
+ self.built = True
838
+ if getattr(self, "fc1", None) is not None:
839
+ with tf.name_scope(self.fc1.name):
840
+ self.fc1.build([None, None, self.hidden_size])
841
+ if getattr(self, "fc2", None) is not None:
842
+ with tf.name_scope(self.fc2.name):
843
+ self.fc2.build([None, None, self.intermediate_size])
844
+
845
+
846
+ class TFGroupViTMixerMLP(TFGroupViTMLP):
847
+ def call(self, x, training: bool = False):
848
+ x = super().call(hidden_states=tf.transpose(x, perm=(0, 2, 1)))
849
+ return tf.transpose(x, perm=(0, 2, 1))
850
+
851
+
852
+ # Adapted from transformers.models.clip.modeling_tf_clip.TFCLIPAttention
853
+ class TFGroupViTAttention(keras.layers.Layer):
854
+ """Multi-headed attention from 'Attention Is All You Need' paper"""
855
+
856
+ def __init__(self, config: GroupViTConfig, **kwargs):
857
+ super().__init__(**kwargs)
858
+
859
+ self.embed_dim = config.hidden_size
860
+ self.num_attention_heads = config.num_attention_heads
861
+ self.attention_head_size = self.embed_dim // self.num_attention_heads
862
+ if self.attention_head_size * self.num_attention_heads != self.embed_dim:
863
+ raise ValueError(
864
+ f"embed_dim must be divisible by num_heads (got `embed_dim`: {self.embed_dim} and `num_heads`:"
865
+ f" {self.num_attention_heads})."
866
+ )
867
+
868
+ factor = config.initializer_factor
869
+ in_proj_std = (self.embed_dim**-0.5) * ((2 * config.num_hidden_layers) ** -0.5) * factor
870
+ out_proj_std = (self.embed_dim**-0.5) * factor
871
+
872
+ self.sqrt_att_head_size = math.sqrt(self.attention_head_size)
873
+
874
+ self.q_proj = keras.layers.Dense(
875
+ units=self.embed_dim, kernel_initializer=get_initializer(in_proj_std), name="q_proj"
876
+ )
877
+ self.k_proj = keras.layers.Dense(
878
+ units=self.embed_dim, kernel_initializer=get_initializer(in_proj_std), name="k_proj"
879
+ )
880
+ self.v_proj = keras.layers.Dense(
881
+ units=self.embed_dim, kernel_initializer=get_initializer(in_proj_std), name="v_proj"
882
+ )
883
+
884
+ self.dropout = keras.layers.Dropout(rate=config.attention_dropout)
885
+
886
+ self.out_proj = keras.layers.Dense(
887
+ units=self.embed_dim, kernel_initializer=get_initializer(out_proj_std), name="out_proj"
888
+ )
889
+
890
+ # Copied from transformers.models.bert.modeling_tf_bert.TFBertSelfAttention.transpose_for_scores
891
+ def transpose_for_scores(self, tensor: tf.Tensor, batch_size: int) -> tf.Tensor:
892
+ # Reshape from [batch_size, seq_length, all_head_size] to [batch_size, seq_length, num_attention_heads, attention_head_size]
893
+ tensor = tf.reshape(tensor=tensor, shape=(batch_size, -1, self.num_attention_heads, self.attention_head_size))
894
+
895
+ # Transpose the tensor from [batch_size, seq_length, num_attention_heads, attention_head_size] to [batch_size, num_attention_heads, seq_length, attention_head_size]
896
+ return tf.transpose(tensor, perm=[0, 2, 1, 3])
897
+
898
+ def call(
899
+ self,
900
+ hidden_states: tf.Tensor,
901
+ attention_mask: tf.Tensor | None = None,
902
+ causal_attention_mask: tf.Tensor | None = None,
903
+ output_attentions: bool | None = None,
904
+ encoder_hidden_states: tf.Tensor | None = None,
905
+ training: bool = False,
906
+ ) -> tuple[tf.Tensor]:
907
+ """Input shape: Batch x Time x Channel"""
908
+
909
+ batch_size = shape_list(hidden_states)[0]
910
+ is_cross_attention = encoder_hidden_states is not None
911
+
912
+ mixed_query_layer = self.q_proj(inputs=hidden_states)
913
+ if is_cross_attention:
914
+ mixed_key_layer = self.k_proj(inputs=encoder_hidden_states)
915
+ mixed_value_layer = self.v_proj(inputs=encoder_hidden_states)
916
+ else:
917
+ mixed_key_layer = self.k_proj(inputs=hidden_states)
918
+ mixed_value_layer = self.v_proj(inputs=hidden_states)
919
+
920
+ query_layer = self.transpose_for_scores(mixed_query_layer, batch_size)
921
+ key_layer = self.transpose_for_scores(mixed_key_layer, batch_size)
922
+ value_layer = self.transpose_for_scores(mixed_value_layer, batch_size)
923
+
924
+ # Take the dot product between "query" and "key" to get the raw attention scores.
925
+ # (batch size, num_heads, seq_len_q, seq_len_k)
926
+ attention_scores = tf.matmul(query_layer, key_layer, transpose_b=True)
927
+ dk = tf.cast(self.sqrt_att_head_size, dtype=attention_scores.dtype)
928
+ attention_scores = tf.divide(attention_scores, dk)
929
+
930
+ # apply the causal_attention_mask first
931
+ if causal_attention_mask is not None:
932
+ # Apply the causal attention mask (precomputed for all layers in TFCLIPModel call() function)
933
+ attention_scores = tf.add(attention_scores, causal_attention_mask)
934
+
935
+ if attention_mask is not None:
936
+ # Apply the attention mask (precomputed for all layers in TFCLIPModel call() function)
937
+ attention_scores = tf.add(attention_scores, attention_mask)
938
+
939
+ # Normalize the attention scores to probabilities.
940
+ _attention_probs = stable_softmax(logits=attention_scores, axis=-1)
941
+
942
+ # This is actually dropping out entire tokens to attend to, which might
943
+ # seem a bit unusual, but is taken from the original Transformer paper.
944
+ attention_probs = self.dropout(inputs=_attention_probs)
945
+
946
+ attention_output = tf.matmul(attention_probs, value_layer)
947
+ attention_output = tf.transpose(attention_output, perm=[0, 2, 1, 3])
948
+
949
+ # (batch_size, seq_len_q, embed_dim)
950
+ attention_output = tf.reshape(tensor=attention_output, shape=(batch_size, -1, self.embed_dim))
951
+
952
+ attention_output = self.out_proj(attention_output)
953
+ # In TFBert, attention weights are returned after dropout.
954
+ # However, in CLIP, they are returned before dropout.
955
+ outputs = (attention_output, _attention_probs) if output_attentions else (attention_output,)
956
+
957
+ return outputs
958
+
959
+ def build(self, input_shape=None):
960
+ if self.built:
961
+ return
962
+ self.built = True
963
+ if getattr(self, "q_proj", None) is not None:
964
+ with tf.name_scope(self.q_proj.name):
965
+ self.q_proj.build([None, None, self.embed_dim])
966
+ if getattr(self, "k_proj", None) is not None:
967
+ with tf.name_scope(self.k_proj.name):
968
+ self.k_proj.build([None, None, self.embed_dim])
969
+ if getattr(self, "v_proj", None) is not None:
970
+ with tf.name_scope(self.v_proj.name):
971
+ self.v_proj.build([None, None, self.embed_dim])
972
+ if getattr(self, "out_proj", None) is not None:
973
+ with tf.name_scope(self.out_proj.name):
974
+ self.out_proj.build([None, None, self.embed_dim])
975
+
976
+
977
+ # Copied from transformers.models.clip.modeling_tf_clip.TFCLIPEncoderLayer with CLIP->GroupViT
978
+ class TFGroupViTEncoderLayer(keras.layers.Layer):
979
+ def __init__(self, config: GroupViTConfig, **kwargs):
980
+ super().__init__(**kwargs)
981
+
982
+ self.embed_dim = config.hidden_size
983
+ self.self_attn = TFGroupViTAttention(config, name="self_attn")
984
+ self.layer_norm1 = keras.layers.LayerNormalization(epsilon=config.layer_norm_eps, name="layer_norm1")
985
+ self.mlp = TFGroupViTMLP(config, name="mlp")
986
+ self.layer_norm2 = keras.layers.LayerNormalization(epsilon=config.layer_norm_eps, name="layer_norm2")
987
+
988
+ def call(
989
+ self,
990
+ hidden_states: tf.Tensor,
991
+ attention_mask: tf.Tensor,
992
+ causal_attention_mask: tf.Tensor,
993
+ output_attentions: bool,
994
+ training: bool = False,
995
+ ) -> tuple[tf.Tensor]:
996
+ """
997
+ Args:
998
+ hidden_states (`tf.Tensor`): input to the layer of shape `(batch, seq_len, embed_dim)`
999
+ attention_mask (`tf.Tensor`): attention mask of size
1000
+ `(batch, 1, tgt_len, src_len)` where padding elements are indicated by very large negative values.
1001
+ causal_attention_mask (`tf.Tensor`): causal attention mask of size
1002
+ `(batch, 1, tgt_len, src_len)` where padding elements are indicated by very large negative values.
1003
+ output_attentions (`bool`):
1004
+ Whether or not to return the attentions tensors of all attention layers. See `outputs` under returned
1005
+ tensors for more detail.
1006
+ """
1007
+ residual = hidden_states
1008
+
1009
+ hidden_states = self.layer_norm1(inputs=hidden_states)
1010
+ attention_outputs = self.self_attn(
1011
+ hidden_states=hidden_states,
1012
+ attention_mask=attention_mask,
1013
+ causal_attention_mask=causal_attention_mask,
1014
+ output_attentions=output_attentions,
1015
+ training=training,
1016
+ )
1017
+ hidden_states = attention_outputs[0]
1018
+ hidden_states = residual + hidden_states
1019
+
1020
+ residual = hidden_states
1021
+ hidden_states = self.layer_norm2(inputs=hidden_states)
1022
+ hidden_states = self.mlp(hidden_states=hidden_states)
1023
+ hidden_states = residual + hidden_states
1024
+
1025
+ outputs = (hidden_states,) + attention_outputs[1:] # add attentions if we output them
1026
+
1027
+ return outputs
1028
+
1029
+ def build(self, input_shape=None):
1030
+ if self.built:
1031
+ return
1032
+ self.built = True
1033
+ if getattr(self, "self_attn", None) is not None:
1034
+ with tf.name_scope(self.self_attn.name):
1035
+ self.self_attn.build(None)
1036
+ if getattr(self, "layer_norm1", None) is not None:
1037
+ with tf.name_scope(self.layer_norm1.name):
1038
+ self.layer_norm1.build([None, None, self.embed_dim])
1039
+ if getattr(self, "mlp", None) is not None:
1040
+ with tf.name_scope(self.mlp.name):
1041
+ self.mlp.build(None)
1042
+ if getattr(self, "layer_norm2", None) is not None:
1043
+ with tf.name_scope(self.layer_norm2.name):
1044
+ self.layer_norm2.build([None, None, self.embed_dim])
1045
+
1046
+
1047
+ # Adapted from transformers.models.clip.modeling_tf_clip.TFGroupViTTextEncoder
1048
+ class TFGroupViTTextEncoder(keras.layers.Layer):
1049
+ def __init__(self, config: GroupViTTextConfig, **kwargs):
1050
+ super().__init__(**kwargs)
1051
+
1052
+ self.layers = [TFGroupViTEncoderLayer(config, name=f"layers_._{i}") for i in range(config.num_hidden_layers)]
1053
+
1054
+ def call(
1055
+ self,
1056
+ hidden_states,
1057
+ attention_mask: tf.Tensor,
1058
+ causal_attention_mask: tf.Tensor,
1059
+ output_attentions: bool,
1060
+ output_hidden_states: bool,
1061
+ return_dict: bool,
1062
+ training: bool = False,
1063
+ ) -> tuple | TFBaseModelOutput:
1064
+ encoder_states = () if output_hidden_states else None
1065
+ all_attentions = () if output_attentions else None
1066
+
1067
+ for idx, encoder_layer in enumerate(self.layers):
1068
+ if output_hidden_states:
1069
+ encoder_states = encoder_states + (hidden_states,)
1070
+
1071
+ layer_outputs = encoder_layer(
1072
+ hidden_states,
1073
+ attention_mask,
1074
+ causal_attention_mask,
1075
+ output_attentions=output_attentions,
1076
+ )
1077
+ hidden_states = layer_outputs[0]
1078
+
1079
+ if output_attentions:
1080
+ all_attentions = all_attentions + (layer_outputs[1],)
1081
+
1082
+ if output_hidden_states:
1083
+ encoder_states = encoder_states + (hidden_states,)
1084
+
1085
+ if not return_dict:
1086
+ return tuple(v for v in [hidden_states, encoder_states, all_attentions] if v is not None)
1087
+ return TFBaseModelOutput(
1088
+ last_hidden_state=hidden_states, hidden_states=encoder_states, attentions=all_attentions
1089
+ )
1090
+
1091
+ def build(self, input_shape=None):
1092
+ if self.built:
1093
+ return
1094
+ self.built = True
1095
+ if getattr(self, "layers", None) is not None:
1096
+ for layer in self.layers:
1097
+ with tf.name_scope(layer.name):
1098
+ layer.build(None)
1099
+
1100
+
1101
+ class TFGroupViTVisionEncoder(keras.layers.Layer):
1102
+ def __init__(self, config: GroupViTVisionConfig, **kwargs) -> None:
1103
+ super().__init__(**kwargs)
1104
+
1105
+ self.stages = [
1106
+ TFGroupViTStage(
1107
+ config=config,
1108
+ depth=config.depths[i],
1109
+ num_group_token=config.num_group_tokens[i],
1110
+ num_output_group=config.num_output_groups[i],
1111
+ num_prev_group_token=config.num_output_groups[i - 1] if i > 0 else 0,
1112
+ name=f"stages_._{i}",
1113
+ )
1114
+ for i in range(len(config.depths))
1115
+ ]
1116
+
1117
+ def call(
1118
+ self,
1119
+ hidden_states: tf.Tensor,
1120
+ output_hidden_states: bool,
1121
+ output_attentions: bool,
1122
+ return_dict: bool,
1123
+ training: bool = False,
1124
+ ) -> tuple | TFBaseModelOutput:
1125
+ all_hidden_states = () if output_hidden_states else None
1126
+ all_groupings = () if output_attentions else None
1127
+
1128
+ group_tokens = None
1129
+
1130
+ for stage in self.stages:
1131
+ if output_hidden_states:
1132
+ all_hidden_states = all_hidden_states + (hidden_states,)
1133
+
1134
+ layer_outputs = stage(hidden_states, group_tokens, output_attentions)
1135
+
1136
+ hidden_states = layer_outputs[0]
1137
+ group_tokens = layer_outputs[1]
1138
+
1139
+ if output_attentions and layer_outputs[2] is not None:
1140
+ all_groupings = all_groupings + (layer_outputs[2],)
1141
+
1142
+ if output_hidden_states:
1143
+ all_hidden_states = all_hidden_states + (hidden_states,)
1144
+
1145
+ if not return_dict:
1146
+ return tuple(v for v in [hidden_states, all_hidden_states, all_groupings] if v is not None)
1147
+ return TFBaseModelOutput(
1148
+ last_hidden_state=hidden_states, hidden_states=all_hidden_states, attentions=all_groupings
1149
+ )
1150
+
1151
+ def build(self, input_shape=None):
1152
+ if self.built:
1153
+ return
1154
+ self.built = True
1155
+ if getattr(self, "stages", None) is not None:
1156
+ for layer in self.stages:
1157
+ with tf.name_scope(layer.name):
1158
+ layer.build(None)
1159
+
1160
+
1161
+ # Copied from transformers.models.clip.modeling_tf_clip.TFCLIPTextTransformer with CLIPText->GroupViTText, CLIPEncoder->GroupViTTextEncoder
1162
+ class TFGroupViTTextTransformer(keras.layers.Layer):
1163
+ def __init__(self, config: GroupViTTextConfig, **kwargs):
1164
+ super().__init__(**kwargs)
1165
+
1166
+ self.embeddings = TFGroupViTTextEmbeddings(config, name="embeddings")
1167
+ self.encoder = TFGroupViTTextEncoder(config, name="encoder")
1168
+ self.final_layer_norm = keras.layers.LayerNormalization(epsilon=config.layer_norm_eps, name="final_layer_norm")
1169
+
1170
+ # For `pooled_output` computation
1171
+ self.eos_token_id = config.eos_token_id
1172
+ self.embed_dim = config.hidden_size
1173
+
1174
+ def call(
1175
+ self,
1176
+ input_ids: TFModelInputType,
1177
+ attention_mask: tf.Tensor,
1178
+ position_ids: tf.Tensor,
1179
+ output_attentions: bool,
1180
+ output_hidden_states: bool,
1181
+ return_dict: bool,
1182
+ training: bool = False,
1183
+ ) -> TFBaseModelOutputWithPooling | tuple[tf.Tensor]:
1184
+ input_shape = shape_list(input_ids)
1185
+
1186
+ embedding_output = self.embeddings(input_ids=input_ids, position_ids=position_ids)
1187
+
1188
+ batch_size, seq_length = input_shape
1189
+ # CLIP's text model uses causal mask, prepare it here.
1190
+ # https://github.com/openai/CLIP/blob/cfcffb90e69f37bf2ff1e988237a0fbe41f33c04/clip/model.py#L324
1191
+ causal_attention_mask = self._build_causal_attention_mask(batch_size, seq_length, dtype=embedding_output.dtype)
1192
+
1193
+ # check attention mask and invert
1194
+ # [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]
1195
+ attention_mask = _expand_mask(attention_mask)
1196
+
1197
+ encoder_outputs = self.encoder(
1198
+ hidden_states=embedding_output,
1199
+ attention_mask=attention_mask,
1200
+ causal_attention_mask=causal_attention_mask,
1201
+ output_attentions=output_attentions,
1202
+ output_hidden_states=output_hidden_states,
1203
+ return_dict=return_dict,
1204
+ training=training,
1205
+ )
1206
+
1207
+ sequence_output = encoder_outputs[0]
1208
+ sequence_output = self.final_layer_norm(inputs=sequence_output)
1209
+
1210
+ if self.eos_token_id == 2:
1211
+ # The `eos_token_id` was incorrect before PR #24773: Let's keep what have been done here.
1212
+ # A CLIP model with such `eos_token_id` in the config can't work correctly with extra new tokens added
1213
+ # ------------------------------------------------------------
1214
+ # text_embeds.shape = [batch_size, n_ctx, transformer.width]
1215
+ # take features from the eot embedding (eot_token is the highest number in each sequence)
1216
+ pooled_output = tf.gather_nd(
1217
+ params=sequence_output,
1218
+ indices=tf.stack(
1219
+ values=(tf.range(input_shape[0], dtype=tf.int64), tf.math.argmax(input_ids, axis=-1)), axis=1
1220
+ ),
1221
+ )
1222
+ else:
1223
+ # The config gets updated `eos_token_id` from PR #24773 (so the use of extra new tokens is possible)
1224
+ pooled_output = tf.gather_nd(
1225
+ params=sequence_output,
1226
+ indices=tf.stack(
1227
+ values=(
1228
+ tf.range(input_shape[0], dtype=tf.int64),
1229
+ tf.math.argmax(tf.cast(input_ids == self.eos_token_id, dtype=tf.int8), axis=-1),
1230
+ ),
1231
+ axis=1,
1232
+ ),
1233
+ )
1234
+
1235
+ if not return_dict:
1236
+ return (sequence_output, pooled_output) + encoder_outputs[1:]
1237
+
1238
+ return TFBaseModelOutputWithPooling(
1239
+ last_hidden_state=sequence_output,
1240
+ pooler_output=pooled_output,
1241
+ hidden_states=encoder_outputs.hidden_states,
1242
+ attentions=encoder_outputs.attentions,
1243
+ )
1244
+
1245
+ def _build_causal_attention_mask(self, batch_size, seq_length, dtype=tf.float32):
1246
+ # It is possible with an unspecified sequence length for seq_length to be
1247
+ # a runtime value, which is unsupported by tf.constant. Per the TensorFlow
1248
+ # docs, tf.fill can handle runtime dynamic shapes:
1249
+ # https://www.tensorflow.org/api_docs/python/tf/fill
1250
+ diag = tf.cast(tf.fill((seq_length,), 0.0), dtype)
1251
+
1252
+ # set an additive 2D attention mask with all places being masked
1253
+ to_mask = tf.cast(tf.fill((seq_length, seq_length), -10000.0), dtype)
1254
+
1255
+ # set diagonal & lower triangular parts to 0 (i.e. the places not to be masked)
1256
+ # TIP: think the 2D matrix as the space of (query_seq, key_seq)
1257
+ to_mask = tf.linalg.band_part(to_mask, 0, -1)
1258
+ # to_mask = tf.linalg.band_part(to_mask, -1, 0)
1259
+ to_mask = tf.linalg.set_diag(to_mask, diagonal=diag)
1260
+
1261
+ return tf.broadcast_to(input=to_mask, shape=(batch_size, 1, seq_length, seq_length))
1262
+
1263
+ def build(self, input_shape=None):
1264
+ if self.built:
1265
+ return
1266
+ self.built = True
1267
+ if getattr(self, "embeddings", None) is not None:
1268
+ with tf.name_scope(self.embeddings.name):
1269
+ self.embeddings.build(None)
1270
+ if getattr(self, "encoder", None) is not None:
1271
+ with tf.name_scope(self.encoder.name):
1272
+ self.encoder.build(None)
1273
+ if getattr(self, "final_layer_norm", None) is not None:
1274
+ with tf.name_scope(self.final_layer_norm.name):
1275
+ self.final_layer_norm.build([None, None, self.embed_dim])
1276
+
1277
+
1278
+ # Adapted from transformers.models.clip.modeling_tf_clip.TFCLIPVisionTransformer
1279
+ class TFGroupViTVisionTransformer(keras.layers.Layer):
1280
+ def __init__(self, config: GroupViTVisionConfig, **kwargs):
1281
+ super().__init__(**kwargs)
1282
+
1283
+ self.embeddings = TFGroupViTVisionEmbeddings(config, name="embeddings")
1284
+ self.encoder = TFGroupViTVisionEncoder(config, name="encoder")
1285
+ self.layernorm = keras.layers.LayerNormalization(epsilon=config.layer_norm_eps, name="layernorm")
1286
+ self.embed_dim = config.hidden_size
1287
+
1288
+ def call(
1289
+ self,
1290
+ pixel_values: TFModelInputType,
1291
+ output_attentions: bool,
1292
+ output_hidden_states: bool,
1293
+ return_dict: bool,
1294
+ training: bool = False,
1295
+ ) -> tuple | TFBaseModelOutputWithPooling:
1296
+ embedding_output = self.embeddings(pixel_values)
1297
+
1298
+ encoder_outputs = self.encoder(
1299
+ hidden_states=embedding_output,
1300
+ output_hidden_states=output_hidden_states,
1301
+ output_attentions=output_attentions,
1302
+ return_dict=return_dict,
1303
+ )
1304
+
1305
+ last_hidden_state = encoder_outputs[0]
1306
+
1307
+ # normalize the last hidden state
1308
+ last_hidden_state = self.layernorm(last_hidden_state)
1309
+ pooled_output = tf.math.reduce_mean(last_hidden_state, axis=1)
1310
+
1311
+ if not return_dict:
1312
+ return (last_hidden_state, pooled_output) + encoder_outputs[1:]
1313
+
1314
+ return TFBaseModelOutputWithPooling(
1315
+ last_hidden_state=last_hidden_state,
1316
+ pooler_output=pooled_output,
1317
+ hidden_states=encoder_outputs.hidden_states,
1318
+ attentions=encoder_outputs.attentions,
1319
+ )
1320
+
1321
+ def build(self, input_shape=None):
1322
+ if self.built:
1323
+ return
1324
+ self.built = True
1325
+ if getattr(self, "embeddings", None) is not None:
1326
+ with tf.name_scope(self.embeddings.name):
1327
+ self.embeddings.build(None)
1328
+ if getattr(self, "encoder", None) is not None:
1329
+ with tf.name_scope(self.encoder.name):
1330
+ self.encoder.build(None)
1331
+ if getattr(self, "layernorm", None) is not None:
1332
+ with tf.name_scope(self.layernorm.name):
1333
+ self.layernorm.build([None, None, self.embed_dim])
1334
+
1335
+
1336
+ @keras_serializable
1337
+ # Copied from transformers.models.clip.modeling_tf_clip.TFCLIPTextMainLayer with CLIP->GroupViT
1338
+ class TFGroupViTTextMainLayer(keras.layers.Layer):
1339
+ config_class = GroupViTTextConfig
1340
+
1341
+ def __init__(self, config: GroupViTTextConfig, **kwargs):
1342
+ super().__init__(**kwargs)
1343
+ self.config = config
1344
+ self.text_model = TFGroupViTTextTransformer(config, name="text_model")
1345
+
1346
+ def get_input_embeddings(self) -> keras.layers.Layer:
1347
+ return self.text_model.embeddings
1348
+
1349
+ def set_input_embeddings(self, value: tf.Variable):
1350
+ self.text_model.embeddings.weight = value
1351
+ self.text_model.embeddings.vocab_size = shape_list(value)[0]
1352
+
1353
+ @unpack_inputs
1354
+ def call(
1355
+ self,
1356
+ input_ids: TFModelInputType | None = None,
1357
+ attention_mask: np.ndarray | tf.Tensor | None = None,
1358
+ position_ids: np.ndarray | tf.Tensor | None = None,
1359
+ output_attentions: bool | None = None,
1360
+ output_hidden_states: bool | None = None,
1361
+ return_dict: bool | None = None,
1362
+ training: bool = False,
1363
+ ) -> TFBaseModelOutputWithPooling | tuple[tf.Tensor]:
1364
+ if input_ids is None:
1365
+ raise ValueError("You have to specify input_ids")
1366
+
1367
+ input_shape = shape_list(input_ids)
1368
+
1369
+ if attention_mask is None:
1370
+ attention_mask = tf.fill(dims=input_shape, value=1)
1371
+
1372
+ text_model_outputs = self.text_model(
1373
+ input_ids=input_ids,
1374
+ attention_mask=attention_mask,
1375
+ position_ids=position_ids,
1376
+ output_attentions=output_attentions,
1377
+ output_hidden_states=output_hidden_states,
1378
+ return_dict=return_dict,
1379
+ training=training,
1380
+ )
1381
+
1382
+ return text_model_outputs
1383
+
1384
+ def build(self, input_shape=None):
1385
+ if self.built:
1386
+ return
1387
+ self.built = True
1388
+ if getattr(self, "text_model", None) is not None:
1389
+ with tf.name_scope(self.text_model.name):
1390
+ self.text_model.build(None)
1391
+
1392
+
1393
+ @keras_serializable
1394
+ # Copied from transformers.models.clip.modeling_tf_clip.TFCLIPVisionMainLayer with CLIP->GroupViT
1395
+ class TFGroupViTVisionMainLayer(keras.layers.Layer):
1396
+ config_class = GroupViTVisionConfig
1397
+
1398
+ def __init__(self, config: GroupViTVisionConfig, **kwargs):
1399
+ super().__init__(**kwargs)
1400
+ self.config = config
1401
+ self.vision_model = TFGroupViTVisionTransformer(config, name="vision_model")
1402
+
1403
+ def get_input_embeddings(self) -> keras.layers.Layer:
1404
+ return self.vision_model.embeddings
1405
+
1406
+ @unpack_inputs
1407
+ def call(
1408
+ self,
1409
+ pixel_values: TFModelInputType | None = None,
1410
+ output_attentions: bool | None = None,
1411
+ output_hidden_states: bool | None = None,
1412
+ return_dict: bool | None = None,
1413
+ training: bool = False,
1414
+ ) -> TFBaseModelOutputWithPooling | tuple[tf.Tensor]:
1415
+ if pixel_values is None:
1416
+ raise ValueError("You have to specify pixel_values")
1417
+
1418
+ vision_model_outputs = self.vision_model(
1419
+ pixel_values=pixel_values,
1420
+ output_attentions=output_attentions,
1421
+ output_hidden_states=output_hidden_states,
1422
+ return_dict=return_dict,
1423
+ training=training,
1424
+ )
1425
+
1426
+ return vision_model_outputs
1427
+
1428
+ def build(self, input_shape=None):
1429
+ if self.built:
1430
+ return
1431
+ self.built = True
1432
+ if getattr(self, "vision_model", None) is not None:
1433
+ with tf.name_scope(self.vision_model.name):
1434
+ self.vision_model.build(None)
1435
+
1436
+
1437
+ @keras_serializable
1438
+ # Adapted from transformers.models.clip.modeling_tf_clip.TFCLIPMainLayer
1439
+ class TFGroupViTMainLayer(keras.layers.Layer):
1440
+ config_class = GroupViTConfig
1441
+
1442
+ def __init__(self, config: GroupViTConfig, **kwargs):
1443
+ super().__init__(**kwargs)
1444
+
1445
+ if not isinstance(config.text_config, GroupViTTextConfig):
1446
+ raise TypeError(
1447
+ "config.text_config is expected to be of type GroupViTTextConfig but is of type"
1448
+ f" {type(config.text_config)}."
1449
+ )
1450
+
1451
+ if not isinstance(config.vision_config, GroupViTVisionConfig):
1452
+ raise TypeError(
1453
+ "config.vision_config is expected to be of type GroupViTVisionConfig but is of type"
1454
+ f" {type(config.vision_config)}."
1455
+ )
1456
+
1457
+ self.config = config
1458
+
1459
+ text_config = config.text_config
1460
+ vision_config = config.vision_config
1461
+
1462
+ self.projection_dim = config.projection_dim
1463
+ self.projection_intermediate_dim = config.projection_intermediate_dim
1464
+ self.text_embed_dim = text_config.hidden_size
1465
+ self.vision_embed_dim = vision_config.hidden_size
1466
+
1467
+ self.text_model = TFGroupViTTextTransformer(text_config, name="text_model")
1468
+ self.vision_model = TFGroupViTVisionTransformer(vision_config, name="vision_model")
1469
+
1470
+ self.visual_projection = [
1471
+ keras.layers.Dense(self.projection_intermediate_dim, name="visual_projection.0"),
1472
+ keras.layers.BatchNormalization(name="visual_projection.1", momentum=0.9, epsilon=1e-5),
1473
+ keras.layers.ReLU(name="visual_projection.2"),
1474
+ keras.layers.Dense(self.projection_dim, name="visual_projection.3"),
1475
+ ]
1476
+ self.text_projection = [
1477
+ keras.layers.Dense(self.projection_intermediate_dim, name="text_projection.0"),
1478
+ keras.layers.BatchNormalization(name="text_projection.1", momentum=0.9, epsilon=1e-5),
1479
+ keras.layers.ReLU(name="text_projection.2"),
1480
+ keras.layers.Dense(self.projection_dim, name="text_projection.3"),
1481
+ ]
1482
+
1483
+ def build(self, input_shape=None):
1484
+ self.logit_scale = self.add_weight(
1485
+ shape=(1,),
1486
+ initializer=keras.initializers.Constant(self.config.logit_scale_init_value),
1487
+ trainable=True,
1488
+ name="logit_scale",
1489
+ )
1490
+
1491
+ if self.built:
1492
+ return
1493
+ self.built = True
1494
+ if getattr(self, "text_model", None) is not None:
1495
+ with tf.name_scope(self.text_model.name):
1496
+ self.text_model.build(None)
1497
+ if getattr(self, "vision_model", None) is not None:
1498
+ with tf.name_scope(self.vision_model.name):
1499
+ self.vision_model.build(None)
1500
+ if getattr(self, "visual_projection", None) is not None:
1501
+ with tf.name_scope(self.visual_projection[0].name):
1502
+ self.visual_projection[0].build([None, None, None, self.vision_embed_dim])
1503
+ with tf.name_scope(self.visual_projection[1].name):
1504
+ self.visual_projection[1].build((None, self.projection_intermediate_dim))
1505
+ with tf.name_scope(self.visual_projection[3].name):
1506
+ self.visual_projection[3].build([None, None, None, self.projection_intermediate_dim])
1507
+ if getattr(self, "text_projection", None) is not None:
1508
+ with tf.name_scope(self.text_projection[0].name):
1509
+ self.text_projection[0].build([None, None, None, self.text_embed_dim])
1510
+ with tf.name_scope(self.text_projection[1].name):
1511
+ self.text_projection[1].build((None, self.projection_intermediate_dim))
1512
+ with tf.name_scope(self.text_projection[3].name):
1513
+ self.text_projection[3].build([None, None, None, self.projection_intermediate_dim])
1514
+
1515
+ @unpack_inputs
1516
+ def get_text_features(
1517
+ self,
1518
+ input_ids: TFModelInputType | None = None,
1519
+ attention_mask: np.ndarray | tf.Tensor | None = None,
1520
+ position_ids: np.ndarray | tf.Tensor | None = None,
1521
+ output_attentions: bool | None = None,
1522
+ output_hidden_states: bool | None = None,
1523
+ return_dict: bool | None = None,
1524
+ training: bool = False,
1525
+ ) -> tf.Tensor:
1526
+ if input_ids is None:
1527
+ raise ValueError("You have to specify either input_ids")
1528
+
1529
+ input_shape = shape_list(input_ids)
1530
+
1531
+ if attention_mask is None:
1532
+ attention_mask = tf.fill(dims=input_shape, value=1)
1533
+
1534
+ text_outputs = self.text_model(
1535
+ input_ids=input_ids,
1536
+ attention_mask=attention_mask,
1537
+ position_ids=position_ids,
1538
+ output_attentions=output_attentions,
1539
+ output_hidden_states=output_hidden_states,
1540
+ return_dict=return_dict,
1541
+ training=training,
1542
+ )
1543
+
1544
+ pooled_output = text_outputs[1]
1545
+ for layer in self.text_projection:
1546
+ pooled_output = layer(pooled_output)
1547
+
1548
+ text_features = pooled_output
1549
+ return text_features
1550
+
1551
+ @unpack_inputs
1552
+ def get_image_features(
1553
+ self,
1554
+ pixel_values: TFModelInputType | None = None,
1555
+ output_attentions: bool | None = None,
1556
+ output_hidden_states: bool | None = None,
1557
+ return_dict: bool | None = None,
1558
+ training: bool = False,
1559
+ ) -> tf.Tensor:
1560
+ if pixel_values is None:
1561
+ raise ValueError("You have to specify pixel_values")
1562
+
1563
+ vision_outputs = self.vision_model(
1564
+ pixel_values=pixel_values,
1565
+ output_attentions=output_attentions,
1566
+ output_hidden_states=output_hidden_states,
1567
+ return_dict=return_dict,
1568
+ training=training,
1569
+ )
1570
+
1571
+ pooled_output = vision_outputs[1]
1572
+ for layer in self.visual_projection:
1573
+ pooled_output = layer(pooled_output)
1574
+
1575
+ image_features = pooled_output
1576
+ return image_features
1577
+
1578
+ @unpack_inputs
1579
+ def call(
1580
+ self,
1581
+ input_ids: TFModelInputType | None = None,
1582
+ pixel_values: TFModelInputType | None = None,
1583
+ attention_mask: np.ndarray | tf.Tensor | None = None,
1584
+ position_ids: np.ndarray | tf.Tensor | None = None,
1585
+ return_loss: bool | None = None,
1586
+ output_attentions: bool | None = None,
1587
+ output_hidden_states: bool | None = None,
1588
+ output_segmentation: bool | None = None,
1589
+ return_dict: bool | None = None,
1590
+ training: bool = False,
1591
+ ) -> TFGroupViTModelOutput | tuple[tf.Tensor]:
1592
+ if input_ids is None:
1593
+ raise ValueError("You have to specify either input_ids")
1594
+ if pixel_values is None:
1595
+ raise ValueError("You have to specify pixel_values")
1596
+
1597
+ input_shape = shape_list(input_ids)
1598
+
1599
+ if attention_mask is None:
1600
+ attention_mask = tf.fill(dims=input_shape, value=1)
1601
+ if output_segmentation:
1602
+ output_attentions = True
1603
+ vision_outputs = self.vision_model(
1604
+ pixel_values=pixel_values,
1605
+ output_attentions=output_attentions,
1606
+ output_hidden_states=output_hidden_states,
1607
+ return_dict=return_dict,
1608
+ training=training,
1609
+ )
1610
+
1611
+ text_outputs = self.text_model(
1612
+ input_ids=input_ids,
1613
+ attention_mask=attention_mask,
1614
+ position_ids=position_ids,
1615
+ output_attentions=output_attentions,
1616
+ output_hidden_states=output_hidden_states,
1617
+ return_dict=return_dict,
1618
+ training=training,
1619
+ )
1620
+
1621
+ image_embeds = vision_outputs[1]
1622
+ for layer in self.visual_projection:
1623
+ image_embeds = layer(image_embeds)
1624
+
1625
+ text_embeds = text_outputs[1]
1626
+ for layer in self.text_projection:
1627
+ text_embeds = layer(text_embeds)
1628
+
1629
+ # normalized features
1630
+ image_embeds = image_embeds / tf.norm(image_embeds, axis=-1, keepdims=True)
1631
+ text_embeds = text_embeds / tf.norm(text_embeds, axis=-1, keepdims=True)
1632
+
1633
+ # cosine similarity as logits
1634
+ logit_scale = tf.math.exp(self.logit_scale)
1635
+ logits_per_text = tf.matmul(text_embeds, image_embeds, transpose_b=True) * logit_scale
1636
+ logits_per_image = tf.transpose(logits_per_text)
1637
+
1638
+ seg_logits = None
1639
+ if output_segmentation:
1640
+ # grouped features
1641
+ # [batch_size_image, num_group, hidden_size]
1642
+ image_group_embeds = vision_outputs[0]
1643
+ # [batch_size_image*num_group, hidden_size]
1644
+ image_group_embeds = tf.reshape(image_group_embeds, shape=(-1, shape_list(image_group_embeds)[-1]))
1645
+ for layer in self.visual_projection:
1646
+ image_group_embeds = layer(image_group_embeds)
1647
+ if output_hidden_states:
1648
+ attentions = vision_outputs[3]
1649
+ else:
1650
+ attentions = vision_outputs[2]
1651
+ # [batch_size_image, num_group, height, width]
1652
+ grouping = get_grouping_from_attentions(attentions, pixel_values.shape[2:])
1653
+
1654
+ # normalized features
1655
+ image_group_embeds = image_group_embeds / tf.norm(
1656
+ tensor=image_group_embeds, ord="euclidean", axis=-1, keepdims=True
1657
+ )
1658
+ # [batch_size_image x num_group, batch_size_text]
1659
+ logits_per_image_group = tf.matmul(image_group_embeds, text_embeds, transpose_b=True) * logit_scale
1660
+ # [batch_size_image, batch_size_text, num_group]
1661
+ logits_per_image_group = tf.reshape(
1662
+ logits_per_image_group, shape=(image_embeds.shape[0], -1, text_embeds.shape[0])
1663
+ )
1664
+ logits_per_image_group = tf.transpose(logits_per_image_group, perm=(0, 2, 1))
1665
+
1666
+ # [batch_size_image, batch_size_text, height x width]
1667
+ flatten_grouping = tf.reshape(grouping, shape=(shape_list(grouping)[0], shape_list(grouping)[1], -1))
1668
+
1669
+ # [batch_size_image, batch_size_text, height, width]
1670
+ seg_logits = tf.matmul(logits_per_image_group, flatten_grouping) * logit_scale
1671
+ seg_logits = tf.reshape(
1672
+ seg_logits, shape=(seg_logits.shape[0], seg_logits.shape[1], grouping.shape[2], grouping.shape[3])
1673
+ )
1674
+
1675
+ loss = None
1676
+ if return_loss:
1677
+ loss = groupvit_loss(logits_per_text)[None, ...]
1678
+
1679
+ if not return_dict:
1680
+ if seg_logits is not None:
1681
+ output = (
1682
+ logits_per_image,
1683
+ logits_per_text,
1684
+ seg_logits,
1685
+ text_embeds,
1686
+ image_embeds,
1687
+ text_outputs,
1688
+ vision_outputs,
1689
+ )
1690
+ else:
1691
+ output = (logits_per_image, logits_per_text, text_embeds, image_embeds, text_outputs, vision_outputs)
1692
+ return ((loss,) + output) if loss is not None else output
1693
+
1694
+ return TFGroupViTModelOutput(
1695
+ loss=loss,
1696
+ logits_per_image=logits_per_image,
1697
+ logits_per_text=logits_per_text,
1698
+ segmentation_logits=seg_logits,
1699
+ text_embeds=text_embeds,
1700
+ image_embeds=image_embeds,
1701
+ text_model_output=text_outputs,
1702
+ vision_model_output=vision_outputs,
1703
+ )
1704
+
1705
+
1706
+ class TFGroupViTPreTrainedModel(TFPreTrainedModel):
1707
+ """
1708
+ An abstract class to handle weights initialization and a simple interface for downloading and loading pretrained
1709
+ models.
1710
+ """
1711
+
1712
+ config_class = GroupViTConfig
1713
+ base_model_prefix = "groupvit"
1714
+
1715
+
1716
+ GROUPVIT_START_DOCSTRING = r"""
1717
+ This model inherits from [`TFPreTrainedModel`]. Check the superclass documentation for the generic methods the
1718
+ library implements for all its model (such as downloading or saving, resizing the input embeddings, pruning heads
1719
+ etc.)
1720
+
1721
+ This model is also a [keras.Model](https://www.tensorflow.org/api_docs/python/tf/keras/Model) subclass. Use it
1722
+ as a regular TF 2.0 Keras Model and refer to the TF 2.0 documentation for all matter related to general usage and
1723
+ behavior.
1724
+
1725
+ <Tip>
1726
+
1727
+ TF 2.0 models accepts two formats as inputs:
1728
+
1729
+ - having all inputs as keyword arguments (like PyTorch models), or
1730
+ - having all inputs as a list, tuple or dict in the first positional arguments.
1731
+
1732
+ This second option is useful when using [`keras.Model.fit`] method which currently requires having all the
1733
+ tensors in the first argument of the model call function: `model(inputs)`.
1734
+
1735
+ If you choose this second option, there are three possibilities you can use to gather all the input Tensors in the
1736
+ first positional argument :
1737
+
1738
+ - a single Tensor with `input_ids` only and nothing else: `model(input_ids)`
1739
+ - a list of varying length with one or several input Tensors IN THE ORDER given in the docstring:
1740
+ `model([input_ids, attention_mask])` or `model([input_ids, attention_mask, token_type_ids])`
1741
+ - a dictionary with one or several input Tensors associated to the input names given in the docstring:
1742
+ `model({"input_ids": input_ids, "token_type_ids": token_type_ids})`
1743
+
1744
+ </Tip>
1745
+
1746
+ Args:
1747
+ config ([`GroupViTConfig`]): Model configuration class with all the parameters of the model.
1748
+ Initializing with a config file does not load the weights associated with the model, only the
1749
+ configuration. Check out the [`~PreTrainedModel.from_pretrained`] method to load the model weights.
1750
+ """
1751
+
1752
+ GROUPVIT_TEXT_INPUTS_DOCSTRING = r"""
1753
+ Args:
1754
+ input_ids (`np.ndarray`, `tf.Tensor`, `list[tf.Tensor]` ``dict[str, tf.Tensor]` or `dict[str, np.ndarray]` and each example must have the shape `({0})`):
1755
+ Indices of input sequence tokens in the vocabulary.
1756
+
1757
+ Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.__call__`] and
1758
+ [`PreTrainedTokenizer.encode`] for details.
1759
+
1760
+ [What are input IDs?](../glossary#input-ids)
1761
+ attention_mask (`np.ndarray` or `tf.Tensor` of shape `({0})`, *optional*):
1762
+ Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:
1763
+
1764
+ - 1 for tokens that are **not masked**,
1765
+ - 0 for tokens that are **masked**.
1766
+
1767
+ [What are attention masks?](../glossary#attention-mask)
1768
+ position_ids (`np.ndarray` or `tf.Tensor` of shape `({0})`, *optional*):
1769
+ Indices of positions of each input sequence tokens in the position embeddings. Selected in the range `[0,
1770
+ config.max_position_embeddings - 1]`.
1771
+
1772
+ [What are position IDs?](../glossary#position-ids)
1773
+ output_attentions (`bool`, *optional*):
1774
+ Whether or not to return the attentions tensors of all attention layers. See `attentions` under returned
1775
+ tensors for more detail. This argument can be used only in eager mode, in graph mode the value in the
1776
+ config will be used instead.
1777
+ output_hidden_states (`bool`, *optional*):
1778
+ Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors for
1779
+ more detail. This argument can be used only in eager mode, in graph mode the value in the config will be
1780
+ used instead.
1781
+ return_dict (`bool`, *optional*):
1782
+ Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple. This argument can be used in
1783
+ eager mode, in graph mode the value will always be set to True.
1784
+ training (`bool`, *optional*, defaults to `False``):
1785
+ Whether or not to use the model in training mode (some modules like dropout modules have different
1786
+ behaviors between training and evaluation).
1787
+ """
1788
+
1789
+ GROUPVIT_VISION_INPUTS_DOCSTRING = r"""
1790
+ Args:
1791
+ pixel_values (`np.ndarray`, `tf.Tensor`, `list[tf.Tensor]`, `dict[str, tf.Tensor]` or `dict[str, np.ndarray]` and each example must have the shape `(batch_size, num_channels, height, width)`):
1792
+ Pixel values. Pixel values can be obtained using [`AutoImageProcessor`]. See
1793
+ [`CLIPImageProcessor.__call__`] for details.
1794
+ output_attentions (`bool`, *optional*):
1795
+ Whether or not to return the attentions tensors of all attention layers. See `attentions` under returned
1796
+ tensors for more detail. This argument can be used only in eager mode, in graph mode the value in the
1797
+ config will be used instead.
1798
+ output_hidden_states (`bool`, *optional*):
1799
+ Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors for
1800
+ more detail. This argument can be used only in eager mode, in graph mode the value in the config will be
1801
+ used instead.
1802
+ return_dict (`bool`, *optional*):
1803
+ Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple. This argument can be used in
1804
+ eager mode, in graph mode the value will always be set to True.
1805
+ training (`bool`, *optional*, defaults to `False``):
1806
+ Whether or not to use the model in training mode (some modules like dropout modules have different
1807
+ behaviors between training and evaluation).
1808
+ """
1809
+
1810
+ GROUPVIT_INPUTS_DOCSTRING = r"""
1811
+ Args:
1812
+ input_ids (`np.ndarray`, `tf.Tensor`, `list[tf.Tensor]` ``dict[str, tf.Tensor]` or `dict[str, np.ndarray]` and each example must have the shape `({0})`):
1813
+ Indices of input sequence tokens in the vocabulary.
1814
+
1815
+ Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.__call__`] and
1816
+ [`PreTrainedTokenizer.encode`] for details.
1817
+
1818
+ [What are input IDs?](../glossary#input-ids)
1819
+ pixel_values (`np.ndarray`, `tf.Tensor`, `list[tf.Tensor]` `dict[str, tf.Tensor]` or `dict[str, np.ndarray]` and each example must have the shape `(batch_size, num_channels, height, width)`):
1820
+ Pixel values. Pixel values can be obtained using [`AutoImageProcessor`]. See
1821
+ [`CLIPImageProcessor.__call__`] for details.
1822
+ attention_mask (`np.ndarray` or `tf.Tensor` of shape `({0})`, *optional*):
1823
+ Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:
1824
+
1825
+ - 1 for tokens that are **not masked**,
1826
+ - 0 for tokens that are **masked**.
1827
+
1828
+ [What are attention masks?](../glossary#attention-mask)
1829
+ position_ids (`np.ndarray` or `tf.Tensor` of shape `({0})`, *optional*):
1830
+ Indices of positions of each input sequence tokens in the position embeddings. Selected in the range `[0,
1831
+ config.max_position_embeddings - 1]`.
1832
+
1833
+ [What are position IDs?](../glossary#position-ids)
1834
+ return_loss (`bool`, *optional*):
1835
+ Whether or not to return the contrastive loss.
1836
+ output_attentions (`bool`, *optional*):
1837
+ Whether or not to return the attentions tensors of all attention layers. See `attentions` under returned
1838
+ tensors for more detail. This argument can be used only in eager mode, in graph mode the value in the
1839
+ config will be used instead.
1840
+ output_hidden_states (`bool`, *optional*):
1841
+ Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors for
1842
+ more detail. This argument can be used only in eager mode, in graph mode the value in the config will be
1843
+ used instead.
1844
+ return_dict (`bool`, *optional*):
1845
+ Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple. This argument can be used in
1846
+ eager mode, in graph mode the value will always be set to True.
1847
+ training (`bool`, *optional*, defaults to `False``):
1848
+ Whether or not to use the model in training mode (some modules like dropout modules have different
1849
+ behaviors between training and evaluation).
1850
+ """
1851
+
1852
+
1853
+ class TFGroupViTTextModel(TFGroupViTPreTrainedModel):
1854
+ config_class = GroupViTTextConfig
1855
+ main_input_name = "input_ids"
1856
+
1857
+ def __init__(self, config: GroupViTTextConfig, *inputs, **kwargs):
1858
+ super().__init__(config, *inputs, **kwargs)
1859
+
1860
+ self.groupvit = TFGroupViTTextMainLayer(config, name="groupvit")
1861
+
1862
+ @unpack_inputs
1863
+ @add_start_docstrings_to_model_forward(GROUPVIT_TEXT_INPUTS_DOCSTRING.format("batch_size, sequence_length"))
1864
+ @replace_return_docstrings(output_type=TFBaseModelOutputWithPooling, config_class=GroupViTTextConfig)
1865
+ def call(
1866
+ self,
1867
+ input_ids: TFModelInputType | None = None,
1868
+ attention_mask: np.ndarray | tf.Tensor | None = None,
1869
+ position_ids: np.ndarray | tf.Tensor | None = None,
1870
+ output_attentions: bool | None = None,
1871
+ output_hidden_states: bool | None = None,
1872
+ return_dict: bool | None = None,
1873
+ training: bool = False,
1874
+ ) -> TFBaseModelOutputWithPooling | tuple[tf.Tensor]:
1875
+ r"""
1876
+ Returns:
1877
+
1878
+ Examples:
1879
+
1880
+ ```python
1881
+ >>> from transformers import CLIPTokenizer, TFGroupViTTextModel
1882
+
1883
+ >>> tokenizer = CLIPTokenizer.from_pretrained("nvidia/groupvit-gcc-yfcc")
1884
+ >>> model = TFGroupViTTextModel.from_pretrained("nvidia/groupvit-gcc-yfcc")
1885
+
1886
+ >>> inputs = tokenizer(["a photo of a cat", "a photo of a dog"], padding=True, return_tensors="tf")
1887
+
1888
+ >>> outputs = model(**inputs)
1889
+ >>> last_hidden_state = outputs.last_hidden_state
1890
+ >>> pooled_output = outputs.pooler_output # pooled (EOS token) states
1891
+ ```"""
1892
+
1893
+ outputs = self.groupvit(
1894
+ input_ids=input_ids,
1895
+ attention_mask=attention_mask,
1896
+ position_ids=position_ids,
1897
+ output_attentions=output_attentions,
1898
+ output_hidden_states=output_hidden_states,
1899
+ return_dict=return_dict,
1900
+ training=training,
1901
+ )
1902
+
1903
+ return outputs
1904
+
1905
+ def build(self, input_shape=None):
1906
+ if self.built:
1907
+ return
1908
+ self.built = True
1909
+ if getattr(self, "groupvit", None) is not None:
1910
+ with tf.name_scope(self.groupvit.name):
1911
+ self.groupvit.build(None)
1912
+
1913
+
1914
+ class TFGroupViTVisionModel(TFGroupViTPreTrainedModel):
1915
+ config_class = GroupViTVisionConfig
1916
+ main_input_name = "pixel_values"
1917
+
1918
+ def __init__(self, config: GroupViTVisionConfig, *inputs, **kwargs):
1919
+ super().__init__(config, *inputs, **kwargs)
1920
+
1921
+ self.groupvit = TFGroupViTVisionMainLayer(config, name="groupvit")
1922
+
1923
+ @unpack_inputs
1924
+ @add_start_docstrings_to_model_forward(GROUPVIT_VISION_INPUTS_DOCSTRING)
1925
+ @replace_return_docstrings(output_type=TFBaseModelOutputWithPooling, config_class=GroupViTVisionConfig)
1926
+ def call(
1927
+ self,
1928
+ pixel_values: TFModelInputType | None = None,
1929
+ output_attentions: bool | None = None,
1930
+ output_hidden_states: bool | None = None,
1931
+ return_dict: bool | None = None,
1932
+ training: bool = False,
1933
+ ) -> TFBaseModelOutputWithPooling | tuple[tf.Tensor]:
1934
+ r"""
1935
+ Returns:
1936
+
1937
+ Examples:
1938
+
1939
+ ```python
1940
+ >>> from PIL import Image
1941
+ >>> import requests
1942
+ >>> from transformers import AutoProcessor, TFGroupViTVisionModel
1943
+
1944
+ >>> processor = AutoProcessor.from_pretrained("nvidia/groupvit-gcc-yfcc")
1945
+ >>> model = TFGroupViTVisionModel.from_pretrained("nvidia/groupvit-gcc-yfcc")
1946
+
1947
+ >>> url = "http://images.cocodataset.org/val2017/000000039769.jpg"
1948
+ >>> image = Image.open(requests.get(url, stream=True).raw)
1949
+
1950
+ >>> inputs = processor(images=image, return_tensors="tf")
1951
+
1952
+ >>> outputs = model(**inputs)
1953
+ >>> last_hidden_state = outputs.last_hidden_state
1954
+ >>> pooled_output = outputs.pooler_output # pooled CLS states
1955
+ ```"""
1956
+
1957
+ outputs = self.groupvit(
1958
+ pixel_values=pixel_values,
1959
+ output_attentions=output_attentions,
1960
+ output_hidden_states=output_hidden_states,
1961
+ return_dict=return_dict,
1962
+ training=training,
1963
+ )
1964
+
1965
+ return outputs
1966
+
1967
+ def build(self, input_shape=None):
1968
+ if self.built:
1969
+ return
1970
+ self.built = True
1971
+ if getattr(self, "groupvit", None) is not None:
1972
+ with tf.name_scope(self.groupvit.name):
1973
+ self.groupvit.build(None)
1974
+
1975
+
1976
+ @add_start_docstrings(GROUPVIT_START_DOCSTRING)
1977
+ class TFGroupViTModel(TFGroupViTPreTrainedModel):
1978
+ config_class = GroupViTConfig
1979
+
1980
+ def __init__(self, config: GroupViTConfig, *inputs, **kwargs):
1981
+ super().__init__(config, *inputs, **kwargs)
1982
+
1983
+ self.groupvit = TFGroupViTMainLayer(config, name="groupvit")
1984
+
1985
+ @unpack_inputs
1986
+ @add_start_docstrings_to_model_forward(GROUPVIT_TEXT_INPUTS_DOCSTRING.format("batch_size, sequence_length"))
1987
+ def get_text_features(
1988
+ self,
1989
+ input_ids: TFModelInputType | None = None,
1990
+ attention_mask: np.ndarray | tf.Tensor | None = None,
1991
+ position_ids: np.ndarray | tf.Tensor | None = None,
1992
+ output_attentions: bool | None = None,
1993
+ output_hidden_states: bool | None = None,
1994
+ return_dict: bool | None = None,
1995
+ training: bool = False,
1996
+ ) -> tf.Tensor:
1997
+ r"""
1998
+ Returns:
1999
+ text_features (`tf.Tensor` of shape `(batch_size, output_dim`): The text embeddings obtained by applying
2000
+ the projection layer to the pooled output of [`TFGroupViTTextModel`].
2001
+
2002
+ Examples:
2003
+
2004
+ ```python
2005
+ >>> from transformers import CLIPTokenizer, TFGroupViTModel
2006
+
2007
+ >>> model = TFGroupViTModel.from_pretrained("nvidia/groupvit-gcc-yfcc")
2008
+ >>> tokenizer = CLIPTokenizer.from_pretrained("nvidia/groupvit-gcc-yfcc")
2009
+
2010
+ >>> inputs = tokenizer(["a photo of a cat", "a photo of a dog"], padding=True, return_tensors="tf")
2011
+ >>> text_features = model.get_text_features(**inputs)
2012
+ ```"""
2013
+
2014
+ text_features = self.groupvit.get_text_features(
2015
+ input_ids=input_ids,
2016
+ attention_mask=attention_mask,
2017
+ position_ids=position_ids,
2018
+ output_attentions=output_attentions,
2019
+ output_hidden_states=output_hidden_states,
2020
+ return_dict=return_dict,
2021
+ training=training,
2022
+ )
2023
+
2024
+ return text_features
2025
+
2026
+ @unpack_inputs
2027
+ @add_start_docstrings_to_model_forward(GROUPVIT_VISION_INPUTS_DOCSTRING)
2028
+ def get_image_features(
2029
+ self,
2030
+ pixel_values: TFModelInputType | None = None,
2031
+ output_attentions: bool | None = None,
2032
+ output_hidden_states: bool | None = None,
2033
+ return_dict: bool | None = None,
2034
+ training: bool = False,
2035
+ ) -> tf.Tensor:
2036
+ r"""
2037
+ Returns:
2038
+ image_features (`tf.Tensor` of shape `(batch_size, output_dim`): The image embeddings obtained by applying
2039
+ the projection layer to the pooled output of [`TFGroupViTVisionModel`].
2040
+
2041
+ Examples:
2042
+
2043
+ ```python
2044
+ >>> from PIL import Image
2045
+ >>> import requests
2046
+ >>> from transformers import AutoProcessor, TFGroupViTModel
2047
+
2048
+ >>> model = TFGroupViTModel.from_pretrained("nvidia/groupvit-gcc-yfcc")
2049
+ >>> processor = AutoProcessor.from_pretrained("nvidia/groupvit-gcc-yfcc")
2050
+
2051
+ >>> url = "http://images.cocodataset.org/val2017/000000039769.jpg"
2052
+ >>> image = Image.open(requests.get(url, stream=True).raw)
2053
+
2054
+ >>> inputs = processor(images=image, return_tensors="tf")
2055
+
2056
+ >>> image_features = model.get_image_features(**inputs)
2057
+ ```"""
2058
+
2059
+ image_features = self.groupvit.get_image_features(
2060
+ pixel_values=pixel_values,
2061
+ output_attentions=output_attentions,
2062
+ output_hidden_states=output_hidden_states,
2063
+ return_dict=return_dict,
2064
+ training=training,
2065
+ )
2066
+
2067
+ return image_features
2068
+
2069
+ @unpack_inputs
2070
+ @add_start_docstrings_to_model_forward(GROUPVIT_INPUTS_DOCSTRING.format("batch_size, sequence_length"))
2071
+ @replace_return_docstrings(output_type=TFGroupViTModelOutput, config_class=GroupViTConfig)
2072
+ def call(
2073
+ self,
2074
+ input_ids: TFModelInputType | None = None,
2075
+ pixel_values: TFModelInputType | None = None,
2076
+ attention_mask: np.ndarray | tf.Tensor | None = None,
2077
+ position_ids: np.ndarray | tf.Tensor | None = None,
2078
+ return_loss: bool | None = None,
2079
+ output_attentions: bool | None = None,
2080
+ output_hidden_states: bool | None = None,
2081
+ output_segmentation: bool | None = None,
2082
+ return_dict: bool | None = None,
2083
+ training: bool = False,
2084
+ ) -> TFGroupViTModelOutput | tuple[tf.Tensor]:
2085
+ r"""
2086
+ Returns:
2087
+
2088
+ Examples:
2089
+
2090
+ ```python
2091
+ >>> from PIL import Image
2092
+ >>> import requests
2093
+ >>> from transformers import AutoProcessor, TFGroupViTModel
2094
+ >>> import tensorflow as tf
2095
+
2096
+ >>> model = TFGroupViTModel.from_pretrained("nvidia/groupvit-gcc-yfcc")
2097
+ >>> processor = AutoProcessor.from_pretrained("nvidia/groupvit-gcc-yfcc")
2098
+
2099
+ >>> url = "http://images.cocodataset.org/val2017/000000039769.jpg"
2100
+ >>> image = Image.open(requests.get(url, stream=True).raw)
2101
+
2102
+ >>> inputs = processor(
2103
+ ... text=["a photo of a cat", "a photo of a dog"], images=image, return_tensors="tf", padding=True
2104
+ ... )
2105
+
2106
+ >>> outputs = model(**inputs)
2107
+ >>> logits_per_image = outputs.logits_per_image # this is the image-text similarity score
2108
+ >>> probs = tf.math.softmax(logits_per_image, axis=1) # we can take the softmax to get the label probabilities
2109
+ ```"""
2110
+
2111
+ outputs = self.groupvit(
2112
+ input_ids=input_ids,
2113
+ pixel_values=pixel_values,
2114
+ attention_mask=attention_mask,
2115
+ position_ids=position_ids,
2116
+ return_loss=return_loss,
2117
+ output_attentions=output_attentions,
2118
+ output_hidden_states=output_hidden_states,
2119
+ output_segmentation=output_segmentation,
2120
+ return_dict=return_dict,
2121
+ training=training,
2122
+ )
2123
+
2124
+ return outputs
2125
+
2126
+ def serving_output(self, output: TFGroupViTModelOutput) -> TFGroupViTModelOutput:
2127
+ # TODO: As is this currently fails with saved_model=True, because
2128
+ # TensorFlow cannot trace through nested dataclasses. Reference:
2129
+ # https://github.com/huggingface/transformers/pull/16886
2130
+ return output
2131
+
2132
+ def build(self, input_shape=None):
2133
+ if self.built:
2134
+ return
2135
+ self.built = True
2136
+ if getattr(self, "groupvit", None) is not None:
2137
+ with tf.name_scope(self.groupvit.name):
2138
+ self.groupvit.build(None)
2139
+
2140
+
2141
+ __all__ = ["TFGroupViTModel", "TFGroupViTPreTrainedModel", "TFGroupViTTextModel", "TFGroupViTVisionModel"]