Prompt48 commited on
Commit
1abf25c
·
verified ·
1 Parent(s): cc10368

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

Browse files
edit//Qwen3-TTS-test//.venv//Lib//site-packages//transformers//models//hiera//modeling_hiera.py ADDED
@@ -0,0 +1,1439 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # coding=utf-8
2
+ # Copyright 2024 Meta and The HuggingFace Inc. team. All rights reserved.
3
+ #
4
+ # Licensed under the Apache License, Version 2.0 (the "License");
5
+ # you may not use this file except in compliance with the License.
6
+ # You may obtain a copy of the License at
7
+ #
8
+ # http://www.apache.org/licenses/LICENSE-2.0
9
+ #
10
+ # Unless required by applicable law or agreed to in writing, software
11
+ # distributed under the License is distributed on an "AS IS" BASIS,
12
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ # See the License for the specific language governing permissions and
14
+ # limitations under the License.
15
+ """PyTorch Hiera model."""
16
+
17
+ import math
18
+ from dataclasses import dataclass
19
+ from typing import Optional, Union
20
+
21
+ import torch
22
+ from torch import nn
23
+
24
+ from ...activations import ACT2FN
25
+ from ...modeling_layers import GradientCheckpointingLayer
26
+ from ...modeling_outputs import (
27
+ BackboneOutput,
28
+ BaseModelOutput,
29
+ BaseModelOutputWithPooling,
30
+ ImageClassifierOutput,
31
+ ModelOutput,
32
+ )
33
+ from ...modeling_utils import PreTrainedModel
34
+ from ...utils import auto_docstring, logging, torch_int
35
+ from ...utils.backbone_utils import BackboneMixin
36
+ from .configuration_hiera import HieraConfig
37
+
38
+
39
+ logger = logging.get_logger(__name__)
40
+
41
+
42
+ @dataclass
43
+ @auto_docstring(
44
+ custom_intro="""
45
+ Hiera encoder's outputs, with potential hidden states and attentions.
46
+ """
47
+ )
48
+ class HieraEncoderOutput(ModelOutput):
49
+ r"""
50
+ reshaped_hidden_states (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`):
51
+ Tuple of `torch.FloatTensor` (one for the output of the embeddings + one for the output of each stage) of
52
+ shape `(batch_size, height, width, hidden_size)`. These are the reshaped and re-rolled hidden states of the model.
53
+
54
+ Hidden-states of the model at the output of each layer plus the initial embedding outputs reshaped to
55
+ include the spatial dimensions.
56
+ """
57
+
58
+ last_hidden_state: Optional[torch.FloatTensor] = None
59
+ hidden_states: Optional[tuple[torch.FloatTensor, ...]] = None
60
+ attentions: Optional[tuple[torch.FloatTensor, ...]] = None
61
+ reshaped_hidden_states: Optional[tuple[torch.FloatTensor, ...]] = None
62
+
63
+
64
+ @dataclass
65
+ @auto_docstring(
66
+ custom_intro="""
67
+ Hiera model's outputs that also contains a pooling of the last hidden states.
68
+ """
69
+ )
70
+ class HieraModelOutput(ModelOutput):
71
+ r"""
72
+ pooler_output (`torch.FloatTensor` of shape `(batch_size, hidden_size)`, *optional*, returned when `add_pooling_layer=True` is passed):
73
+ Average pooling of the last layer hidden-state.
74
+ bool_masked_pos (`torch.BoolTensor` of shape `(batch_size, sequence_length)`):
75
+ Tensor indicating which patches are masked (0) and which are not (1).
76
+ ids_restore (`torch.LongTensor` of shape `(batch_size, sequence_length)`):
77
+ Tensor containing the original index of the (shuffled) masked patches.
78
+ reshaped_hidden_states (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`):
79
+ Tuple of `torch.FloatTensor` (one for the output of the embeddings + one for the output of each stage) of
80
+ shape `(batch_size, height, width, hidden_size)`. These are the reshaped and re-rolled hidden states of the model.
81
+
82
+ Hidden-states of the model at the output of each layer plus the initial embedding outputs reshaped to
83
+ include the spatial dimensions.
84
+ """
85
+
86
+ last_hidden_state: Optional[torch.FloatTensor] = None
87
+ pooler_output: Optional[torch.FloatTensor] = None
88
+ bool_masked_pos: Optional[torch.BoolTensor] = None
89
+ ids_restore: Optional[torch.LongTensor] = None
90
+ hidden_states: Optional[tuple[torch.FloatTensor, ...]] = None
91
+ attentions: Optional[tuple[torch.FloatTensor, ...]] = None
92
+ reshaped_hidden_states: Optional[tuple[torch.FloatTensor, ...]] = None
93
+
94
+
95
+ @dataclass
96
+ @auto_docstring(
97
+ custom_intro="""
98
+ Hiera image classification outputs.
99
+ """
100
+ )
101
+ class HieraForImageClassificationOutput(ImageClassifierOutput):
102
+ r"""
103
+ loss (`torch.FloatTensor` of shape `(1,)`, `optional`):
104
+ Loss value for the training task.
105
+ logits (`torch.FloatTensor` of shape `(batch_size, num_labels)`):
106
+ Prediction scores of the classification head (logits of the output layer).
107
+ hidden_states (`tuple(torch.FloatTensor)`, `optional`):
108
+ Tuple of `torch.FloatTensor` (one for the output of the embeddings + one for the output of each stage) of
109
+ shape `(batch_size, sequence_length, hidden_size)`. These are the unrolled hidden states of the model.
110
+
111
+ Hidden-states of the model at the output of each layer plus the initial embedding outputs.
112
+ attentions (`tuple(torch.FloatTensor)`, `optional`):
113
+ Tuple of `torch.FloatTensor` (one for each stage) of shape `(batch_size, num_heads, sequence_length,
114
+ sequence_length)`.
115
+
116
+ Attentions weights after the attention softmax, used to compute the weighted average in the self-attention
117
+ heads.
118
+ reshaped_hidden_states (`tuple(torch.FloatTensor)`, `optional`):
119
+ Tuple of `torch.FloatTensor` (one for the output of the embeddings + one for the output of each stage) of
120
+ shape `(batch_size, height, width, hidden_size)`. These are the reshaped and re-rolled hidden states of the model.
121
+
122
+ Hidden-states of the model at the output of each layer plus the initial embedding outputs reshaped to
123
+ include the spatial dimensions.
124
+ """
125
+
126
+ loss: Optional[torch.FloatTensor] = None
127
+ logits: Optional[torch.FloatTensor] = None
128
+ hidden_states: Optional[tuple[torch.FloatTensor, ...]] = None
129
+ attentions: Optional[tuple[torch.FloatTensor, ...]] = None
130
+ reshaped_hidden_states: Optional[tuple[torch.FloatTensor, ...]] = None
131
+
132
+
133
+ @dataclass
134
+ @auto_docstring(
135
+ custom_intro="""
136
+ Class for HieraForPreTraining's outputs, with potential hidden states and attentions.
137
+ """
138
+ )
139
+ class HieraForPreTrainingOutput(ModelOutput):
140
+ r"""
141
+ loss (`torch.FloatTensor` of shape `(1,)`):
142
+ Pixel reconstruction loss.
143
+ logits (`torch.FloatTensor` of shape `(batch_size, sequence_length, patch_size ** 2 * num_channels)`):
144
+ Pixel reconstruction logits.
145
+ bool_masked_pos (`torch.BoolTensor` of shape `(batch_size, sequence_length)`):
146
+ Tensor indicating which patches are masked (0) and which are not (1).
147
+ ids_restore (`torch.LongTensor` of shape `(batch_size, sequence_length)`):
148
+ Tensor containing the original index of the (shuffled) masked patches.
149
+ reshaped_hidden_states (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`):
150
+ Tuple of `torch.FloatTensor` (one for the output of the embeddings + one for the output of each layer) of
151
+ shape `(batch_size, height, width, hidden_size)`. Hidden-states of the model at the output of each layer
152
+ plus the initial embedding outputs reshaped to include the spatial dimensions.
153
+ """
154
+
155
+ loss: Optional[torch.FloatTensor] = None
156
+ logits: Optional[torch.FloatTensor] = None
157
+ bool_masked_pos: Optional[torch.BoolTensor] = None
158
+ ids_restore: Optional[torch.LongTensor] = None
159
+ hidden_states: Optional[tuple[torch.FloatTensor]] = None
160
+ attentions: Optional[tuple[torch.FloatTensor]] = None
161
+ reshaped_hidden_states: Optional[tuple[torch.FloatTensor]] = None
162
+
163
+
164
+ class HieraPatchEmbeddings(nn.Module):
165
+ """
166
+ This class turns `pixel_values` of shape `(batch_size, num_channels, height, width)` into the initial
167
+ `hidden_states` (patch embeddings) of shape `(batch_size, seq_length, hidden_size)` to be consumed by a
168
+ Transformer.
169
+ """
170
+
171
+ def __init__(self, config, is_mae: bool = False):
172
+ super().__init__()
173
+
174
+ # Support any number of spatial dimensions
175
+ self.spatial_dims = len(config.patch_size)
176
+ if self.spatial_dims != 2:
177
+ raise ValueError(f"The number of dimensions of the input image should be 2, but got {self.spatial_dims}.")
178
+ self.num_channels = config.num_channels
179
+ self.image_size = config.image_size[-2:]
180
+ self.tokens_spatial_shape = [i // s for i, s in zip(config.image_size, config.patch_stride)]
181
+ self.mask_spatial_shape = [i // s for i, s in zip(self.tokens_spatial_shape, config.masked_unit_size)]
182
+ self.mask_ratio = config.mask_ratio
183
+ self.is_mae = is_mae
184
+ self.projection = nn.Conv2d(
185
+ self.num_channels,
186
+ config.embed_dim,
187
+ kernel_size=config.patch_size,
188
+ stride=config.patch_stride,
189
+ padding=config.patch_padding,
190
+ )
191
+
192
+ def masked_conv(
193
+ self, pixel_values: torch.FloatTensor, bool_masked_pos: Optional[torch.BoolTensor] = None
194
+ ) -> torch.Tensor:
195
+ """Zero-out the masked regions of the input before conv.
196
+ Prevents leakage of masked regions when using overlapping kernels.
197
+ """
198
+ if bool_masked_pos is None:
199
+ return self.projection(pixel_values)
200
+
201
+ target_size = pixel_values.shape[2:]
202
+ # Reshape bool_masked_pos to (batch_size, 1, mask_unit_height, mask_unit_width)
203
+ bool_masked_pos = bool_masked_pos.view(pixel_values.shape[0], 1, *self.mask_spatial_shape)
204
+
205
+ bool_masked_pos = nn.functional.interpolate(bool_masked_pos.float(), size=target_size)
206
+
207
+ return self.projection(pixel_values * bool_masked_pos)
208
+
209
+ def random_masking(
210
+ self, pixel_values: torch.FloatTensor, noise: Optional[torch.FloatTensor] = None
211
+ ) -> tuple[torch.BoolTensor, torch.LongTensor]:
212
+ """
213
+ Perform per-sample random masking by per-sample shuffling. Per-sample shuffling is done by argsort random
214
+ noise.
215
+
216
+ Args:
217
+ pixel_values (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)`)
218
+ noise (`torch.FloatTensor` of shape `(batch_size, num_mask_units)`, *optional*) which is
219
+ mainly used for testing purposes to control randomness and maintain the reproducibility
220
+ """
221
+ batch_size = pixel_values.shape[0]
222
+ # Tokens selected for masking at mask unit level
223
+ num_windows = math.prod(self.mask_spatial_shape)
224
+ len_keep = int(num_windows * (1 - self.mask_ratio))
225
+
226
+ if noise is None:
227
+ noise = torch.rand(batch_size, num_windows, device=pixel_values.device)
228
+
229
+ # Sort noise for each sample
230
+ ids_shuffle = torch.argsort(noise, dim=1)
231
+ # ascend: small is keep, large is remove
232
+ ids_restore = torch.argsort(ids_shuffle, dim=1).to(pixel_values.device)
233
+
234
+ # Generate the binary bool_masked_pos: 1 is *keep*, 0 is *remove*
235
+ # Note this is opposite to original MAE
236
+ bool_masked_pos = torch.zeros([batch_size, num_windows], device=pixel_values.device)
237
+ bool_masked_pos[:, :len_keep] = 1
238
+ # Unshuffle to get the binary bool_masked_pos
239
+ bool_masked_pos = torch.gather(bool_masked_pos, dim=1, index=ids_restore).bool()
240
+
241
+ return bool_masked_pos, ids_restore
242
+
243
+ def forward(
244
+ self,
245
+ pixel_values: torch.FloatTensor,
246
+ noise: Optional[torch.FloatTensor] = None,
247
+ ) -> tuple[torch.Tensor, Optional[torch.BoolTensor], Optional[torch.LongTensor]]:
248
+ (bool_masked_pos, ids_restore) = (
249
+ self.random_masking(pixel_values, noise=noise) if self.is_mae else (None, None)
250
+ )
251
+
252
+ embeddings = self.masked_conv(pixel_values, bool_masked_pos)
253
+ embeddings = embeddings.flatten(2).transpose(2, 1)
254
+
255
+ return embeddings, bool_masked_pos, ids_restore
256
+
257
+
258
+ class HieraEmbeddings(nn.Module):
259
+ """
260
+ Construct position and patch embeddings.
261
+ """
262
+
263
+ def __init__(self, config: HieraConfig, is_mae: bool = False) -> None:
264
+ super().__init__()
265
+ self.patch_stride = config.patch_stride
266
+ tokens_spatial_shape = [i // s for i, s in zip(config.image_size, config.patch_stride)]
267
+ self.mask_spatial_shape = [i // s for i, s in zip(tokens_spatial_shape, config.masked_unit_size)]
268
+ self.num_tokens = math.prod(tokens_spatial_shape)
269
+ self.is_mae = is_mae
270
+
271
+ self.patch_embeddings = HieraPatchEmbeddings(config, is_mae=is_mae)
272
+
273
+ self.position_embeddings = nn.Parameter(torch.zeros(1, self.num_tokens, config.embed_dim))
274
+
275
+ def interpolate_pos_encoding(
276
+ self, embeddings: torch.Tensor, pos_embeds: torch.Tensor, height: int, width: int
277
+ ) -> torch.Tensor:
278
+ """
279
+ This method allows to interpolate the pre-trained position encodings, to be able to use the model on higher resolution
280
+ images. This method is also adapted to support torch.jit tracing, no class embeddings, and different patch strides.
281
+
282
+ Adapted from:
283
+ - https://github.com/facebookresearch/dino/blob/de9ee3df6cf39fac952ab558447af1fa1365362a/vision_transformer.py#L174-L194, and
284
+ - https://github.com/facebookresearch/dinov2/blob/e1277af2ba9496fbadf7aec6eba56e8d882d1e35/dinov2/models/vision_transformer.py#L179-L211
285
+ """
286
+
287
+ num_patches = embeddings.shape[1]
288
+ num_positions = pos_embeds.shape[1]
289
+
290
+ # always interpolate when tracing to ensure the exported model works for dynamic input shapes
291
+ if not torch.jit.is_tracing() and num_patches == num_positions and height == width:
292
+ return pos_embeds
293
+
294
+ dim = embeddings.shape[-1]
295
+
296
+ new_height = height // self.patch_stride[0]
297
+ new_width = width // self.patch_stride[1]
298
+
299
+ sqrt_num_positions = torch_int(num_positions**0.5)
300
+ pos_embeds = pos_embeds.reshape(1, sqrt_num_positions, sqrt_num_positions, dim)
301
+ pos_embeds = pos_embeds.permute(0, 3, 1, 2)
302
+
303
+ pos_embeds = nn.functional.interpolate(
304
+ pos_embeds,
305
+ size=(new_height, new_width),
306
+ mode="bicubic",
307
+ align_corners=False,
308
+ )
309
+
310
+ pos_embeds = pos_embeds.permute(0, 2, 3, 1).view(1, -1, dim)
311
+ return pos_embeds
312
+
313
+ def get_position_embedding(
314
+ self, embeddings: torch.Tensor, height: int, width: int, interpolate_pos_encoding: bool
315
+ ) -> torch.FloatTensor:
316
+ return (
317
+ self.interpolate_pos_encoding(embeddings, self.position_embeddings, height, width)
318
+ if interpolate_pos_encoding
319
+ else self.position_embeddings
320
+ )
321
+
322
+ def forward(
323
+ self,
324
+ pixel_values: torch.FloatTensor,
325
+ noise: Optional[torch.FloatTensor] = None,
326
+ interpolate_pos_encoding: bool = False,
327
+ ) -> tuple[torch.Tensor, Optional[torch.BoolTensor], Optional[torch.LongTensor]]:
328
+ height, width = pixel_values.shape[-2:]
329
+ embeddings, bool_masked_pos, ids_restore = self.patch_embeddings(pixel_values, noise=noise)
330
+ embeddings = embeddings + self.get_position_embedding(embeddings, height, width, interpolate_pos_encoding)
331
+ return embeddings, bool_masked_pos, ids_restore
332
+
333
+
334
+ class HieraMaskUnitAttention(nn.Module):
335
+ """
336
+ Computes either Mask Unit or Global Attention. Also is able to perform query pooling.
337
+
338
+ Note: this assumes the tokens have already been flattened and unrolled into mask units.
339
+ """
340
+
341
+ def __init__(
342
+ self,
343
+ hidden_size: int,
344
+ hidden_size_output: int,
345
+ num_heads: int,
346
+ query_stride: int = 1,
347
+ window_size: int = 0,
348
+ use_mask_unit_attn: bool = False,
349
+ ) -> None:
350
+ super().__init__()
351
+ self.num_heads = num_heads
352
+ self.query_stride = query_stride
353
+ self.hidden_size_output = hidden_size_output
354
+
355
+ self.head_dim = hidden_size_output // num_heads
356
+ self.scale = (self.head_dim) ** -0.5
357
+
358
+ self.qkv = nn.Linear(hidden_size, 3 * hidden_size_output)
359
+ self.proj = nn.Linear(hidden_size_output, hidden_size_output)
360
+
361
+ self.window_size = window_size
362
+ self.use_mask_unit_attn = use_mask_unit_attn
363
+
364
+ def forward(
365
+ self,
366
+ hidden_states: torch.Tensor,
367
+ head_mask: Optional[torch.FloatTensor] = None,
368
+ output_attentions: bool = False,
369
+ ) -> tuple[torch.Tensor, Optional[torch.Tensor]]:
370
+ """Input should be of shape [batch, tokens, channels]."""
371
+ batch_size, seq_len, _ = hidden_states.shape
372
+
373
+ num_windows = 1
374
+ if self.use_mask_unit_attn:
375
+ num_windows = seq_len // (self.query_stride * self.window_size)
376
+
377
+ qkv = self.qkv(hidden_states)
378
+ qkv = qkv.reshape(batch_size, -1, num_windows, 3, self.num_heads, self.head_dim)
379
+ qkv = qkv.permute(3, 0, 4, 2, 1, 5)
380
+
381
+ query, key, value = qkv.unbind(0)
382
+
383
+ if self.query_stride > 1:
384
+ # Refer to unroll to see how this performs a maxpool-Nd
385
+ query = query.view(batch_size, self.num_heads, num_windows, self.query_stride, -1, self.head_dim)
386
+ query = query.max(dim=3).values
387
+
388
+ attn_weights = (query * self.scale) @ key.transpose(-1, -2)
389
+ attn_weights = attn_weights.softmax(dim=-1)
390
+
391
+ # Mask heads if we want to
392
+ if head_mask is not None:
393
+ attn_weights = attn_weights * head_mask
394
+
395
+ attn_output = attn_weights @ value
396
+ attn_output = attn_output.transpose(1, 3).reshape(batch_size, -1, self.hidden_size_output)
397
+ attn_output = self.proj(attn_output)
398
+
399
+ return (attn_output, attn_weights) if output_attentions else (attn_output, None)
400
+
401
+
402
+ # Copied from transformers.models.beit.modeling_beit.drop_path
403
+ def drop_path(input: torch.Tensor, drop_prob: float = 0.0, training: bool = False) -> torch.Tensor:
404
+ """
405
+ Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks).
406
+
407
+ Comment by Ross Wightman: This is the same as the DropConnect impl I created for EfficientNet, etc networks,
408
+ however, the original name is misleading as 'Drop Connect' is a different form of dropout in a separate paper...
409
+ See discussion: https://github.com/tensorflow/tpu/issues/494#issuecomment-532968956 ... I've opted for changing the
410
+ layer and argument names to 'drop path' rather than mix DropConnect as a layer name and use 'survival rate' as the
411
+ argument.
412
+ """
413
+ if drop_prob == 0.0 or not training:
414
+ return input
415
+ keep_prob = 1 - drop_prob
416
+ shape = (input.shape[0],) + (1,) * (input.ndim - 1) # work with diff dim tensors, not just 2D ConvNets
417
+ random_tensor = keep_prob + torch.rand(shape, dtype=input.dtype, device=input.device)
418
+ random_tensor.floor_() # binarize
419
+ output = input.div(keep_prob) * random_tensor
420
+ return output
421
+
422
+
423
+ # Copied from transformers.models.beit.modeling_beit.BeitDropPath with Beit->Hiera
424
+ class HieraDropPath(nn.Module):
425
+ """Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks)."""
426
+
427
+ def __init__(self, drop_prob: Optional[float] = None) -> None:
428
+ super().__init__()
429
+ self.drop_prob = drop_prob
430
+
431
+ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
432
+ return drop_path(hidden_states, self.drop_prob, self.training)
433
+
434
+ def extra_repr(self) -> str:
435
+ return f"p={self.drop_prob}"
436
+
437
+
438
+ class HieraMlp(nn.Module):
439
+ def __init__(self, config, dim: int) -> None:
440
+ super().__init__()
441
+ self.activation_fn = ACT2FN[config.hidden_act]
442
+ self.fc1 = nn.Linear(dim, int(dim * config.mlp_ratio))
443
+ self.fc2 = nn.Linear(int(dim * config.mlp_ratio), dim)
444
+
445
+ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
446
+ hidden_states = self.fc1(hidden_states)
447
+ hidden_states = self.activation_fn(hidden_states)
448
+ hidden_states = self.fc2(hidden_states)
449
+ return hidden_states
450
+
451
+
452
+ class HieraLayer(nn.Module):
453
+ def __init__(
454
+ self,
455
+ config,
456
+ hidden_size: int,
457
+ hidden_size_output: int,
458
+ num_heads: int,
459
+ drop_path: float = 0.0,
460
+ query_stride: int = 1,
461
+ window_size: int = 0,
462
+ use_mask_unit_attn: bool = False,
463
+ ) -> None:
464
+ super().__init__()
465
+
466
+ self.hidden_size = hidden_size
467
+ self.hidden_size_output = hidden_size_output
468
+ self.query_stride = query_stride
469
+
470
+ self.layernorm_before = nn.LayerNorm(hidden_size, eps=config.layer_norm_eps)
471
+ self.attn = HieraMaskUnitAttention(
472
+ hidden_size=hidden_size,
473
+ hidden_size_output=hidden_size_output,
474
+ num_heads=num_heads,
475
+ query_stride=query_stride,
476
+ window_size=window_size,
477
+ use_mask_unit_attn=use_mask_unit_attn,
478
+ )
479
+
480
+ self.layernorm_after = nn.LayerNorm(hidden_size_output, eps=config.layer_norm_eps)
481
+ self.mlp = HieraMlp(config, hidden_size_output)
482
+
483
+ self.drop_path = HieraDropPath(drop_path) if drop_path > 0 else nn.Identity()
484
+ if hidden_size != hidden_size_output:
485
+ self.proj = nn.Linear(hidden_size, hidden_size_output)
486
+
487
+ def forward(
488
+ self,
489
+ hidden_states: torch.Tensor,
490
+ head_mask: Optional[torch.FloatTensor] = None,
491
+ output_attentions: bool = False,
492
+ ) -> tuple[torch.Tensor, Optional[torch.Tensor]]:
493
+ batch_size, seq_len, _ = hidden_states.shape
494
+ # Attention + Q Pooling
495
+ hidden_states_norm = self.layernorm_before(hidden_states)
496
+ if self.hidden_size != self.hidden_size_output:
497
+ hidden_states = self.proj(hidden_states_norm)
498
+ # Refer to unroll to see how this performs a maxpool-Nd
499
+ hidden_states = (
500
+ hidden_states.view(batch_size, self.query_stride, -1, self.hidden_size_output).max(dim=1).values
501
+ )
502
+
503
+ (hidden_states_norm, attn_weights) = self.attn(
504
+ hidden_states_norm, head_mask, output_attentions=output_attentions
505
+ )
506
+ hidden_states = hidden_states + self.drop_path(hidden_states_norm)
507
+
508
+ residual = hidden_states
509
+ hidden_states = self.layernorm_after(hidden_states)
510
+ hidden_states = self.mlp(hidden_states)
511
+ hidden_states = residual + self.drop_path(hidden_states)
512
+
513
+ return (hidden_states, attn_weights)
514
+
515
+
516
+ class HieraStage(GradientCheckpointingLayer):
517
+ def __init__(
518
+ self,
519
+ config,
520
+ depth: int,
521
+ hidden_size: int,
522
+ hidden_size_output: int,
523
+ num_heads: int,
524
+ drop_path: list[float],
525
+ query_stride: list[int],
526
+ window_size: int,
527
+ use_mask_unit_attn: bool,
528
+ stage_num: Optional[int] = None,
529
+ ) -> None:
530
+ super().__init__()
531
+ # we need to know if the previous stage used masked attention
532
+ # mask unit or global attention.
533
+ # lag by 1 layer, so that global attention,
534
+ # applied post pooling on lower resolution
535
+ previous_stage_used_masked_attention = False
536
+ if stage_num is not None:
537
+ previous_stage_used_masked_attention = config.masked_unit_attention[stage_num - 1 if stage_num > 0 else 0]
538
+ self.layers = nn.ModuleList(
539
+ [
540
+ HieraLayer(
541
+ config=config,
542
+ hidden_size=hidden_size if i == 0 else hidden_size_output,
543
+ hidden_size_output=hidden_size_output,
544
+ num_heads=num_heads,
545
+ drop_path=drop_path[i],
546
+ query_stride=query_stride[i],
547
+ window_size=window_size,
548
+ use_mask_unit_attn=use_mask_unit_attn or (previous_stage_used_masked_attention and i == 0),
549
+ )
550
+ for i in range(depth)
551
+ ]
552
+ )
553
+
554
+ def forward(
555
+ self, hidden_states: torch.Tensor, head_mask: Optional[torch.FloatTensor], output_attentions: bool = False
556
+ ) -> tuple[torch.Tensor, Optional[torch.Tensor]]:
557
+ for i, layer_module in enumerate(self.layers):
558
+ layer_head_mask = head_mask[i] if head_mask is not None else None
559
+ (hidden_states, attn_weights) = layer_module(
560
+ hidden_states, layer_head_mask, output_attentions=output_attentions
561
+ )
562
+
563
+ return hidden_states, attn_weights
564
+
565
+
566
+ def undo_windowing(hidden_states: torch.Tensor, shape: list[int], mask_unit_shape: list[int]) -> torch.Tensor:
567
+ """
568
+ Restore spatial organization by undoing windowed organization of mask units.
569
+
570
+ Args:
571
+ hidden_states (`torch.Tensor`): The hidden states tensor of shape `[batch_size, num_mask_unit_height*num_mask_unit_width, hidden_size]`.
572
+ shape (`list[int]`): The original shape of the hidden states tensor before windowing.
573
+ mask_unit_shape (`list[int]`): The shape of the mask units used for windowing.
574
+
575
+ Returns:
576
+ torch.Tensor: The restored hidden states tensor of shape [batch_size, num_mask_unit_height*mask_unit_height, num_mask_unit_width*mask_unit_width, hidden_size].
577
+ """
578
+ batch_size, hidden_size = hidden_states.shape[0], hidden_states.shape[-1]
579
+ # From: [batch_size, num_mask_unit_height*num_mask_unit_width, hidden_size]
580
+ # To: [batch_size, num_mask_unit_height, num_mask_unit_width, mask_unit_height, mask_unit_width, hidden_size]
581
+ num_mask_units = [s // mu for s, mu in zip(shape, mask_unit_shape)]
582
+ hidden_states = hidden_states.view(batch_size, *num_mask_units, *mask_unit_shape, hidden_size)
583
+
584
+ # From: [batch_size, num_mask_unit_height, num_mask_unit_width, mask_unit_height, mask_unit_width, hidden_size]
585
+ # To: [batch_size, num_mask_unit_height*mask_unit_height, num_mask_unit_width*mask_unit_width, hidden_size]
586
+ hidden_states = hidden_states.permute(0, 1, 3, 2, 4, 5)
587
+ hidden_states = hidden_states.reshape(batch_size, *shape, hidden_size)
588
+
589
+ return hidden_states
590
+
591
+
592
+ class HieraEncoder(nn.Module):
593
+ def __init__(self, config: HieraConfig) -> None:
594
+ super().__init__()
595
+ total_depth = sum(config.depths)
596
+ # stochastic depth decay rule
597
+ dpr = [x.item() for x in torch.linspace(0, config.drop_path_rate, total_depth, device="cpu")]
598
+ # query strides rule
599
+ cumulative_depths = torch.tensor(config.depths, device="cpu").cumsum(0).tolist()
600
+ query_pool_layer = cumulative_depths[: config.num_query_pool]
601
+ query_strides = [math.prod(config.query_stride) if i in query_pool_layer else 1 for i in range(total_depth)]
602
+
603
+ # Transformer blocks
604
+ self.stages = nn.ModuleList()
605
+ hidden_size = config.embed_dim
606
+ stage_ends = [0] + cumulative_depths
607
+ masked_unit_area = math.prod(config.masked_unit_size)
608
+ query_stride_area = math.prod(config.query_stride)
609
+ for idx_stage, depth in enumerate(config.depths):
610
+ hidden_size_output = int(config.embed_dim * config.embed_dim_multiplier**idx_stage)
611
+
612
+ stage = HieraStage(
613
+ config=config,
614
+ depth=depth,
615
+ hidden_size=hidden_size,
616
+ hidden_size_output=hidden_size_output,
617
+ num_heads=config.num_heads[idx_stage],
618
+ drop_path=dpr[stage_ends[idx_stage] : stage_ends[idx_stage + 1]],
619
+ query_stride=query_strides[stage_ends[idx_stage] : stage_ends[idx_stage + 1]],
620
+ window_size=int(masked_unit_area * query_stride_area**-idx_stage),
621
+ use_mask_unit_attn=config.masked_unit_attention[idx_stage],
622
+ stage_num=idx_stage,
623
+ )
624
+
625
+ hidden_size = hidden_size_output
626
+ self.stages.append(stage)
627
+
628
+ # Setting reroll schedule
629
+ # The first stage has to reverse everything
630
+ # The next stage has to reverse all but the first unroll, etc.
631
+ stage_size = [i // s for i, s in zip(config.image_size, config.patch_stride)]
632
+ unroll_schedule = [config.query_stride] * len(config.depths[:-1])
633
+
634
+ self.schedule = {}
635
+ for idx_stage in range(len(config.depths)):
636
+ self.schedule[idx_stage] = unroll_schedule, stage_size
637
+ if idx_stage < config.num_query_pool:
638
+ stage_size = [i // s for i, s in zip(stage_size, config.query_stride)]
639
+ unroll_schedule = unroll_schedule[1:]
640
+
641
+ self.gradient_checkpointing = False
642
+
643
+ def reroll(
644
+ self, hidden_states: torch.Tensor, stage_idx: int, bool_masked_pos: Optional[torch.BoolTensor] = None
645
+ ) -> torch.Tensor:
646
+ """
647
+ Roll the given tensor back up to spatial order assuming it's from the given block.
648
+
649
+ If no bool_masked_pos is provided returns:
650
+ - [batch_size, height, width, hidden_size]
651
+ If a bool_masked_pos is provided returns:
652
+ - [batch_size, num_mask_units, mask_unit_height, mask_unit_width, hidden_size]
653
+ """
654
+ schedule, size = self.schedule[stage_idx]
655
+ batch_size, seq_len, hidden_size = hidden_states.shape
656
+
657
+ num_dim = len(size)
658
+ mask_unit_shape = [1] * num_dim
659
+
660
+ for strides in schedule:
661
+ # Extract the current patch from seq_len
662
+ hidden_states = hidden_states.view(
663
+ batch_size, *strides, seq_len // math.prod(strides), *mask_unit_shape, hidden_size
664
+ )
665
+
666
+ # Move that patch into the current MU
667
+ # Input: [batch_size, stride, stride, seq_len//(stride*stride), mask_unit_height, mask_unit_width, hidden_size]
668
+ # Output: [batch_size, seq_len//(stride*stride), stride, mask_unit_height, stride, mask_unit_width, hidden_size]
669
+ hidden_states = hidden_states.permute(0, 3, 1, 4, 2, 5, 6)
670
+
671
+ # Reshape to [batch_size, seq_len//(stride*stride), *mask_units, hidden_size]
672
+ for i in range(num_dim):
673
+ mask_unit_shape[i] *= strides[i]
674
+ hidden_states = hidden_states.reshape(batch_size, -1, *mask_unit_shape, hidden_size)
675
+ seq_len = hidden_states.shape[1]
676
+
677
+ # Current shape (e.g., 2d: [batch_size, #num_mask_units_height*#num_mask_units_width, mask_unit_height, mask_unit_width, hidden_size])
678
+ hidden_states = hidden_states.view(batch_size, seq_len, *mask_unit_shape, hidden_size)
679
+
680
+ # If masked, return [batch_size, num_mask_units, mask_unit_height, mask_unit_width, hidden_size]
681
+ if bool_masked_pos is not None:
682
+ return hidden_states
683
+
684
+ # If not masked, we can return [batch_size, height, width, hidden_size]
685
+ hidden_states = undo_windowing(hidden_states, size, mask_unit_shape)
686
+
687
+ return hidden_states
688
+
689
+ def forward(
690
+ self,
691
+ hidden_states: torch.Tensor,
692
+ bool_masked_pos: Optional[torch.BoolTensor] = None,
693
+ head_mask: Optional[torch.FloatTensor] = None,
694
+ output_attentions: bool = False,
695
+ output_hidden_states: bool = False,
696
+ return_dict: bool = True,
697
+ ) -> Union[tuple, BaseModelOutput]:
698
+ all_hidden_states = () if output_hidden_states else None
699
+ all_reshaped_hidden_states = () if output_hidden_states else None
700
+ all_self_attentions = () if output_attentions else None
701
+
702
+ if output_hidden_states:
703
+ all_hidden_states = all_hidden_states + (hidden_states,)
704
+ reshaped_hidden_states = self.reroll(hidden_states, stage_idx=0, bool_masked_pos=bool_masked_pos)
705
+ all_reshaped_hidden_states = all_reshaped_hidden_states + (reshaped_hidden_states,)
706
+
707
+ for i, stage_module in enumerate(self.stages):
708
+ layer_head_mask = head_mask[i] if head_mask is not None else None
709
+
710
+ layer_outputs = stage_module(hidden_states, layer_head_mask, output_attentions)
711
+
712
+ hidden_states = layer_outputs[0]
713
+
714
+ if output_attentions:
715
+ all_self_attentions = all_self_attentions + (layer_outputs[1],)
716
+
717
+ if output_hidden_states:
718
+ all_hidden_states = all_hidden_states + (hidden_states,)
719
+ reshaped_hidden_states = self.reroll(hidden_states, stage_idx=i, bool_masked_pos=bool_masked_pos)
720
+ all_reshaped_hidden_states = all_reshaped_hidden_states + (reshaped_hidden_states,)
721
+
722
+ if not return_dict:
723
+ return tuple(
724
+ v
725
+ for v in [hidden_states, all_hidden_states, all_self_attentions, all_reshaped_hidden_states]
726
+ if v is not None
727
+ )
728
+ return HieraEncoderOutput(
729
+ last_hidden_state=hidden_states,
730
+ hidden_states=all_hidden_states,
731
+ attentions=all_self_attentions,
732
+ reshaped_hidden_states=all_reshaped_hidden_states,
733
+ )
734
+
735
+
736
+ def unroll(
737
+ hidden_states: torch.Tensor, image_shape: tuple[int, int], patch_stride: tuple[int, int], schedule: list[list[int]]
738
+ ) -> torch.Tensor:
739
+ """
740
+ Reorders the tokens such that patches are contiguous in memory.
741
+ E.g., given [batch_size, (height, width), hidden_size] and stride of (stride, stride), this will re-order the tokens as
742
+ [batch_size, (stride, stride, height // stride, width // stride), hidden_size]
743
+
744
+ This allows operations like Max2d to be computed as x.view(batch_size, stride*stride, -1, hidden_size).max(dim=1).
745
+ Not only is this faster, but it also makes it easy to support inputs of arbitrary
746
+ dimensions in addition to patch-wise sparsity.
747
+
748
+ Performing this operation multiple times in sequence puts entire windows as contiguous
749
+ in memory. For instance, if you applied the stride (2, 2) 3 times, entire windows of
750
+ size 8x8 would be contiguous in memory, allowing operations like mask unit attention
751
+ computed easily and efficiently, while also allowing max to be applied sequentially.
752
+
753
+ Note: This means that intermediate values of the model are not in height x width order, so they
754
+ need to be re-rolled if you want to use the intermediate values as a height x width feature map.
755
+ The last block of the network is fine though, since by then the strides are all consumed.
756
+ """
757
+ batch_size, _, hidden_size = hidden_states.shape
758
+
759
+ size = [i // s for i, s in zip(image_shape, patch_stride)]
760
+
761
+ current_size = size
762
+ hidden_states = hidden_states.view(*([batch_size] + current_size + [hidden_size]))
763
+
764
+ for strides in schedule:
765
+ # Move patches with the given strides to the batch dimension
766
+
767
+ # Create a view of the tensor with the patch stride as separate dims
768
+ # For example in 2d: [batch_size, height // stride, stride, width // stride, stride, C]
769
+ current_size = [i // s for i, s in zip(current_size, strides)]
770
+ # initialize new_shape with [height // stride, stride, width // stride, stride]
771
+ new_shape = [item for pair in zip(current_size, strides) for item in pair]
772
+ # add batch_size and hidden_size to new_shape
773
+ new_shape = [batch_size] + new_shape + [hidden_size]
774
+ hidden_states = hidden_states.view(new_shape)
775
+
776
+ # Move the patch stride into the batch dimension
777
+ # For example in 2d: [batch_size, stride, stride, height // stride, width // stride, hidden_size]
778
+ num_dims = len(new_shape)
779
+ permute = [0] + list(range(2, num_dims - 1, 2)) + list(range(1, num_dims - 1, 2)) + [num_dims - 1]
780
+ hidden_states = hidden_states.permute(permute)
781
+
782
+ # Now finally flatten the relevant dims into the batch dimension
783
+ hidden_states = hidden_states.flatten(0, len(strides))
784
+ batch_size *= math.prod(strides)
785
+
786
+ hidden_states = hidden_states.reshape(-1, math.prod(size), hidden_size)
787
+ return hidden_states
788
+
789
+
790
+ @auto_docstring
791
+ class HieraPreTrainedModel(PreTrainedModel):
792
+ config: HieraConfig
793
+ base_model_prefix = "hiera"
794
+ main_input_name = "pixel_values"
795
+ supports_gradient_checkpointing = True
796
+
797
+ def _init_weights(self, module) -> None:
798
+ """Initialize the weights"""
799
+ std = self.config.initializer_range
800
+
801
+ if isinstance(module, HieraEmbeddings):
802
+ nn.init.trunc_normal_(module.position_embeddings, std=std)
803
+
804
+ elif isinstance(module, HieraDecoder):
805
+ nn.init.trunc_normal_(module.mask_token, std=std)
806
+ nn.init.trunc_normal_(module.decoder_position_embeddings, std=std)
807
+
808
+ elif isinstance(module, (nn.Linear, nn.Conv1d, nn.Conv2d)):
809
+ nn.init.trunc_normal_(module.weight, std=std)
810
+ if module.bias is not None:
811
+ nn.init.constant_(module.bias, std)
812
+
813
+ elif isinstance(module, nn.LayerNorm):
814
+ nn.init.constant_(module.bias, std)
815
+ nn.init.constant_(module.weight, self.config.layer_norm_init)
816
+
817
+
818
+ class HieraPooler(nn.Module):
819
+ def __init__(self, config: HieraConfig):
820
+ super().__init__()
821
+ num_features = int(config.embed_dim * config.embed_dim_multiplier ** (len(config.depths) - 1))
822
+ self.layernorm = nn.LayerNorm(num_features, eps=config.layer_norm_eps)
823
+ self.pooler = nn.AdaptiveAvgPool1d(1)
824
+
825
+ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
826
+ hidden_states = hidden_states.transpose(1, 2)
827
+ pooled_output = self.pooler(hidden_states)
828
+ pooled_output = torch.flatten(pooled_output, 1)
829
+ pooled_output = self.layernorm(pooled_output)
830
+ return pooled_output
831
+
832
+
833
+ @auto_docstring
834
+ class HieraModel(HieraPreTrainedModel):
835
+ def __init__(self, config: HieraConfig, add_pooling_layer: bool = True, is_mae: bool = False):
836
+ r"""
837
+ add_pooling_layer (`bool`, *optional*, defaults to `True`):
838
+ Whether or not to apply pooling layer.
839
+ is_mae (`bool`, *optional*, defaults to `False`):
840
+ Whether or not to run the model on MAE mode.
841
+ """
842
+ super().__init__(config)
843
+ self.num_features = int(config.embed_dim * config.embed_dim_multiplier ** (len(config.depths) - 1))
844
+
845
+ self.embeddings = HieraEmbeddings(config, is_mae=is_mae)
846
+ self.encoder = HieraEncoder(config)
847
+
848
+ self.unroll_schedule = [config.query_stride] * len(config.depths[:-1])
849
+
850
+ self.pooler = HieraPooler(config) if add_pooling_layer else None
851
+
852
+ # Initialize weights and apply final processing
853
+ self.post_init()
854
+
855
+ def get_input_embeddings(self) -> HieraPatchEmbeddings:
856
+ return self.embeddings.patch_embeddings
857
+
858
+ def _prune_heads(self, heads_to_prune: dict[int, list[int]]) -> None:
859
+ """
860
+ Prunes heads of the model. heads_to_prune: dict of {layer_num: list of heads to prune in this layer} See base
861
+ class PreTrainedModel
862
+ """
863
+ for layer, heads in heads_to_prune.items():
864
+ self.encoder.layer[layer].attention.prune_heads(heads)
865
+
866
+ @auto_docstring
867
+ def forward(
868
+ self,
869
+ pixel_values: Optional[torch.Tensor] = None,
870
+ noise: Optional[torch.FloatTensor] = None,
871
+ head_mask: Optional[torch.Tensor] = None,
872
+ output_attentions: Optional[bool] = None,
873
+ output_hidden_states: Optional[bool] = None,
874
+ interpolate_pos_encoding: Optional[bool] = None,
875
+ return_dict: Optional[bool] = None,
876
+ ) -> Union[tuple, BaseModelOutputWithPooling]:
877
+ r"""
878
+ noise (`torch.FloatTensor` of shape `(batch_size, num_mask_units)`, *optional*):
879
+ Mainly used for testing purposes to control randomness and maintain the reproducibility
880
+ """
881
+ output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
882
+ output_hidden_states = (
883
+ output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
884
+ )
885
+ return_dict = return_dict if return_dict is not None else self.config.use_return_dict
886
+
887
+ if pixel_values is None:
888
+ raise ValueError("You have to specify pixel_values")
889
+
890
+ # Prepare head mask if needed
891
+ # 1.0 in head_mask indicate we keep the head
892
+ # attention_probs has shape bsz x n_heads x N x N
893
+ # input head_mask has shape [num_heads] or [num_hidden_layers x num_heads]
894
+ # and head_mask is converted to shape [num_hidden_layers x batch x num_heads x seq_length x seq_length]
895
+ head_mask = self.get_head_mask(head_mask, len(self.config.depths))
896
+
897
+ embedding_output, bool_masked_pos, ids_restore = self.embeddings(
898
+ pixel_values, interpolate_pos_encoding=interpolate_pos_encoding, noise=noise
899
+ )
900
+
901
+ image_shape = (pixel_values.shape[-2], pixel_values.shape[-1])
902
+ hidden_states = unroll(
903
+ embedding_output,
904
+ image_shape=image_shape,
905
+ patch_stride=self.config.patch_stride,
906
+ schedule=self.unroll_schedule,
907
+ )
908
+
909
+ # Discard masked tokens if bool_masked_pos is provided
910
+ if bool_masked_pos is not None:
911
+ mask_unit_area = math.prod(self.config.masked_unit_size)
912
+ batch_size, _, hidden_size = hidden_states.shape
913
+ positions = bool_masked_pos.unsqueeze(-1).tile(1, mask_unit_area, hidden_size)
914
+ hidden_states = hidden_states[positions]
915
+ hidden_states = hidden_states.view(batch_size, -1, hidden_size)
916
+
917
+ encoder_outputs = self.encoder(
918
+ hidden_states,
919
+ bool_masked_pos=bool_masked_pos,
920
+ head_mask=head_mask,
921
+ output_attentions=output_attentions,
922
+ output_hidden_states=output_hidden_states,
923
+ return_dict=return_dict,
924
+ )
925
+ sequence_output = encoder_outputs[0]
926
+ pooled_output = None
927
+ if self.pooler is not None:
928
+ pooled_output = self.pooler(sequence_output)
929
+
930
+ if not return_dict:
931
+ head_outputs = (sequence_output, pooled_output) if pooled_output is not None else (sequence_output,)
932
+ head_outputs = (
933
+ head_outputs + (bool_masked_pos, ids_restore) if bool_masked_pos is not None else head_outputs
934
+ )
935
+ return head_outputs + encoder_outputs[1:]
936
+
937
+ return HieraModelOutput(
938
+ last_hidden_state=sequence_output,
939
+ pooler_output=pooled_output,
940
+ bool_masked_pos=bool_masked_pos,
941
+ ids_restore=ids_restore,
942
+ hidden_states=encoder_outputs.hidden_states,
943
+ attentions=encoder_outputs.attentions,
944
+ reshaped_hidden_states=encoder_outputs.reshaped_hidden_states,
945
+ )
946
+
947
+
948
+ class HieraDecoder(nn.Module):
949
+ def __init__(self, config: HieraConfig):
950
+ super().__init__()
951
+ num_features = int(config.embed_dim * config.embed_dim_multiplier ** (len(config.depths) - 1))
952
+ tokens_spatial_shape = [i // s for i, s in zip(config.image_size, config.patch_stride)]
953
+ self.tokens_spatial_shape_final = [
954
+ i // s ** (config.num_query_pool) for i, s in zip(tokens_spatial_shape, config.query_stride)
955
+ ]
956
+ self.mask_unit_spatial_shape_final = [
957
+ i // s ** (config.num_query_pool) for i, s in zip(config.masked_unit_size, config.query_stride)
958
+ ]
959
+
960
+ self.decoder_embeddings = nn.Linear(num_features, config.decoder_hidden_size)
961
+
962
+ self.mask_token = nn.Parameter(torch.zeros(1, 1, config.decoder_hidden_size))
963
+
964
+ self.decoder_position_embeddings = nn.Parameter(
965
+ torch.zeros(1, math.prod(self.tokens_spatial_shape_final), config.decoder_hidden_size)
966
+ )
967
+
968
+ self.decoder_block = HieraStage(
969
+ config=config,
970
+ hidden_size=config.decoder_hidden_size,
971
+ hidden_size_output=config.decoder_hidden_size,
972
+ num_heads=config.decoder_num_heads,
973
+ depth=config.decoder_depth,
974
+ use_mask_unit_attn=False,
975
+ drop_path=[0.0] * config.decoder_depth,
976
+ query_stride=[1] * config.decoder_depth,
977
+ window_size=0,
978
+ )
979
+
980
+ self.decoder_norm = nn.LayerNorm(config.decoder_hidden_size, eps=config.layer_norm_eps)
981
+
982
+ # patch stride of prediction
983
+ self.pred_stride = config.patch_stride[-1] * (config.query_stride[-1] ** config.num_query_pool)
984
+ pred_dim = (self.pred_stride ** len(config.query_stride)) * config.num_channels
985
+
986
+ self.decoder_pred = nn.Linear(config.decoder_hidden_size, pred_dim)
987
+
988
+ def forward(
989
+ self,
990
+ encoder_hidden_states: torch.Tensor,
991
+ bool_masked_pos: torch.BoolTensor,
992
+ head_mask: Optional[torch.Tensor] = None,
993
+ output_attentions: bool = False,
994
+ ) -> tuple[torch.Tensor, torch.BoolTensor]:
995
+ # Embed tokens
996
+ hidden_states = self.decoder_embeddings(encoder_hidden_states)
997
+
998
+ # Combine visible and bool_masked_pos tokens
999
+
1000
+ # hidden_states : [batch_size, num_mask_units_visible, *mask_unit_spatial_shape_final, decoder_hidden_size]
1001
+ # bool_masked_pos: [batch_size, num_mask_units]
1002
+ mask_unit_height, mask_unit_width, decoder_hidden_size = hidden_states.shape[2:]
1003
+ batch_size, num_mask_units = bool_masked_pos.shape
1004
+
1005
+ decoder_hidden_states = torch.zeros(
1006
+ batch_size,
1007
+ num_mask_units,
1008
+ mask_unit_height,
1009
+ mask_unit_width,
1010
+ decoder_hidden_size,
1011
+ device=hidden_states.device,
1012
+ dtype=hidden_states.dtype,
1013
+ )
1014
+ mask_tokens = self.mask_token.view(1, 1, 1, 1, -1)
1015
+ bool_masked_pos = bool_masked_pos.reshape(batch_size, num_mask_units, 1, 1, 1)
1016
+ bool_masked_pos = bool_masked_pos.expand(-1, -1, mask_unit_height, mask_unit_width, decoder_hidden_size)
1017
+ decoder_hidden_states[bool_masked_pos] = hidden_states.flatten()
1018
+ decoder_hidden_states = (
1019
+ 1 - bool_masked_pos.float()
1020
+ ) * mask_tokens + bool_masked_pos.float() * decoder_hidden_states
1021
+
1022
+ # Get back spatial order
1023
+ hidden_states = undo_windowing(
1024
+ decoder_hidden_states,
1025
+ self.tokens_spatial_shape_final,
1026
+ self.mask_unit_spatial_shape_final,
1027
+ )
1028
+ bool_masked_pos = undo_windowing(
1029
+ bool_masked_pos[..., 0:1],
1030
+ self.tokens_spatial_shape_final,
1031
+ self.mask_unit_spatial_shape_final,
1032
+ )
1033
+
1034
+ # Flatten
1035
+ hidden_states = hidden_states.reshape(hidden_states.shape[0], -1, hidden_states.shape[-1])
1036
+ bool_masked_pos = bool_masked_pos.view(hidden_states.shape[0], -1)
1037
+
1038
+ # Add pos embed
1039
+ hidden_states = hidden_states + self.decoder_position_embeddings
1040
+
1041
+ # Apply decoder blocks
1042
+ hidden_states, attn_weights = self.decoder_block(
1043
+ hidden_states, head_mask=head_mask, output_attentions=output_attentions
1044
+ )
1045
+ hidden_states = self.decoder_norm(hidden_states)
1046
+
1047
+ # Predictor projection
1048
+ hidden_states = self.decoder_pred(hidden_states)
1049
+
1050
+ return hidden_states, bool_masked_pos
1051
+
1052
+
1053
+ class HieraMultiScaleHead(nn.Module):
1054
+ def __init__(self, config: HieraConfig):
1055
+ super().__init__()
1056
+ self.mask_unit_spatial_shape_final = [
1057
+ i // s ** (config.num_query_pool) for i, s in zip(config.masked_unit_size, config.query_stride)
1058
+ ]
1059
+ self.stage_dimensions = [
1060
+ int(config.embed_dim * config.embed_dim_multiplier**i) for i in range(len(config.depths))
1061
+ ]
1062
+ current_masked_unit_size = config.masked_unit_size
1063
+ self.multi_scale_fusion_heads = nn.ModuleList()
1064
+
1065
+ for idx in range(config.num_query_pool):
1066
+ kernel = [i // s for i, s in zip(current_masked_unit_size, self.mask_unit_spatial_shape_final)]
1067
+ current_masked_unit_size = [i // s for i, s in zip(current_masked_unit_size, config.query_stride)]
1068
+ self.multi_scale_fusion_heads.append(
1069
+ nn.Conv2d(
1070
+ self.stage_dimensions[idx],
1071
+ self.stage_dimensions[-1],
1072
+ kernel_size=kernel,
1073
+ stride=kernel,
1074
+ )
1075
+ )
1076
+ self.multi_scale_fusion_heads.append(nn.Identity())
1077
+
1078
+ def apply_fusion_head(self, head: nn.Module, hidden_states: torch.Tensor) -> torch.Tensor:
1079
+ if isinstance(head, nn.Identity):
1080
+ return hidden_states
1081
+
1082
+ # Doing explicit to avoid problems with torch.fx
1083
+ batch_size, num_mask_units, mask_unit_height, mask_unit_width, hidden_size = hidden_states.shape
1084
+ # From: [batch_size, num_mask_units, mask_unit_height, mask_unit_width, hidden_size]
1085
+ # To: head([batch_size * num_mask_units, hidden_size, mask_unit_height, mask_unit_width])
1086
+ hidden_states = hidden_states.reshape(
1087
+ batch_size * num_mask_units, mask_unit_height, mask_unit_width, hidden_size
1088
+ )
1089
+ hidden_states = hidden_states.permute(0, 3, 1, 2)
1090
+ hidden_states = head(hidden_states)
1091
+
1092
+ # Restore original layout
1093
+ hidden_states = hidden_states.permute(0, 2, 3, 1)
1094
+ mask_unit_height_final, mask_unit_width_final, hidden_size = hidden_states.shape[1:]
1095
+ hidden_states = hidden_states.reshape(
1096
+ batch_size, num_mask_units, mask_unit_height_final, mask_unit_width_final, hidden_size
1097
+ )
1098
+
1099
+ return hidden_states
1100
+
1101
+ def forward(self, feature_maps: list[torch.Tensor]) -> torch.Tensor:
1102
+ # Multi-scale fusion
1103
+ hidden_states = 0.0
1104
+ for head, feature_map in zip(self.multi_scale_fusion_heads, feature_maps):
1105
+ hidden_states = hidden_states + self.apply_fusion_head(head, feature_map)
1106
+
1107
+ return hidden_states
1108
+
1109
+
1110
+ @auto_docstring(
1111
+ custom_intro="""
1112
+ The Hiera Model transformer with the decoder on top for self-supervised pre-training.
1113
+
1114
+ <Tip>
1115
+
1116
+ Note that we provide a script to pre-train this model on custom data in our [examples
1117
+ directory](https://github.com/huggingface/transformers/tree/main/examples/pytorch/image-pretraining).
1118
+
1119
+ </Tip>
1120
+ """
1121
+ )
1122
+ class HieraForPreTraining(HieraPreTrainedModel):
1123
+ def __init__(self, config: HieraConfig) -> None:
1124
+ super().__init__(config)
1125
+ # Encoder
1126
+ self.hiera = HieraModel(config, add_pooling_layer=False, is_mae=True)
1127
+ self.encoder_norm = nn.LayerNorm(self.hiera.num_features, eps=config.layer_norm_eps)
1128
+ # Multi-scale fusion heads
1129
+ self.multiscale_fusion = HieraMultiScaleHead(config)
1130
+ # Decoder
1131
+ self.decoder = HieraDecoder(config)
1132
+ self.pred_stride = self.decoder.pred_stride
1133
+
1134
+ # Initialize weights and apply final processing
1135
+ self.post_init()
1136
+
1137
+ def get_pixel_label_2d(self, pixel_values: torch.Tensor, bool_masked_pos: torch.BoolTensor) -> torch.Tensor:
1138
+ # bool_masked_pos (boolean tensor): True means *masked*
1139
+ pixel_values = pixel_values.permute(0, 2, 3, 1)
1140
+
1141
+ size = self.pred_stride
1142
+ label = pixel_values.unfold(1, size, size).unfold(2, size, size)
1143
+ label = label.flatten(1, 2).flatten(2)
1144
+ label = label[bool_masked_pos]
1145
+ if self.config.normalize_pixel_loss:
1146
+ mean = label.mean(dim=-1, keepdim=True)
1147
+ var = label.var(dim=-1, keepdim=True)
1148
+ label = (label - mean) / (var + 1.0e-6) ** 0.5
1149
+
1150
+ return label
1151
+
1152
+ def forward_loss(self, pixel_values: torch.Tensor, logits: torch.Tensor, bool_masked_pos: torch.BoolTensor):
1153
+ # We invert the bool_masked_pos such that 1.0 is *masked*
1154
+ bool_masked_pos = ~bool_masked_pos
1155
+ label = self.get_pixel_label_2d(pixel_values, bool_masked_pos)
1156
+
1157
+ logits = logits[bool_masked_pos]
1158
+ loss = (logits - label) ** 2
1159
+ loss = loss.mean()
1160
+
1161
+ return loss
1162
+
1163
+ @auto_docstring
1164
+ def forward(
1165
+ self,
1166
+ pixel_values: Optional[torch.Tensor] = None,
1167
+ noise: Optional[torch.FloatTensor] = None,
1168
+ head_mask: Optional[torch.Tensor] = None,
1169
+ output_attentions: Optional[bool] = None,
1170
+ output_hidden_states: Optional[bool] = None,
1171
+ interpolate_pos_encoding: Optional[bool] = None,
1172
+ return_dict: Optional[bool] = None,
1173
+ ) -> Union[tuple, HieraForPreTrainingOutput]:
1174
+ r"""
1175
+ noise (`torch.FloatTensor` of shape `(batch_size, num_mask_units)`, *optional*):
1176
+ Mainly used for testing purposes to control randomness and maintain the reproducibility
1177
+
1178
+ Examples:
1179
+ ```python
1180
+ >>> from transformers import AutoImageProcessor, HieraForPreTraining
1181
+ >>> import torch
1182
+ >>> from PIL import Image
1183
+ >>> import requests
1184
+
1185
+ >>> url = "http://images.cocodataset.org/val2017/000000039769.jpg"
1186
+ >>> image = Image.open(requests.get(url, stream=True).raw)
1187
+
1188
+ >>> image_processor = AutoImageProcessor.from_pretrained("facebook/hiera-tiny-224-mae-hf")
1189
+ >>> model = HieraForPreTraining.from_pretrained("facebook/hiera-tiny-224-mae-hf")
1190
+
1191
+ >>> inputs = image_processor(images=image, return_tensors="pt")
1192
+
1193
+ >>> outputs = model(**inputs)
1194
+ >>> logits = outputs.logits
1195
+ >>> loss = outputs.loss
1196
+ >>> print(list(logits.shape))
1197
+ [1, 196, 768]
1198
+ ```"""
1199
+ return_dict = return_dict if return_dict is not None else self.config.use_return_dict
1200
+ output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
1201
+ output_hidden_states = (
1202
+ output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
1203
+ )
1204
+
1205
+ outputs = self.hiera(
1206
+ pixel_values,
1207
+ noise=noise,
1208
+ head_mask=head_mask,
1209
+ output_attentions=output_attentions,
1210
+ output_hidden_states=True,
1211
+ interpolate_pos_encoding=interpolate_pos_encoding,
1212
+ return_dict=return_dict,
1213
+ )
1214
+
1215
+ feature_maps = outputs[-1]
1216
+ bool_masked_pos = outputs[1]
1217
+ ids_to_restore = outputs[2]
1218
+ # Take only the query pooled and last hidden states
1219
+ feature_maps = feature_maps[1 : self.hiera.config.num_query_pool + 1] + (feature_maps[-1],)
1220
+ fused_hidden_states = self.multiscale_fusion(feature_maps)
1221
+ fused_hidden_states = self.encoder_norm(fused_hidden_states)
1222
+
1223
+ # Reconstruct pixel values
1224
+ logits, bool_masked_pos = self.decoder(
1225
+ fused_hidden_states,
1226
+ bool_masked_pos=bool_masked_pos,
1227
+ head_mask=head_mask,
1228
+ output_attentions=output_attentions,
1229
+ )
1230
+
1231
+ loss = self.forward_loss(pixel_values, logits, bool_masked_pos)
1232
+
1233
+ if not return_dict:
1234
+ output = (logits, bool_masked_pos, ids_to_restore)
1235
+ if output_hidden_states:
1236
+ output = output + (outputs[3],)
1237
+ if output_attentions:
1238
+ output = output + (outputs[4],)
1239
+ if output_hidden_states:
1240
+ output = output + (outputs[-1],)
1241
+ return ((loss,) + output) if loss is not None else output
1242
+
1243
+ return HieraForPreTrainingOutput(
1244
+ loss=loss,
1245
+ logits=logits,
1246
+ bool_masked_pos=bool_masked_pos,
1247
+ ids_restore=ids_to_restore,
1248
+ hidden_states=outputs.hidden_states if output_hidden_states else None,
1249
+ attentions=outputs.attentions,
1250
+ reshaped_hidden_states=outputs.reshaped_hidden_states if output_hidden_states else None,
1251
+ )
1252
+
1253
+
1254
+ @auto_docstring(
1255
+ custom_intro="""
1256
+ Hiera Model transformer with an image classification head on top (a linear layer on top of the final hidden state with
1257
+ average pooling) e.g. for ImageNet.
1258
+
1259
+ <Tip>
1260
+
1261
+ Note that it's possible to fine-tune Hiera on higher resolution images than the ones it has been trained on, by
1262
+ setting `interpolate_pos_encoding` to `True` in the forward of the model. This will interpolate the pre-trained
1263
+ position embeddings to the higher resolution.
1264
+
1265
+ </Tip>
1266
+ """
1267
+ )
1268
+ class HieraForImageClassification(HieraPreTrainedModel):
1269
+ def __init__(self, config: HieraConfig) -> None:
1270
+ super().__init__(config)
1271
+
1272
+ self.num_labels = config.num_labels
1273
+ self.hiera = HieraModel(config, add_pooling_layer=True, is_mae=False)
1274
+
1275
+ # Classifier head
1276
+ self.classifier = (
1277
+ nn.Linear(self.hiera.num_features, config.num_labels) if config.num_labels > 0 else nn.Identity()
1278
+ )
1279
+
1280
+ # Initialize weights and apply final processing
1281
+ self.post_init()
1282
+
1283
+ @auto_docstring
1284
+ def forward(
1285
+ self,
1286
+ pixel_values,
1287
+ head_mask: Optional[torch.Tensor] = None,
1288
+ labels: Optional[torch.Tensor] = None,
1289
+ output_attentions: Optional[bool] = None,
1290
+ output_hidden_states: Optional[bool] = None,
1291
+ interpolate_pos_encoding: Optional[bool] = None,
1292
+ return_dict: Optional[bool] = None,
1293
+ ) -> Union[tuple, HieraForImageClassificationOutput]:
1294
+ r"""
1295
+ labels (`torch.LongTensor` of shape `(batch_size,)`, *optional*):
1296
+ Labels for computing the image classification/regression loss. Indices should be in `[0, ...,
1297
+ config.num_labels - 1]`. If `config.num_labels == 1` a regression loss is computed (Mean-Square loss), If
1298
+ `config.num_labels > 1` a classification loss is computed (Cross-Entropy).
1299
+ """
1300
+ return_dict = return_dict if return_dict is not None else self.config.use_return_dict
1301
+ output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
1302
+ output_hidden_states = (
1303
+ output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
1304
+ )
1305
+
1306
+ outputs = self.hiera(
1307
+ pixel_values,
1308
+ head_mask=head_mask,
1309
+ output_attentions=output_attentions,
1310
+ output_hidden_states=output_hidden_states,
1311
+ interpolate_pos_encoding=interpolate_pos_encoding,
1312
+ return_dict=return_dict,
1313
+ )
1314
+
1315
+ pooled_output = outputs[1]
1316
+
1317
+ logits = self.classifier(pooled_output)
1318
+
1319
+ loss = None
1320
+ if labels is not None:
1321
+ loss = self.loss_function(labels, logits, self.config)
1322
+
1323
+ if not return_dict:
1324
+ output = (logits,) + outputs[2:]
1325
+ return ((loss,) + output) if loss is not None else output
1326
+
1327
+ return HieraForImageClassificationOutput(
1328
+ loss=loss,
1329
+ logits=logits,
1330
+ hidden_states=outputs.hidden_states,
1331
+ attentions=outputs.attentions,
1332
+ reshaped_hidden_states=outputs.reshaped_hidden_states,
1333
+ )
1334
+
1335
+
1336
+ @auto_docstring(
1337
+ custom_intro="""
1338
+ Hiera backbone, to be used with frameworks like DETR and MaskFormer.
1339
+ """
1340
+ )
1341
+ class HieraBackbone(HieraPreTrainedModel, BackboneMixin):
1342
+ def __init__(self, config: HieraConfig):
1343
+ super().__init__(config)
1344
+ super()._init_backbone(config)
1345
+
1346
+ self.num_features = [config.embed_dim] + [
1347
+ int(config.embed_dim * config.embed_dim_multiplier**i) for i in range(len(config.depths))
1348
+ ]
1349
+ self.embeddings = HieraEmbeddings(config, is_mae=False)
1350
+ self.encoder = HieraEncoder(config)
1351
+
1352
+ # Add layer norms to hidden states of out_features
1353
+ hidden_states_norms = {}
1354
+ for stage, num_channels in zip(self._out_features, self.channels):
1355
+ hidden_states_norms[stage] = nn.LayerNorm(num_channels)
1356
+ self.hidden_states_norms = nn.ModuleDict(hidden_states_norms)
1357
+
1358
+ # Initialize weights and apply final processing
1359
+ self.post_init()
1360
+
1361
+ def get_input_embeddings(self):
1362
+ return self.embeddings.patch_embeddings
1363
+
1364
+ def forward(
1365
+ self,
1366
+ pixel_values: torch.Tensor,
1367
+ output_hidden_states: Optional[bool] = None,
1368
+ output_attentions: Optional[bool] = None,
1369
+ return_dict: Optional[bool] = None,
1370
+ ) -> BackboneOutput:
1371
+ """
1372
+ Returns:
1373
+
1374
+ Examples:
1375
+
1376
+ ```python
1377
+ >>> from transformers import AutoImageProcessor, AutoBackbone
1378
+ >>> import torch
1379
+ >>> from PIL import Image
1380
+ >>> import requests
1381
+
1382
+ >>> url = "http://images.cocodataset.org/val2017/000000039769.jpg"
1383
+ >>> image = Image.open(requests.get(url, stream=True).raw)
1384
+
1385
+ >>> processor = AutoImageProcessor.from_pretrained("facebook/hiera-tiny-224-hf")
1386
+ >>> model = AutoBackbone.from_pretrained(
1387
+ ... "facebook/hiera-tiny-224-hf", out_features=["stage1", "stage2", "stage3", "stage4"]
1388
+ ... )
1389
+
1390
+ >>> inputs = processor(image, return_tensors="pt")
1391
+ >>> outputs = model(**inputs)
1392
+ >>> feature_maps = outputs.feature_maps
1393
+ >>> list(feature_maps[-1].shape)
1394
+ [1, 768, 7, 7]
1395
+ ```"""
1396
+ return_dict = return_dict if return_dict is not None else self.config.use_return_dict
1397
+ output_hidden_states = (
1398
+ output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
1399
+ )
1400
+ output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
1401
+
1402
+ embedding_output, _, _ = self.embeddings(pixel_values)
1403
+
1404
+ outputs = self.encoder(
1405
+ embedding_output,
1406
+ head_mask=None,
1407
+ output_attentions=output_attentions,
1408
+ output_hidden_states=True,
1409
+ return_dict=return_dict,
1410
+ )
1411
+
1412
+ hidden_states = outputs[-1]
1413
+
1414
+ feature_maps = ()
1415
+ for stage, hidden_state in zip(self.stage_names, hidden_states):
1416
+ if stage in self.out_features:
1417
+ batch_size, height, width, num_channels = hidden_state.shape
1418
+ hidden_state = hidden_state.view(batch_size, height * width, num_channels)
1419
+ hidden_state = self.hidden_states_norms[stage](hidden_state)
1420
+ hidden_state = hidden_state.view(batch_size, height, width, num_channels)
1421
+ hidden_state = hidden_state.permute(0, 3, 1, 2).contiguous()
1422
+ feature_maps += (hidden_state,)
1423
+
1424
+ if not return_dict:
1425
+ output = (feature_maps,)
1426
+ if output_hidden_states:
1427
+ output += (outputs[1],)
1428
+ if output_attentions:
1429
+ output += (outputs[2],)
1430
+ return output
1431
+
1432
+ return BackboneOutput(
1433
+ feature_maps=feature_maps,
1434
+ hidden_states=outputs[1] if output_hidden_states else None,
1435
+ attentions=outputs[2] if output_attentions else None,
1436
+ )
1437
+
1438
+
1439
+ __all__ = ["HieraForImageClassification", "HieraForPreTraining", "HieraBackbone", "HieraModel", "HieraPreTrainedModel"]