Aluode/PerceptionLabPortable
0
1# coding=utf-82# Copyright 2022 Meta Platforms 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 at7#8# http://www.apache.org/licenses/LICENSE-2.09#10# Unless required by applicable law or agreed to in writing, software11# 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 and14# limitations under the License.15"""PyTorch Data2VecVision model."""16 17import collections.abc18import math19import warnings20from dataclasses import dataclass21from typing import Optional, Union22 23import torch24from torch import nn25from torch.nn import CrossEntropyLoss26 27from ...activations import ACT2FN28from ...modeling_layers import GradientCheckpointingLayer29from ...modeling_outputs import (30 BaseModelOutput,31 BaseModelOutputWithPooling,32 ImageClassifierOutput,33 SemanticSegmenterOutput,34)35from ...modeling_utils import PreTrainedModel36from ...pytorch_utils import compile_compatible_method_lru_cache, find_pruneable_heads_and_indices, prune_linear_layer37from ...utils import auto_docstring, logging, torch_int38from .configuration_data2vec_vision import Data2VecVisionConfig39 40 41logger = logging.get_logger(__name__)42 43 44@dataclass45@auto_docstring(46 custom_intro="""47 Class for outputs of [`Data2VecVisionModel`].48 """49)50# Copied from transformers.models.beit.modeling_beit.BeitModelOutputWithPooling with Beit->Data2VecVision51class Data2VecVisionModelOutputWithPooling(BaseModelOutputWithPooling):52 r"""53 pooler_output (`torch.FloatTensor` of shape `(batch_size, hidden_size)`):54 Average of the last layer hidden states of the patch tokens (excluding the *[CLS]* token) if55 *config.use_mean_pooling* is set to True. If set to False, then the final hidden state of the *[CLS]* token56 will be returned.57 """58 59 60# Copied from transformers.models.beit.modeling_beit.drop_path61def drop_path(input: torch.Tensor, drop_prob: float = 0.0, training: bool = False) -> torch.Tensor:62 """63 Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks).64 65 Comment by Ross Wightman: This is the same as the DropConnect impl I created for EfficientNet, etc networks,66 however, the original name is misleading as 'Drop Connect' is a different form of dropout in a separate paper...67 See discussion: https://github.com/tensorflow/tpu/issues/494#issuecomment-532968956 ... I've opted for changing the68 layer and argument names to 'drop path' rather than mix DropConnect as a layer name and use 'survival rate' as the69 argument.70 """71 if drop_prob == 0.0 or not training:72 return input73 keep_prob = 1 - drop_prob74 shape = (input.shape[0],) + (1,) * (input.ndim - 1) # work with diff dim tensors, not just 2D ConvNets75 random_tensor = keep_prob + torch.rand(shape, dtype=input.dtype, device=input.device)76 random_tensor.floor_() # binarize77 output = input.div(keep_prob) * random_tensor78 return output79 80 81# Copied from transformers.models.beit.modeling_beit.BeitDropPath with Beit->Data2VecVision82class Data2VecVisionDropPath(nn.Module):83 """Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks)."""84 85 def __init__(self, drop_prob: Optional[float] = None) -> None:86 super().__init__()87 self.drop_prob = drop_prob88 89 def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:90 return drop_path(hidden_states, self.drop_prob, self.training)91 92 def extra_repr(self) -> str:93 return f"p={self.drop_prob}"94 95 96# Copied from transformers.models.beit.modeling_beit.BeitEmbeddings with Beit->Data2VecVision97class Data2VecVisionEmbeddings(nn.Module):98 """99 Construct the CLS token, position and patch embeddings. Optionally, also the mask token.100 101 """102 103 def __init__(self, config: Data2VecVisionConfig) -> None:104 super().__init__()105 106 self.cls_token = nn.Parameter(torch.zeros(1, 1, config.hidden_size))107 if config.use_mask_token:108 self.mask_token = nn.Parameter(torch.zeros(1, 1, config.hidden_size))109 else:110 self.mask_token = None111 self.patch_embeddings = Data2VecVisionPatchEmbeddings(config)112 self.patch_size = config.patch_size113 self.image_size = (114 config.image_size115 if isinstance(config.image_size, collections.abc.Iterable)116 else (config.image_size, config.image_size)117 )118 num_patches = self.patch_embeddings.num_patches119 if config.use_absolute_position_embeddings:120 self.position_embeddings = nn.Parameter(torch.zeros(1, num_patches + 1, config.hidden_size))121 else:122 self.position_embeddings = None123 self.dropout = nn.Dropout(config.hidden_dropout_prob)124 125 # Copied from transformers.models.vit.modeling_vit.ViTEmbeddings.interpolate_pos_encoding126 def interpolate_pos_encoding(self, embeddings: torch.Tensor, height: int, width: int) -> torch.Tensor:127 """128 This method allows to interpolate the pre-trained position encodings, to be able to use the model on higher resolution129 images. This method is also adapted to support torch.jit tracing.130 131 Adapted from:132 - https://github.com/facebookresearch/dino/blob/de9ee3df6cf39fac952ab558447af1fa1365362a/vision_transformer.py#L174-L194, and133 - https://github.com/facebookresearch/dinov2/blob/e1277af2ba9496fbadf7aec6eba56e8d882d1e35/dinov2/models/vision_transformer.py#L179-L211134 """135 136 num_patches = embeddings.shape[1] - 1137 num_positions = self.position_embeddings.shape[1] - 1138 139 # always interpolate when tracing to ensure the exported model works for dynamic input shapes140 if not torch.jit.is_tracing() and num_patches == num_positions and height == width:141 return self.position_embeddings142 143 class_pos_embed = self.position_embeddings[:, :1]144 patch_pos_embed = self.position_embeddings[:, 1:]145 146 dim = embeddings.shape[-1]147 148 new_height = height // self.patch_size149 new_width = width // self.patch_size150 151 sqrt_num_positions = torch_int(num_positions**0.5)152 patch_pos_embed = patch_pos_embed.reshape(1, sqrt_num_positions, sqrt_num_positions, dim)153 patch_pos_embed = patch_pos_embed.permute(0, 3, 1, 2)154 155 patch_pos_embed = nn.functional.interpolate(156 patch_pos_embed,157 size=(new_height, new_width),158 mode="bicubic",159 align_corners=False,160 )161 162 patch_pos_embed = patch_pos_embed.permute(0, 2, 3, 1).view(1, -1, dim)163 164 return torch.cat((class_pos_embed, patch_pos_embed), dim=1)165 166 def forward(167 self,168 pixel_values: torch.Tensor,169 bool_masked_pos: Optional[torch.BoolTensor] = None,170 interpolate_pos_encoding: Optional[bool] = None,171 ) -> torch.Tensor:172 if self.position_embeddings is not None and interpolate_pos_encoding is not None:173 warnings.warn(174 "`interpolate_pos_encoding` argument has no effect for BEiTEmbeddings, embeddings are always "175 "interpolated to the input image size. The argument will be removed in transformers v4.51.0."176 )177 178 _, _, height, width = pixel_values.shape179 embeddings, (patch_height, patch_width) = self.patch_embeddings(pixel_values)180 batch_size, seq_len, _ = embeddings.size()181 182 if bool_masked_pos is not None:183 mask_tokens = self.mask_token.expand(batch_size, seq_len, -1)184 # replace the masked visual tokens by mask_tokens185 w = bool_masked_pos.unsqueeze(-1).type_as(mask_tokens)186 embeddings = embeddings * (1 - w) + mask_tokens * w187 188 cls_tokens = self.cls_token.expand(batch_size, -1, -1)189 embeddings = torch.cat((cls_tokens, embeddings), dim=1)190 191 if self.position_embeddings is not None:192 embeddings = embeddings + self.interpolate_pos_encoding(embeddings, height, width)193 194 embeddings = self.dropout(embeddings)195 196 return embeddings, (patch_height, patch_width)197 198 199# Copied from transformers.models.beit.modeling_beit.BeitPatchEmbeddings with Beit->Data2VecVision200class Data2VecVisionPatchEmbeddings(nn.Module):201 """202 This class turns `pixel_values` of shape `(batch_size, num_channels, height, width)` into the initial203 `hidden_states` (patch embeddings) of shape `(batch_size, seq_length, hidden_size)` to be consumed by a204 Transformer.205 """206 207 def __init__(self, config):208 super().__init__()209 image_size, patch_size = config.image_size, config.patch_size210 num_channels, hidden_size = config.num_channels, config.hidden_size211 212 image_size = image_size if isinstance(image_size, collections.abc.Iterable) else (image_size, image_size)213 patch_size = patch_size if isinstance(patch_size, collections.abc.Iterable) else (patch_size, patch_size)214 num_patches = (image_size[1] // patch_size[1]) * (image_size[0] // patch_size[0])215 patch_shape = (image_size[0] // patch_size[0], image_size[1] // patch_size[1])216 self.image_size = image_size217 self.patch_size = patch_size218 self.num_channels = num_channels219 self.num_patches = num_patches220 self.patch_shape = patch_shape221 222 self.projection = nn.Conv2d(num_channels, hidden_size, kernel_size=patch_size, stride=patch_size)223 224 def forward(self, pixel_values: torch.Tensor) -> torch.Tensor:225 batch_size, num_channels, height, width = pixel_values.shape226 if num_channels != self.num_channels:227 raise ValueError(228 "Make sure that the channel dimension of the pixel values match with the one set in the configuration."229 )230 231 embeddings = self.projection(pixel_values)232 patch_height, patch_width = embeddings.shape[2], embeddings.shape[3]233 embeddings = embeddings.flatten(2).transpose(1, 2)234 235 return embeddings, (patch_height, patch_width)236 237 238# Copied from transformers.models.beit.modeling_beit.BeitSelfAttention with Beit->Data2VecVision239class Data2VecVisionSelfAttention(nn.Module):240 def __init__(self, config: Data2VecVisionConfig, window_size: Optional[tuple] = None) -> None:241 super().__init__()242 self.config = config243 if config.hidden_size % config.num_attention_heads != 0 and not hasattr(config, "embedding_size"):244 raise ValueError(245 f"The hidden size {config.hidden_size} is not a multiple of the number of attention "246 f"heads {config.num_attention_heads}."247 )248 249 self.num_attention_heads = config.num_attention_heads250 self.attention_head_size = int(config.hidden_size / config.num_attention_heads)251 self.all_head_size = self.num_attention_heads * self.attention_head_size252 253 self.query = nn.Linear(config.hidden_size, self.all_head_size)254 self.key = nn.Linear(config.hidden_size, self.all_head_size, bias=False)255 self.value = nn.Linear(config.hidden_size, self.all_head_size)256 257 self.dropout = nn.Dropout(config.attention_probs_dropout_prob)258 259 self.has_relative_position_bias = bool(window_size)260 if self.has_relative_position_bias:261 self.relative_position_bias = Data2VecVisionRelativePositionBias(config, window_size=window_size)262 263 def forward(264 self,265 hidden_states: torch.Tensor,266 head_mask: Optional[torch.Tensor] = None,267 output_attentions: bool = False,268 relative_position_bias: Optional[torch.Tensor] = None,269 interpolate_pos_encoding: bool = False,270 resolution: Optional[tuple[int]] = None,271 ) -> Union[tuple[torch.Tensor], tuple[torch.Tensor, torch.Tensor]]:272 batch_size, seq_length, _ = hidden_states.shape273 query_layer = (274 self.query(hidden_states)275 .view(batch_size, -1, self.num_attention_heads, self.attention_head_size)276 .transpose(1, 2)277 )278 key_layer = (279 self.key(hidden_states)280 .view(batch_size, -1, self.num_attention_heads, self.attention_head_size)281 .transpose(1, 2)282 )283 value_layer = (284 self.value(hidden_states)285 .view(batch_size, -1, self.num_attention_heads, self.attention_head_size)286 .transpose(1, 2)287 )288 289 # Take the dot product between "query" and "key" to get the raw attention scores.290 attention_scores = torch.matmul(query_layer, key_layer.transpose(-1, -2))291 292 attention_scores = attention_scores / math.sqrt(self.attention_head_size)293 294 # Add relative position bias if present.295 if self.has_relative_position_bias:296 height, width = resolution297 window_size = (height // self.config.patch_size, width // self.config.patch_size)298 attention_scores = attention_scores + self.relative_position_bias(299 window_size, interpolate_pos_encoding, dim_size=hidden_states.shape[1]300 )301 302 # Add shared relative position bias if provided.303 if relative_position_bias is not None:304 attention_scores = attention_scores + relative_position_bias305 306 # Normalize the attention scores to probabilities.307 attention_probs = nn.functional.softmax(attention_scores, dim=-1)308 309 # This is actually dropping out entire tokens to attend to, which might310 # seem a bit unusual, but is taken from the original Transformer paper.311 attention_probs = self.dropout(attention_probs)312 313 # Mask heads if we want to314 if head_mask is not None:315 attention_probs = attention_probs * head_mask316 317 context_layer = torch.matmul(attention_probs, value_layer)318 319 context_layer = context_layer.permute(0, 2, 1, 3).contiguous()320 new_context_layer_shape = context_layer.size()[:-2] + (self.all_head_size,)321 context_layer = context_layer.view(*new_context_layer_shape)322 323 outputs = (context_layer, attention_probs) if output_attentions else (context_layer,)324 325 return outputs326 327 328# Copied from transformers.models.beit.modeling_beit.BeitSdpaSelfAttention with Beit->Data2VecVision329class Data2VecVisionSdpaSelfAttention(Data2VecVisionSelfAttention):330 def forward(331 self,332 hidden_states: torch.Tensor,333 head_mask: Optional[torch.Tensor] = None,334 output_attentions: bool = False,335 relative_position_bias: Optional[torch.Tensor] = None,336 interpolate_pos_encoding: bool = False,337 resolution: Optional[tuple[int]] = None,338 ) -> Union[tuple[torch.Tensor], tuple[torch.Tensor, torch.Tensor]]:339 if output_attentions or head_mask is not None:340 logger.warning_once(341 "`Data2VecVisionSdpaSelfAttention` is used but `torch.nn.functional.scaled_dot_product_attention` does not "342 "support `output_attentions=True` or `head_mask`. Falling back to the manual attention implementation, "343 "but specifying the manual implementation will be required from Transformers version v5.0.0 onwards. "344 'This warning can be removed using the argument `attn_implementation="eager"` when loading the model.'345 )346 return super().forward(347 hidden_states=hidden_states,348 head_mask=head_mask,349 output_attentions=output_attentions,350 relative_position_bias=relative_position_bias,351 interpolate_pos_encoding=interpolate_pos_encoding,352 resolution=resolution,353 )354 355 batch_size, seq_length, _ = hidden_states.shape356 query_layer = (357 self.query(hidden_states)358 .view(batch_size, -1, self.num_attention_heads, self.attention_head_size)359 .transpose(1, 2)360 )361 key_layer = (362 self.key(hidden_states)363 .view(batch_size, -1, self.num_attention_heads, self.attention_head_size)364 .transpose(1, 2)365 )366 value_layer = (367 self.value(hidden_states)368 .view(batch_size, -1, self.num_attention_heads, self.attention_head_size)369 .transpose(1, 2)370 )371 372 attn_bias = None373 if self.has_relative_position_bias:374 height, width = resolution375 window_size = (height // self.config.patch_size, width // self.config.patch_size)376 attn_bias = self.relative_position_bias(377 window_size, interpolate_pos_encoding, dim_size=hidden_states.shape[1]378 )379 380 # Add shared relative position bias if provided.381 if relative_position_bias is not None:382 if attn_bias is None:383 attn_bias = relative_position_bias384 else:385 attn_bias += relative_position_bias386 387 scaling = 1 / math.sqrt(self.attention_head_size)388 context_layer = torch.nn.functional.scaled_dot_product_attention(389 query_layer,390 key_layer,391 value_layer,392 attn_mask=attn_bias,393 dropout_p=self.config.attention_probs_dropout_prob if self.training else 0.0,394 is_causal=False,395 scale=scaling,396 )397 context_layer = context_layer.permute(0, 2, 1, 3).contiguous()398 new_context_layer_shape = context_layer.size()[:-2] + (self.all_head_size,)399 context_layer = context_layer.view(*new_context_layer_shape)400 return context_layer, None401 402 403# Copied from transformers.models.beit.modeling_beit.BeitSelfOutput with Beit->Data2VecVision404class Data2VecVisionSelfOutput(nn.Module):405 """406 The residual connection is defined in Data2VecVisionLayer instead of here (as is the case with other models), due to the407 layernorm applied before each block.408 """409 410 def __init__(self, config: Data2VecVisionConfig) -> None:411 super().__init__()412 self.dense = nn.Linear(config.hidden_size, config.hidden_size)413 self.dropout = nn.Dropout(config.hidden_dropout_prob)414 415 def forward(self, hidden_states: torch.Tensor, input_tensor: torch.Tensor, gamma=None) -> torch.Tensor:416 hidden_states = self.dense(hidden_states)417 hidden_states = self.dropout(hidden_states)418 419 return hidden_states420 421 422DATA2VEC_VISION_SELF_ATTENTION_CLASSES = {423 "eager": Data2VecVisionSelfAttention,424 "sdpa": Data2VecVisionSdpaSelfAttention,425}426 427 428# Copied from tests.models.beit.modeling_beit.BeitAttention with Beit->Data2VecVision, BEIT->DATA2VEC_VISION429class Data2VecVisionAttention(nn.Module):430 def __init__(self, config: Data2VecVisionConfig, window_size: Optional[tuple] = None) -> None:431 super().__init__()432 self.attention = DATA2VEC_VISION_SELF_ATTENTION_CLASSES[config._attn_implementation](433 config, window_size=window_size434 )435 self.output = Data2VecVisionSelfOutput(config)436 self.pruned_heads = set()437 438 def prune_heads(self, heads):439 if len(heads) == 0:440 return441 heads, index = find_pruneable_heads_and_indices(442 heads, self.attention.num_attention_heads, self.attention.attention_head_size, self.pruned_heads443 )444 445 # Prune linear layers446 self.attention.query = prune_linear_layer(self.attention.query, index)447 self.attention.key = prune_linear_layer(self.attention.key, index)448 self.attention.value = prune_linear_layer(self.attention.value, index)449 self.output.dense = prune_linear_layer(self.output.dense, index, dim=1)450 451 # Update hyper params and store pruned heads452 self.attention.num_attention_heads = self.attention.num_attention_heads - len(heads)453 self.attention.all_head_size = self.attention.attention_head_size * self.attention.num_attention_heads454 self.pruned_heads = self.pruned_heads.union(heads)455 456 def forward(457 self,458 hidden_states: torch.Tensor,459 head_mask: Optional[torch.Tensor] = None,460 output_attentions: bool = False,461 relative_position_bias: Optional["Data2VecVisionRelativePositionBias"] = None,462 interpolate_pos_encoding: bool = False,463 resolution: Optional[tuple[int]] = None,464 ) -> Union[tuple[torch.Tensor], tuple[torch.Tensor, torch.Tensor]]:465 self_outputs = self.attention(466 hidden_states, head_mask, output_attentions, relative_position_bias, interpolate_pos_encoding, resolution467 )468 469 attention_output = self.output(self_outputs[0], hidden_states)470 471 outputs = (attention_output,) + self_outputs[1:] # add attentions if we output them472 return outputs473 474 475# Copied from transformers.models.beit.modeling_beit.BeitIntermediate with Beit->Data2VecVision476class Data2VecVisionIntermediate(nn.Module):477 def __init__(self, config: Data2VecVisionConfig) -> None:478 super().__init__()479 self.dense = nn.Linear(config.hidden_size, config.intermediate_size)480 if isinstance(config.hidden_act, str):481 self.intermediate_act_fn = ACT2FN[config.hidden_act]482 else:483 self.intermediate_act_fn = config.hidden_act484 485 def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:486 hidden_states = self.dense(hidden_states)487 hidden_states = self.intermediate_act_fn(hidden_states)488 489 return hidden_states490 491 492# Copied from transformers.models.beit.modeling_beit.BeitOutput with Beit->Data2VecVision493class Data2VecVisionOutput(nn.Module):494 def __init__(self, config: Data2VecVisionConfig) -> None:495 super().__init__()496 self.dense = nn.Linear(config.intermediate_size, config.hidden_size)497 self.dropout = nn.Dropout(config.hidden_dropout_prob)498 499 def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:500 hidden_states = self.dense(hidden_states)501 hidden_states = self.dropout(hidden_states)502 503 return hidden_states504 505 506# Copied from transformers.models.beit.modeling_beit.BeitLayer with Beit->Data2VecVision,BEiT->Data2VecVision507class Data2VecVisionLayer(GradientCheckpointingLayer):508 """This corresponds to the Block class in the timm implementation."""509 510 def __init__(511 self, config: Data2VecVisionConfig, window_size: Optional[tuple] = None, drop_path_rate: float = 0.0512 ) -> None:513 super().__init__()514 self.chunk_size_feed_forward = config.chunk_size_feed_forward515 self.seq_len_dim = 1516 self.attention = Data2VecVisionAttention(config, window_size=window_size)517 self.intermediate = Data2VecVisionIntermediate(config)518 self.output = Data2VecVisionOutput(config)519 self.layernorm_before = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)520 self.drop_path = Data2VecVisionDropPath(drop_path_rate) if drop_path_rate > 0.0 else nn.Identity()521 self.layernorm_after = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)522 523 init_values = config.layer_scale_init_value524 if init_values > 0:525 self.lambda_1 = nn.Parameter(init_values * torch.ones(config.hidden_size), requires_grad=True)526 self.lambda_2 = nn.Parameter(init_values * torch.ones(config.hidden_size), requires_grad=True)527 else:528 self.lambda_1, self.lambda_2 = None, None529 530 def forward(531 self,532 hidden_states: torch.Tensor,533 head_mask: Optional[torch.Tensor] = None,534 output_attentions: bool = False,535 relative_position_bias: Optional[torch.Tensor] = None,536 interpolate_pos_encoding: bool = False,537 resolution: Optional[tuple[int, int]] = None,538 ) -> Union[tuple[torch.Tensor], tuple[torch.Tensor, torch.Tensor]]:539 self_attention_outputs = self.attention(540 self.layernorm_before(hidden_states), # in Data2VecVision, layernorm is applied before self-attention541 head_mask,542 output_attentions=output_attentions,543 relative_position_bias=relative_position_bias,544 interpolate_pos_encoding=interpolate_pos_encoding,545 resolution=resolution,546 )547 attention_output = self_attention_outputs[0]548 outputs = self_attention_outputs[1:] # add self attentions if we output attention weights549 550 # apply lambda_1 if present551 if self.lambda_1 is not None:552 attention_output = self.lambda_1 * attention_output553 554 # first residual connection555 hidden_states = self.drop_path(attention_output) + hidden_states556 557 # in Data2VecVision, layernorm is also applied after self-attention558 layer_output = self.layernorm_after(hidden_states)559 560 layer_output = self.intermediate(layer_output)561 layer_output = self.output(layer_output)562 563 if self.lambda_2 is not None:564 layer_output = self.lambda_2 * layer_output565 566 # second residual connection567 layer_output = self.drop_path(layer_output) + hidden_states568 569 outputs = (layer_output,) + outputs570 571 return outputs572 573 574# Copied from transformers.models.beit.modeling_beit.BeitRelativePositionBias with Beit->Data2VecVision575class Data2VecVisionRelativePositionBias(nn.Module):576 def __init__(self, config: Data2VecVisionConfig, window_size: tuple) -> None:577 super().__init__()578 self.window_size = window_size579 self.num_relative_distance = (2 * window_size[0] - 1) * (2 * window_size[1] - 1) + 3580 self.relative_position_bias_table = nn.Parameter(581 torch.zeros(self.num_relative_distance, config.num_attention_heads)582 ) # 2*Wh-1 * 2*Ww-1, nH583 # cls to token & token 2 cls & cls to cls584 585 @compile_compatible_method_lru_cache(maxsize=10)586 def generate_relative_position_index(self, window_size: tuple[int, int]) -> torch.Tensor:587 """588 This method creates the relative position index, modified to support arbitrary window sizes,589 as introduced in [MiDaS v3.1](https://huggingface.co/papers/2307.14460).590 """591 num_relative_distance = (2 * window_size[0] - 1) * (2 * window_size[1] - 1) + 3592 # cls to token & token 2 cls & cls to cls593 # get pair-wise relative position index for each token inside the window594 window_area = window_size[0] * window_size[1]595 grid = torch.meshgrid(torch.arange(window_size[0]), torch.arange(window_size[1]), indexing="ij")596 coords = torch.stack(grid) # 2, Wh, Ww597 coords_flatten = torch.flatten(coords, 1) # 2, Wh*Ww598 relative_coords = coords_flatten[:, :, None] - coords_flatten[:, None, :] # 2, Wh*Ww, Wh*Ww599 relative_coords = relative_coords.permute(1, 2, 0).contiguous() # Wh*Ww, Wh*Ww, 2600 relative_coords[:, :, 0] += window_size[0] - 1 # shift to start from 0601 relative_coords[:, :, 1] += window_size[1] - 1602 relative_coords[:, :, 0] *= 2 * window_size[1] - 1603 relative_position_index = torch.zeros(size=(window_area + 1,) * 2, dtype=relative_coords.dtype)604 relative_position_index[1:, 1:] = relative_coords.sum(-1) # Wh*Ww, Wh*Ww605 relative_position_index[0, 0:] = num_relative_distance - 3606 relative_position_index[0:, 0] = num_relative_distance - 2607 relative_position_index[0, 0] = num_relative_distance - 1608 return relative_position_index609 610 def forward(self, window_size, interpolate_pos_encoding: bool = False, dim_size=None) -> torch.Tensor:611 """612 Modification of timm.models.beit.py: Attention._get_rel_pos_bias to support arbitrary window sizes.613 """614 old_height = 2 * self.window_size[0] - 1615 old_width = 2 * self.window_size[1] - 1616 617 new_height = 2 * window_size[0] - 1618 new_width = 2 * window_size[1] - 1619 620 old_relative_position_bias_table = self.relative_position_bias_table621 622 old_num_relative_distance = self.num_relative_distance623 new_num_relative_distance = new_height * new_width + 3624 625 old_sub_table = old_relative_position_bias_table[: old_num_relative_distance - 3]626 627 old_sub_table = old_sub_table.reshape(1, old_width, old_height, -1).permute(0, 3, 1, 2)628 new_sub_table = nn.functional.interpolate(629 old_sub_table, size=(torch_int(new_height), torch_int(new_width)), mode="bilinear"630 )631 new_sub_table = new_sub_table.permute(0, 2, 3, 1).reshape(new_num_relative_distance - 3, -1)632 633 new_relative_position_bias_table = torch.cat(634 [new_sub_table, old_relative_position_bias_table[old_num_relative_distance - 3 :]]635 )636 637 relative_position_index = self.generate_relative_position_index(window_size)638 relative_position_bias = new_relative_position_bias_table[relative_position_index.view(-1)]639 640 # patch_size*num_patches_height, patch_size*num_patches_width, num_attention_heads641 relative_position_bias = relative_position_bias.view(642 window_size[0] * window_size[1] + 1, window_size[0] * window_size[1] + 1, -1643 )644 # num_attention_heads, patch_size*num_patches_width, patch_size*num_patches_height645 relative_position_bias = relative_position_bias.permute(2, 0, 1).contiguous()646 647 if interpolate_pos_encoding:648 relative_position_bias = nn.functional.interpolate(649 relative_position_bias.unsqueeze(1),650 size=(dim_size, dim_size),651 mode="bilinear",652 align_corners=False,653 ).squeeze(1)654 655 return relative_position_bias.unsqueeze(0)656 657 658# Copied from transformers.models.beit.modeling_beit.BeitEncoder with Beit->Data2VecVision659class Data2VecVisionEncoder(nn.Module):660 def __init__(self, config: Data2VecVisionConfig, window_size: Optional[tuple] = None) -> None:661 super().__init__()662 self.config = config663 self.has_relative_position_bias = config.use_shared_relative_position_bias664 if self.has_relative_position_bias:665 self.relative_position_bias = Data2VecVisionRelativePositionBias(config, window_size=window_size)666 667 # stochastic depth decay rule668 dpr = [x.item() for x in torch.linspace(0, config.drop_path_rate, config.num_hidden_layers, device="cpu")]669 self.layer = nn.ModuleList(670 [671 Data2VecVisionLayer(672 config,673 window_size=window_size if config.use_relative_position_bias else None,674 drop_path_rate=dpr[i],675 )676 for i in range(config.num_hidden_layers)677 ]678 )679 self.gradient_checkpointing = False680 681 def forward(682 self,683 hidden_states: torch.Tensor,684 head_mask: Optional[torch.Tensor] = None,685 output_attentions: bool = False,686 output_hidden_states: bool = False,687 interpolate_pos_encoding: bool = False,688 resolution: Optional[tuple[int, int]] = None,689 return_dict: bool = True,690 ) -> Union[tuple, BaseModelOutput]:691 all_hidden_states = () if output_hidden_states else None692 all_self_attentions = () if output_attentions else None693 694 for i, layer_module in enumerate(self.layer):695 if output_hidden_states:696 all_hidden_states = all_hidden_states + (hidden_states,)697 698 if self.has_relative_position_bias:699 height, width = resolution700 window_size = (height // self.config.patch_size, width // self.config.patch_size)701 relative_position_bias = self.relative_position_bias(702 window_size, interpolate_pos_encoding=interpolate_pos_encoding, dim_size=hidden_states.shape[1]703 )704 else:705 relative_position_bias = None706 707 layer_head_mask = head_mask[i] if head_mask is not None else None708 709 layer_outputs = layer_module(710 hidden_states,711 head_mask=layer_head_mask,712 output_attentions=output_attentions,713 relative_position_bias=relative_position_bias,714 interpolate_pos_encoding=interpolate_pos_encoding,715 resolution=resolution,716 )717 718 hidden_states = layer_outputs[0]719 720 if output_attentions:721 all_self_attentions = all_self_attentions + (layer_outputs[1],)722 723 if output_hidden_states:724 all_hidden_states = all_hidden_states + (hidden_states,)725 726 if not return_dict:727 return tuple(v for v in [hidden_states, all_hidden_states, all_self_attentions] if v is not None)728 return BaseModelOutput(729 last_hidden_state=hidden_states,730 hidden_states=all_hidden_states,731 attentions=all_self_attentions,732 )733 734 735@auto_docstring736# Copied from transformers.models.beit.modeling_beit.BeitPreTrainedModel with Beit->Data2VecVision,beit->data2vec_vision737class Data2VecVisionPreTrainedModel(PreTrainedModel):738 config: Data2VecVisionConfig739 base_model_prefix = "data2vec_vision"740 main_input_name = "pixel_values"741 supports_gradient_checkpointing = True742 _no_split_modules = ["Data2VecVisionLayer"]743 _keys_to_ignore_on_load_unexpected = [r".*relative_position_index.*"]744 _supports_sdpa = True745 746 def _init_weights(self, module):747 """Initialize the weights"""748 if isinstance(module, (nn.Linear, nn.Conv2d, nn.ConvTranspose2d)):749 # Slightly different from the TF version which uses truncated_normal for initialization750 # cf https://github.com/pytorch/pytorch/pull/5617751 module.weight.data.normal_(mean=0.0, std=self.config.initializer_range)752 if module.bias is not None:753 module.bias.data.zero_()754 elif isinstance(module, nn.Embedding):755 module.weight.data.normal_(mean=0.0, std=self.config.initializer_range)756 if module.padding_idx is not None:757 module.weight.data[module.padding_idx].zero_()758 elif isinstance(module, nn.LayerNorm):759 module.bias.data.zero_()760 module.weight.data.fill_(1.0)761 elif isinstance(module, Data2VecVisionEmbeddings):762 module.cls_token.data.zero_()763 if module.mask_token is not None:764 module.mask_token.data.zero_()765 if module.position_embeddings is not None:766 module.position_embeddings.data.zero_()767 elif isinstance(module, Data2VecVisionRelativePositionBias):768 module.relative_position_bias_table.data.zero_()769 elif isinstance(module, Data2VecVisionLayer):770 if module.lambda_1 is not None:771 module.lambda_1.data.fill_(self.config.layer_scale_init_value)772 module.lambda_2.data.fill_(self.config.layer_scale_init_value)773 774 775@auto_docstring776# Copied from transformers.models.beit.modeling_beit.BeitModel with BEIT->DATA2VEC_VISION,Beit->Data2VecVision,True->False777class Data2VecVisionModel(Data2VecVisionPreTrainedModel):778 def __init__(self, config: Data2VecVisionConfig, add_pooling_layer: bool = False) -> None:779 r"""780 add_pooling_layer (bool, *optional*, defaults to `False`):781 Whether to add a pooling layer782 """783 super().__init__(config)784 self.config = config785 786 self.embeddings = Data2VecVisionEmbeddings(config)787 self.encoder = Data2VecVisionEncoder(config, window_size=self.embeddings.patch_embeddings.patch_shape)788 789 self.layernorm = (790 nn.Identity() if config.use_mean_pooling else nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)791 )792 self.pooler = Data2VecVisionPooler(config) if add_pooling_layer else None793 794 # Initialize weights and apply final processing795 self.post_init()796 797 def get_input_embeddings(self):798 return self.embeddings.patch_embeddings799 800 def _prune_heads(self, heads_to_prune):801 """802 Prunes heads of the model. heads_to_prune: dict of {layer_num: list of heads to prune in this layer} See base803 class PreTrainedModel804 """805 for layer, heads in heads_to_prune.items():806 self.encoder.layer[layer].attention.prune_heads(heads)807 808 @auto_docstring809 def forward(810 self,811 pixel_values: torch.Tensor,812 bool_masked_pos: Optional[torch.BoolTensor] = None,813 head_mask: Optional[torch.Tensor] = None,814 output_attentions: Optional[bool] = None,815 output_hidden_states: Optional[bool] = None,816 interpolate_pos_encoding: bool = False,817 return_dict: Optional[bool] = None,818 ) -> Union[tuple, Data2VecVisionModelOutputWithPooling]:819 r"""820 bool_masked_pos (`torch.BoolTensor` of shape `(batch_size, num_patches)`, *optional*):821 Boolean masked positions. Indicates which patches are masked (1) and which aren't (0).822 """823 output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions824 output_hidden_states = (825 output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states826 )827 return_dict = return_dict if return_dict is not None else self.config.use_return_dict828 829 # Prepare head mask if needed830 # 1.0 in head_mask indicate we keep the head831 # attention_probs has shape bsz x n_heads x N x N832 # input head_mask has shape [num_heads] or [num_hidden_layers x num_heads]833 # and head_mask is converted to shape [num_hidden_layers x batch x num_heads x seq_length x seq_length]834 head_mask = self.get_head_mask(head_mask, self.config.num_hidden_layers)835 836 embedding_output, _ = self.embeddings(pixel_values, bool_masked_pos=bool_masked_pos)837 resolution = pixel_values.shape[2:]838 839 encoder_outputs = self.encoder(840 embedding_output,841 head_mask=head_mask,842 output_attentions=output_attentions,843 output_hidden_states=output_hidden_states,844 resolution=resolution,845 return_dict=return_dict,846 interpolate_pos_encoding=interpolate_pos_encoding,847 )848 sequence_output = encoder_outputs[0]849 sequence_output = self.layernorm(sequence_output)850 pooled_output = self.pooler(sequence_output) if self.pooler is not None else None851 852 if not return_dict:853 head_outputs = (sequence_output, pooled_output) if pooled_output is not None else (sequence_output,)854 return head_outputs + encoder_outputs[1:]855 856 return Data2VecVisionModelOutputWithPooling(857 last_hidden_state=sequence_output,858 pooler_output=pooled_output,859 hidden_states=encoder_outputs.hidden_states,860 attentions=encoder_outputs.attentions,861 )862 863 864# Copied from transformers.models.beit.modeling_beit.BeitPooler with Beit->Data2VecVision865class Data2VecVisionPooler(nn.Module):866 def __init__(self, config: Data2VecVisionConfig) -> None:867 super().__init__()868 self.layernorm = (869 nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps) if config.use_mean_pooling else None870 )871 872 def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:873 if self.layernorm is not None:874 # Mean pool the final hidden states of the patch tokens875 patch_tokens = hidden_states[:, 1:, :]876 pooled_output = self.layernorm(patch_tokens.mean(1))877 else:878 # Pool by simply taking the final hidden state of the [CLS] token879 pooled_output = hidden_states[:, 0]880 881 return pooled_output882 883 884@auto_docstring(885 custom_intro="""886 Data2VecVision Model transformer with an image classification head on top (a linear layer on top of the average of887 the final hidden states of the patch tokens) e.g. for ImageNet.888 """889)890# Copied from transformers.models.beit.modeling_beit.BeitForImageClassification with BEIT->DATA2VEC_VISION,Beit->Data2VecVision,beit->data2vec_vision891class Data2VecVisionForImageClassification(Data2VecVisionPreTrainedModel):892 def __init__(self, config: Data2VecVisionConfig) -> None:893 super().__init__(config)894 895 self.num_labels = config.num_labels896 self.data2vec_vision = Data2VecVisionModel(config, add_pooling_layer=True)897 898 # Classifier head899 self.classifier = nn.Linear(config.hidden_size, config.num_labels) if config.num_labels > 0 else nn.Identity()900 901 # Initialize weights and apply final processing902 self.post_init()903 904 @auto_docstring905 def forward(906 self,907 pixel_values: Optional[torch.Tensor] = None,908 head_mask: Optional[torch.Tensor] = None,909 labels: Optional[torch.Tensor] = None,910 output_attentions: Optional[bool] = None,911 output_hidden_states: Optional[bool] = None,912 interpolate_pos_encoding: bool = False,913 return_dict: Optional[bool] = None,914 ) -> Union[tuple, ImageClassifierOutput]:915 r"""916 labels (`torch.LongTensor` of shape `(batch_size,)`, *optional*):917 Labels for computing the image classification/regression loss. Indices should be in `[0, ...,918 config.num_labels - 1]`. If `config.num_labels == 1` a regression loss is computed (Mean-Square loss), If919 `config.num_labels > 1` a classification loss is computed (Cross-Entropy).920 """921 return_dict = return_dict if return_dict is not None else self.config.use_return_dict922 outputs = self.data2vec_vision(923 pixel_values,924 head_mask=head_mask,925 output_attentions=output_attentions,926 output_hidden_states=output_hidden_states,927 interpolate_pos_encoding=interpolate_pos_encoding,928 return_dict=return_dict,929 )930 931 pooled_output = outputs.pooler_output if return_dict else outputs[1]932 933 logits = self.classifier(pooled_output)934 935 loss = None936 if labels is not None:937 loss = self.loss_function(labels, logits, self.config)938 939 if not return_dict:940 output = (logits,) + outputs[2:]941 return ((loss,) + output) if loss is not None else output942 943 return ImageClassifierOutput(944 loss=loss,945 logits=logits,946 hidden_states=outputs.hidden_states,947 attentions=outputs.attentions,948 )949 950 951# Copied from transformers.models.beit.modeling_beit.BeitConvModule with Beit->Data2VecVision952class Data2VecVisionConvModule(nn.Module):953 """954 A convolutional block that bundles conv/norm/activation layers. This block simplifies the usage of convolution955 layers, which are commonly used with a norm layer (e.g., BatchNorm) and activation layer (e.g., ReLU).956 957 Based on OpenMMLab's implementation, found in https://github.com/open-mmlab/mmsegmentation.958 """959 960 def __init__(961 self,962 in_channels: int,963 out_channels: int,964 kernel_size: Union[int, tuple[int, int]],965 padding: Union[int, tuple[int, int], str] = 0,966 bias: bool = False,967 dilation: Union[int, tuple[int, int]] = 1,968 ) -> None:969 super().__init__()970 self.conv = nn.Conv2d(971 in_channels=in_channels,972 out_channels=out_channels,973 kernel_size=kernel_size,974 padding=padding,975 bias=bias,976 dilation=dilation,977 )978 self.bn = nn.BatchNorm2d(out_channels)979 self.activation = nn.ReLU()980 981 def forward(self, input: torch.Tensor) -> torch.Tensor:982 output = self.conv(input)983 output = self.bn(output)984 output = self.activation(output)985 986 return output987 988 989# Copied from transformers.models.beit.modeling_beit.BeitPyramidPoolingBlock with Beit->Data2VecVision990class Data2VecVisionPyramidPoolingBlock(nn.Module):991 def __init__(self, pool_scale: int, in_channels: int, channels: int) -> None:992 super().__init__()993 self.layers = [994 nn.AdaptiveAvgPool2d(pool_scale),995 Data2VecVisionConvModule(in_channels, channels, kernel_size=1),996 ]997 for i, layer in enumerate(self.layers):998 self.add_module(str(i), layer)999 1000 def forward(self, input: torch.Tensor) -> torch.Tensor:1001 hidden_state = input1002 for layer in self.layers:1003 hidden_state = layer(hidden_state)1004 return hidden_state1005 1006 1007# Copied from transformers.models.beit.modeling_beit.BeitPyramidPoolingModule with Beit->Data2VecVision1008class Data2VecVisionPyramidPoolingModule(nn.Module):1009 """1010 Pyramid Pooling Module (PPM) used in PSPNet.1011 1012 Args:1013 pool_scales (tuple[int]): Pooling scales used in Pooling Pyramid1014 Module.1015 in_channels (int): Input channels.1016 channels (int): Channels after modules, before conv_seg.1017 align_corners (bool): align_corners argument of F.interpolate.1018 1019 Based on OpenMMLab's implementation, found in https://github.com/open-mmlab/mmsegmentation.1020 """1021 1022 def __init__(self, pool_scales: tuple[int, ...], in_channels: int, channels: int, align_corners: bool) -> None:1023 super().__init__()1024 self.pool_scales = pool_scales1025 self.align_corners = align_corners1026 self.in_channels = in_channels1027 self.channels = channels1028 self.blocks = []1029 for i, pool_scale in enumerate(pool_scales):1030 block = Data2VecVisionPyramidPoolingBlock(1031 pool_scale=pool_scale, in_channels=in_channels, channels=channels1032 )1033 self.blocks.append(block)1034 self.add_module(str(i), block)1035 1036 def forward(self, x: torch.Tensor) -> list[torch.Tensor]:1037 ppm_outs = []1038 for ppm in self.blocks:1039 ppm_out = ppm(x)1040 upsampled_ppm_out = nn.functional.interpolate(1041 ppm_out, size=x.size()[2:], mode="bilinear", align_corners=self.align_corners1042 )1043 ppm_outs.append(upsampled_ppm_out)1044 return ppm_outs1045 1046 1047# Copied from transformers.models.beit.modeling_beit.BeitUperHead with Beit->Data2VecVision1048class Data2VecVisionUperHead(nn.Module):1049 """1050 Unified Perceptual Parsing for Scene Understanding. This head is the implementation of1051 [UPerNet](https://huggingface.co/papers/1807.10221).1052 1053 Based on OpenMMLab's implementation, found in https://github.com/open-mmlab/mmsegmentation.1054 """1055 1056 def __init__(self, config: Data2VecVisionConfig) -> None:1057 super().__init__()1058 1059 self.pool_scales = config.pool_scales # e.g. (1, 2, 3, 6)1060 self.in_channels = [config.hidden_size] * 4 # e.g. [768, 768, 768, 768]1061 self.channels = config.hidden_size1062 self.align_corners = False1063 self.classifier = nn.Conv2d(self.channels, config.num_labels, kernel_size=1)1064 1065 # PSP Module1066 self.psp_modules = Data2VecVisionPyramidPoolingModule(1067 self.pool_scales,1068 self.in_channels[-1],1069 self.channels,1070 align_corners=self.align_corners,1071 )1072 self.bottleneck = Data2VecVisionConvModule(1073 self.in_channels[-1] + len(self.pool_scales) * self.channels,1074 self.channels,1075 kernel_size=3,1076 padding=1,1077 )1078 # FPN Module1079 self.lateral_convs = nn.ModuleList()1080 self.fpn_convs = nn.ModuleList()1081 for in_channels in self.in_channels[:-1]: # skip the top layer1082 l_conv = Data2VecVisionConvModule(in_channels, self.channels, kernel_size=1)1083 fpn_conv = Data2VecVisionConvModule(self.channels, self.channels, kernel_size=3, padding=1)1084 self.lateral_convs.append(l_conv)1085 self.fpn_convs.append(fpn_conv)1086 1087 self.fpn_bottleneck = Data2VecVisionConvModule(1088 len(self.in_channels) * self.channels,1089 self.channels,1090 kernel_size=3,1091 padding=1,1092 )1093 1094 def psp_forward(self, inputs):1095 x = inputs[-1]1096 psp_outs = [x]1097 psp_outs.extend(self.psp_modules(x))1098 psp_outs = torch.cat(psp_outs, dim=1)1099 output = self.bottleneck(psp_outs)1100 1101 return output1102 1103 def forward(self, encoder_hidden_states: torch.Tensor) -> torch.Tensor:1104 # build laterals1105 laterals = [lateral_conv(encoder_hidden_states[i]) for i, lateral_conv in enumerate(self.lateral_convs)]1106 1107 laterals.append(self.psp_forward(encoder_hidden_states))1108 1109 # build top-down path1110 used_backbone_levels = len(laterals)1111 for i in range(used_backbone_levels - 1, 0, -1):1112 prev_shape = laterals[i - 1].shape[2:]1113 laterals[i - 1] = laterals[i - 1] + nn.functional.interpolate(1114 laterals[i], size=prev_shape, mode="bilinear", align_corners=self.align_corners1115 )1116 1117 # build outputs1118 fpn_outs = [self.fpn_convs[i](laterals[i]) for i in range(used_backbone_levels - 1)]1119 # append psp feature1120 fpn_outs.append(laterals[-1])1121 1122 for i in range(used_backbone_levels - 1, 0, -1):1123 fpn_outs[i] = nn.functional.interpolate(1124 fpn_outs[i], size=fpn_outs[0].shape[2:], mode="bilinear", align_corners=self.align_corners1125 )1126 fpn_outs = torch.cat(fpn_outs, dim=1)1127 output = self.fpn_bottleneck(fpn_outs)1128 output = self.classifier(output)1129 1130 return output1131 1132 1133# Copied from transformers.models.beit.modeling_beit.BeitFCNHead with Beit->Data2VecVision1134class Data2VecVisionFCNHead(nn.Module):1135 """1136 Fully Convolution Networks for Semantic Segmentation. This head is implemented of1137 [FCNNet](https://huggingface.co/papers/1411.4038>).1138 1139 Args:1140 config (Data2VecVisionConfig): Configuration.1141 in_channels1142 kernel_size (int): The kernel size for convs in the head. Default: 3.1143 dilation (int): The dilation rate for convs in the head. Default: 1.1144 1145 1146 Based on OpenMMLab's implementation, found in https://github.com/open-mmlab/mmsegmentation.1147 """1148 1149 def __init__(1150 self,1151 config: Data2VecVisionConfig,1152 in_index: int = 2,1153 kernel_size: int = 3,1154 dilation: Union[int, tuple[int, int]] = 1,1155 ) -> None:1156 super().__init__()1157 self.in_channels = config.hidden_size1158 self.channels = config.auxiliary_channels1159 self.num_convs = config.auxiliary_num_convs1160 self.concat_input = config.auxiliary_concat_input1161 self.in_index = in_index1162 1163 conv_padding = (kernel_size // 2) * dilation1164 convs = []1165 convs.append(1166 Data2VecVisionConvModule(1167 self.in_channels, self.channels, kernel_size=kernel_size, padding=conv_padding, dilation=dilation1168 )1169 )1170 for i in range(self.num_convs - 1):1171 convs.append(1172 Data2VecVisionConvModule(1173 self.channels, self.channels, kernel_size=kernel_size, padding=conv_padding, dilation=dilation1174 )1175 )1176 if self.num_convs == 0:1177 self.convs = nn.Identity()1178 else:1179 self.convs = nn.Sequential(*convs)1180 if self.concat_input:1181 self.conv_cat = Data2VecVisionConvModule(1182 self.in_channels + self.channels, self.channels, kernel_size=kernel_size, padding=kernel_size // 21183 )1184 1185 self.classifier = nn.Conv2d(self.channels, config.num_labels, kernel_size=1)1186 1187 def forward(self, encoder_hidden_states: torch.Tensor) -> torch.Tensor:1188 # just take the relevant feature maps1189 hidden_states = encoder_hidden_states[self.in_index]1190 output = self.convs(hidden_states)1191 if self.concat_input:1192 output = self.conv_cat(torch.cat([hidden_states, output], dim=1))1193 output = self.classifier(output)1194 return output1195 1196 1197@auto_docstring1198# Copied from transformers.models.beit.modeling_beit.BeitForSemanticSegmentation with BEIT->DATA2VEC_VISION,Beit->Data2VecVision,microsoft/beit-base-finetuned-ade-640-640->facebook/data2vec-vision-base,beit->data2vec_vision1199class Data2VecVisionForSemanticSegmentation(Data2VecVisionPreTrainedModel):1200 def __init__(self, config: Data2VecVisionConfig) -> None: