radna commited on
Commit
456c716
·
verified ·
1 Parent(s): cd3db47

Update modeling_intern_vit.py

Browse files
Files changed (1) hide show
  1. modeling_intern_vit.py +415 -415
modeling_intern_vit.py CHANGED
@@ -1,415 +1,415 @@
1
- # --------------------------------------------------------
2
- # InternVL
3
- # Copyright (c) 2023 OpenGVLab
4
- # Licensed under The MIT License [see LICENSE for details]
5
- # --------------------------------------------------------
6
- from typing import Optional, Tuple, Union
7
-
8
- import torch
9
- import torch.nn.functional as F
10
- import torch.utils.checkpoint
11
- from einops import rearrange
12
- from timm.models.layers import DropPath
13
- from torch import nn
14
- from transformers.activations import ACT2FN
15
- from transformers.modeling_outputs import BaseModelOutput, BaseModelOutputWithPooling
16
- from transformers.modeling_utils import PreTrainedModel
17
- from transformers.utils import logging
18
-
19
- from .configuration_intern_vit import InternVisionConfig
20
-
21
- try:
22
- from .triton_flash_attn import attention
23
-
24
- has_flash_attn = True
25
- except:
26
- print("attention is not installed.")
27
- has_flash_attn = False
28
-
29
-
30
- logger = logging.get_logger(__name__)
31
-
32
-
33
- class InternRMSNorm(nn.Module):
34
- def __init__(self, hidden_size, eps=1e-6):
35
- super().__init__()
36
- self.weight = nn.Parameter(torch.ones(hidden_size))
37
- self.variance_epsilon = eps
38
-
39
- def forward(self, hidden_states):
40
- input_dtype = hidden_states.dtype
41
- hidden_states = hidden_states.to(torch.float32)
42
- variance = hidden_states.pow(2).mean(-1, keepdim=True)
43
- hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)
44
- return self.weight * hidden_states.to(input_dtype)
45
-
46
-
47
- try:
48
- from apex.normalization import FusedRMSNorm
49
-
50
- InternRMSNorm = FusedRMSNorm # noqa
51
-
52
- logger.info(
53
- "Discovered apex.normalization.FusedRMSNorm - will use it instead of InternRMSNorm"
54
- )
55
- except ImportError:
56
- # using the normal InternRMSNorm
57
- pass
58
- except Exception:
59
- logger.warning(
60
- "discovered apex but it failed to load, falling back to InternRMSNorm"
61
- )
62
- pass
63
-
64
-
65
- class InternVisionEmbeddings(nn.Module):
66
- def __init__(self, config: InternVisionConfig):
67
- super().__init__()
68
- self.config = config
69
- self.embed_dim = config.hidden_size
70
- self.image_size = config.image_size
71
- self.patch_size = config.patch_size
72
-
73
- self.class_embedding = nn.Parameter(
74
- torch.randn(1, 1, self.embed_dim),
75
- )
76
-
77
- self.patch_embedding = nn.Conv2d(
78
- in_channels=3,
79
- out_channels=self.embed_dim,
80
- kernel_size=self.patch_size,
81
- stride=self.patch_size,
82
- )
83
-
84
- self.num_patches = (self.image_size // self.patch_size) ** 2
85
- self.num_positions = self.num_patches + 1
86
-
87
- self.position_embedding = nn.Parameter(
88
- torch.randn(1, self.num_positions, self.embed_dim)
89
- )
90
-
91
- def forward(self, pixel_values: torch.FloatTensor) -> torch.Tensor:
92
- batch_size = pixel_values.shape[0]
93
- target_dtype = self.patch_embedding.weight.dtype
94
- patch_embeds = self.patch_embedding(
95
- pixel_values
96
- ) # shape = [*, width, grid, grid]
97
- patch_embeds = patch_embeds.flatten(2).transpose(1, 2)
98
- class_embeds = self.class_embedding.expand(batch_size, 1, -1).to(target_dtype)
99
- embeddings = torch.cat([class_embeds, patch_embeds], dim=1)
100
- embeddings = embeddings + self.position_embedding.to(target_dtype)
101
- return embeddings
102
-
103
-
104
- class InternAttention(nn.Module):
105
- """Multi-headed attention from 'Attention Is All You Need' paper"""
106
-
107
- def __init__(self, config: InternVisionConfig):
108
- super().__init__()
109
- self.config = config
110
- self.embed_dim = config.hidden_size
111
- self.num_heads = config.num_attention_heads
112
- self.use_flash_attn = config.use_flash_attn and has_flash_attn
113
- if config.use_flash_attn and not has_flash_attn:
114
- print(
115
- "Warning: Flash Attention is not available, use_flash_attn is set to False."
116
- )
117
- self.head_dim = self.embed_dim // self.num_heads
118
- if self.head_dim * self.num_heads != self.embed_dim:
119
- raise ValueError(
120
- f"embed_dim must be divisible by num_heads (got `embed_dim`: {self.embed_dim} and `num_heads`:"
121
- f" {self.num_heads})."
122
- )
123
-
124
- self.scale = self.head_dim**-0.5
125
- self.qkv = nn.Linear(self.embed_dim, 3 * self.embed_dim, bias=config.qkv_bias)
126
- self.attn_drop = nn.Dropout(config.attention_dropout)
127
- self.proj_drop = nn.Dropout(config.dropout)
128
-
129
- self.qk_normalization = config.qk_normalization
130
-
131
- if self.qk_normalization:
132
- self.q_norm = InternRMSNorm(self.embed_dim, eps=config.layer_norm_eps)
133
- self.k_norm = InternRMSNorm(self.embed_dim, eps=config.layer_norm_eps)
134
-
135
- if self.use_flash_attn:
136
- self.inner_attn = attention(attention_dropout=config.attention_dropout)
137
- self.proj = nn.Linear(self.embed_dim, self.embed_dim)
138
-
139
- def _naive_attn(self, x):
140
- B, N, C = x.shape
141
- qkv = (
142
- self.qkv(x)
143
- .reshape(B, N, 3, self.num_heads, C // self.num_heads)
144
- .permute(2, 0, 3, 1, 4)
145
- )
146
- q, k, v = qkv.unbind(0) # make torchscript happy (cannot use tensor as tuple)
147
-
148
- if self.qk_normalization:
149
- B_, H_, N_, D_ = q.shape
150
- q = (
151
- self.q_norm(q.transpose(1, 2).flatten(-2, -1))
152
- .view(B_, N_, H_, D_)
153
- .transpose(1, 2)
154
- )
155
- k = (
156
- self.k_norm(k.transpose(1, 2).flatten(-2, -1))
157
- .view(B_, N_, H_, D_)
158
- .transpose(1, 2)
159
- )
160
-
161
- attn = (q * self.scale) @ k.transpose(-2, -1)
162
- attn = attn.softmax(dim=-1)
163
- attn = self.attn_drop(attn)
164
-
165
- x = (attn @ v).transpose(1, 2).reshape(B, N, C)
166
- x = self.proj(x)
167
- x = self.proj_drop(x)
168
- return x
169
-
170
- def _flash_attn(self, x, key_padding_mask=None, need_weights=False):
171
- qkv = self.qkv(x)
172
- qkv = rearrange(
173
- qkv, "b s (three h d) -> b s three h d", three=3, h=self.num_heads
174
- )
175
-
176
- if self.qk_normalization:
177
- q, k, v = qkv.unbind(2)
178
- q = self.q_norm(q.flatten(-2, -1)).view(q.shape)
179
- k = self.k_norm(k.flatten(-2, -1)).view(k.shape)
180
- qkv = torch.stack([q, k, v], dim=2)
181
-
182
- context, _ = self.inner_attn(
183
- qkv,
184
- key_padding_mask=key_padding_mask,
185
- need_weights=need_weights,
186
- causal=False,
187
- )
188
- outs = self.proj(rearrange(context, "b s h d -> b s (h d)"))
189
- outs = self.proj_drop(outs)
190
- return outs
191
-
192
- def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
193
- x = (
194
- self._naive_attn(hidden_states)
195
- if not self.use_flash_attn
196
- else self._flash_attn(hidden_states)
197
- )
198
- return x
199
-
200
-
201
- class InternMLP(nn.Module):
202
- def __init__(self, config: InternVisionConfig):
203
- super().__init__()
204
- self.config = config
205
- self.act = ACT2FN[config.hidden_act]
206
- self.fc1 = nn.Linear(config.hidden_size, config.intermediate_size)
207
- self.fc2 = nn.Linear(config.intermediate_size, config.hidden_size)
208
-
209
- def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
210
- hidden_states = self.fc1(hidden_states)
211
- hidden_states = self.act(hidden_states)
212
- hidden_states = self.fc2(hidden_states)
213
- return hidden_states
214
-
215
-
216
- class InternVisionEncoderLayer(nn.Module):
217
- def __init__(self, config: InternVisionConfig, drop_path_rate: float):
218
- super().__init__()
219
- self.embed_dim = config.hidden_size
220
- self.intermediate_size = config.intermediate_size
221
-
222
- self.attn = InternAttention(config)
223
- self.mlp = InternMLP(config)
224
- self.norm1 = InternRMSNorm(self.embed_dim, eps=config.layer_norm_eps)
225
- self.norm2 = InternRMSNorm(self.embed_dim, eps=config.layer_norm_eps)
226
-
227
- self.ls1 = nn.Parameter(config.initializer_factor * torch.ones(self.embed_dim))
228
- self.ls2 = nn.Parameter(config.initializer_factor * torch.ones(self.embed_dim))
229
- self.drop_path1 = (
230
- DropPath(drop_path_rate) if drop_path_rate > 0.0 else nn.Identity()
231
- )
232
- self.drop_path2 = (
233
- DropPath(drop_path_rate) if drop_path_rate > 0.0 else nn.Identity()
234
- )
235
-
236
- def forward(
237
- self,
238
- hidden_states: torch.Tensor,
239
- ) -> Tuple[
240
- torch.FloatTensor,
241
- Optional[torch.FloatTensor],
242
- Optional[Tuple[torch.FloatTensor]],
243
- ]:
244
- """
245
- Args:
246
- hidden_states (`Tuple[torch.FloatTensor, Optional[torch.FloatTensor]]`): input to the layer of shape `(batch, seq_len, embed_dim)`
247
- """
248
- hidden_states = hidden_states + self.drop_path1(
249
- self.attn(self.norm1(hidden_states)) * self.ls1
250
- )
251
-
252
- hidden_states = hidden_states + self.drop_path2(
253
- self.mlp(self.norm2(hidden_states)) * self.ls2
254
- )
255
-
256
- return hidden_states
257
-
258
-
259
- class InternVisionEncoder(nn.Module):
260
- """
261
- Transformer encoder consisting of `config.num_hidden_layers` self attention layers. Each layer is a
262
- [`InternEncoderLayer`].
263
-
264
- Args:
265
- config (`InternConfig`):
266
- The corresponding vision configuration for the `InternEncoder`.
267
- """
268
-
269
- def __init__(self, config: InternVisionConfig):
270
- super().__init__()
271
- self.config = config
272
- # stochastic depth decay rule
273
- dpr = [
274
- x.item()
275
- for x in torch.linspace(0, config.drop_path_rate, config.num_hidden_layers)
276
- ]
277
- self.layers = nn.ModuleList(
278
- [
279
- InternVisionEncoderLayer(config, dpr[idx])
280
- for idx in range(config.num_hidden_layers)
281
- ]
282
- )
283
- self.gradient_checkpointing = True
284
-
285
- def forward(
286
- self,
287
- inputs_embeds,
288
- output_hidden_states: Optional[bool] = None,
289
- return_dict: Optional[bool] = None,
290
- ) -> Union[Tuple, BaseModelOutput]:
291
- r"""
292
- Args:
293
- inputs_embeds (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`):
294
- Embedded representation of the inputs. Should be float, not int tokens.
295
- output_hidden_states (`bool`, *optional*):
296
- Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors
297
- for more detail.
298
- return_dict (`bool`, *optional*):
299
- Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.
300
- """
301
- output_hidden_states = (
302
- output_hidden_states
303
- if output_hidden_states is not None
304
- else self.config.output_hidden_states
305
- )
306
- return_dict = (
307
- return_dict if return_dict is not None else self.config.use_return_dict
308
- )
309
-
310
- encoder_states = () if output_hidden_states else None
311
- hidden_states = inputs_embeds
312
-
313
- for idx, encoder_layer in enumerate(self.layers):
314
- if output_hidden_states:
315
- encoder_states = encoder_states + (hidden_states,)
316
- if self.gradient_checkpointing and self.training:
317
- layer_outputs = torch.utils.checkpoint.checkpoint(
318
- encoder_layer, hidden_states
319
- )
320
- else:
321
- layer_outputs = encoder_layer(
322
- hidden_states,
323
- )
324
- hidden_states = layer_outputs
325
-
326
- if output_hidden_states:
327
- encoder_states = encoder_states + (hidden_states,)
328
-
329
- if not return_dict:
330
- return tuple(v for v in [hidden_states, encoder_states] if v is not None)
331
- return BaseModelOutput(
332
- last_hidden_state=hidden_states, hidden_states=encoder_states
333
- )
334
-
335
-
336
- class InternVisionModel(PreTrainedModel):
337
- main_input_name = "pixel_values"
338
- config_class = InternVisionConfig
339
- _no_split_modules = ["InternVisionEncoderLayer"]
340
-
341
- def __init__(self, config: InternVisionConfig):
342
- super().__init__(config)
343
- self.config = config
344
-
345
- self.embeddings = InternVisionEmbeddings(config)
346
- self.encoder = InternVisionEncoder(config)
347
-
348
- def resize_pos_embeddings(self, old_size, new_size, patch_size):
349
- pos_emb = self.embeddings.position_embedding
350
- _, num_positions, embed_dim = pos_emb.shape
351
- cls_emb = pos_emb[:, :1, :]
352
- pos_emb = (
353
- pos_emb[:, 1:, :]
354
- .reshape(1, old_size // patch_size, old_size // patch_size, -1)
355
- .permute(0, 3, 1, 2)
356
- )
357
- pos_emb = F.interpolate(
358
- pos_emb.float(),
359
- size=new_size // patch_size,
360
- mode="bicubic",
361
- align_corners=False,
362
- )
363
- pos_emb = pos_emb.to(cls_emb.dtype).reshape(1, embed_dim, -1).permute(0, 2, 1)
364
- pos_emb = torch.cat([cls_emb, pos_emb], dim=1)
365
- self.embeddings.position_embedding = nn.Parameter(pos_emb)
366
- logger.info(
367
- "Resized position embeddings from {} to {}".format(old_size, new_size)
368
- )
369
-
370
- def get_input_embeddings(self):
371
- return self.embeddings
372
-
373
- def forward(
374
- self,
375
- pixel_values: Optional[torch.FloatTensor] = None,
376
- output_hidden_states: Optional[bool] = None,
377
- return_dict: Optional[bool] = None,
378
- pixel_embeds: Optional[torch.FloatTensor] = None,
379
- ) -> Union[Tuple, BaseModelOutputWithPooling]:
380
- output_hidden_states = (
381
- output_hidden_states
382
- if output_hidden_states is not None
383
- else self.config.output_hidden_states
384
- )
385
- return_dict = (
386
- return_dict if return_dict is not None else self.config.use_return_dict
387
- )
388
-
389
- if pixel_values is None and pixel_embeds is None:
390
- raise ValueError("You have to specify pixel_values or pixel_embeds")
391
-
392
- if pixel_embeds is not None:
393
- hidden_states = pixel_embeds
394
- else:
395
- if len(pixel_values.shape) == 4:
396
- hidden_states = self.embeddings(pixel_values)
397
- else:
398
- raise ValueError(f"wrong pixel_values size: {pixel_values.shape}")
399
- encoder_outputs = self.encoder(
400
- inputs_embeds=hidden_states,
401
- output_hidden_states=output_hidden_states,
402
- return_dict=return_dict,
403
- )
404
- last_hidden_state = encoder_outputs.last_hidden_state
405
- pooled_output = last_hidden_state[:, 0, :]
406
-
407
- if not return_dict:
408
- return (last_hidden_state, pooled_output) + encoder_outputs[1:]
409
-
410
- return BaseModelOutputWithPooling(
411
- last_hidden_state=last_hidden_state,
412
- pooler_output=pooled_output,
413
- hidden_states=encoder_outputs.hidden_states,
414
- attentions=encoder_outputs.attentions,
415
- )
 
1
+ # --------------------------------------------------------
2
+ # InternVL
3
+ # Copyright (c) 2023 OpenGVLab
4
+ # Licensed under The MIT License [see LICENSE for details]
5
+ # --------------------------------------------------------
6
+ from typing import Optional, Tuple, Union
7
+
8
+ import torch
9
+ import torch.nn.functional as F
10
+ import torch.utils.checkpoint
11
+ from einops import rearrange
12
+ from timm.models.layers import DropPath
13
+ from torch import nn
14
+ from transformers.activations import ACT2FN
15
+ from transformers.modeling_outputs import BaseModelOutput, BaseModelOutputWithPooling
16
+ from transformers.modeling_utils import PreTrainedModel
17
+ from transformers.utils import logging
18
+
19
+ from .configuration_intern_vit import InternVisionConfig
20
+
21
+ try:
22
+ from .triton_flash_attn import _attention
23
+
24
+ has_flash_attn = True
25
+ except:
26
+ print("attention is not installed.")
27
+ has_flash_attn = False
28
+
29
+
30
+ logger = logging.get_logger(__name__)
31
+
32
+
33
+ class InternRMSNorm(nn.Module):
34
+ def __init__(self, hidden_size, eps=1e-6):
35
+ super().__init__()
36
+ self.weight = nn.Parameter(torch.ones(hidden_size))
37
+ self.variance_epsilon = eps
38
+
39
+ def forward(self, hidden_states):
40
+ input_dtype = hidden_states.dtype
41
+ hidden_states = hidden_states.to(torch.float32)
42
+ variance = hidden_states.pow(2).mean(-1, keepdim=True)
43
+ hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)
44
+ return self.weight * hidden_states.to(input_dtype)
45
+
46
+
47
+ try:
48
+ from apex.normalization import FusedRMSNorm
49
+
50
+ InternRMSNorm = FusedRMSNorm # noqa
51
+
52
+ logger.info(
53
+ "Discovered apex.normalization.FusedRMSNorm - will use it instead of InternRMSNorm"
54
+ )
55
+ except ImportError:
56
+ # using the normal InternRMSNorm
57
+ pass
58
+ except Exception:
59
+ logger.warning(
60
+ "discovered apex but it failed to load, falling back to InternRMSNorm"
61
+ )
62
+ pass
63
+
64
+
65
+ class InternVisionEmbeddings(nn.Module):
66
+ def __init__(self, config: InternVisionConfig):
67
+ super().__init__()
68
+ self.config = config
69
+ self.embed_dim = config.hidden_size
70
+ self.image_size = config.image_size
71
+ self.patch_size = config.patch_size
72
+
73
+ self.class_embedding = nn.Parameter(
74
+ torch.randn(1, 1, self.embed_dim),
75
+ )
76
+
77
+ self.patch_embedding = nn.Conv2d(
78
+ in_channels=3,
79
+ out_channels=self.embed_dim,
80
+ kernel_size=self.patch_size,
81
+ stride=self.patch_size,
82
+ )
83
+
84
+ self.num_patches = (self.image_size // self.patch_size) ** 2
85
+ self.num_positions = self.num_patches + 1
86
+
87
+ self.position_embedding = nn.Parameter(
88
+ torch.randn(1, self.num_positions, self.embed_dim)
89
+ )
90
+
91
+ def forward(self, pixel_values: torch.FloatTensor) -> torch.Tensor:
92
+ batch_size = pixel_values.shape[0]
93
+ target_dtype = self.patch_embedding.weight.dtype
94
+ patch_embeds = self.patch_embedding(
95
+ pixel_values
96
+ ) # shape = [*, width, grid, grid]
97
+ patch_embeds = patch_embeds.flatten(2).transpose(1, 2)
98
+ class_embeds = self.class_embedding.expand(batch_size, 1, -1).to(target_dtype)
99
+ embeddings = torch.cat([class_embeds, patch_embeds], dim=1)
100
+ embeddings = embeddings + self.position_embedding.to(target_dtype)
101
+ return embeddings
102
+
103
+
104
+ class InternAttention(nn.Module):
105
+ """Multi-headed attention from 'Attention Is All You Need' paper"""
106
+
107
+ def __init__(self, config: InternVisionConfig):
108
+ super().__init__()
109
+ self.config = config
110
+ self.embed_dim = config.hidden_size
111
+ self.num_heads = config.num_attention_heads
112
+ self.use_flash_attn = config.use_flash_attn and has_flash_attn
113
+ if config.use_flash_attn and not has_flash_attn:
114
+ print(
115
+ "Warning: Flash Attention is not available, use_flash_attn is set to False."
116
+ )
117
+ self.head_dim = self.embed_dim // self.num_heads
118
+ if self.head_dim * self.num_heads != self.embed_dim:
119
+ raise ValueError(
120
+ f"embed_dim must be divisible by num_heads (got `embed_dim`: {self.embed_dim} and `num_heads`:"
121
+ f" {self.num_heads})."
122
+ )
123
+
124
+ self.scale = self.head_dim**-0.5
125
+ self.qkv = nn.Linear(self.embed_dim, 3 * self.embed_dim, bias=config.qkv_bias)
126
+ self.attn_drop = nn.Dropout(config.attention_dropout)
127
+ self.proj_drop = nn.Dropout(config.dropout)
128
+
129
+ self.qk_normalization = config.qk_normalization
130
+
131
+ if self.qk_normalization:
132
+ self.q_norm = InternRMSNorm(self.embed_dim, eps=config.layer_norm_eps)
133
+ self.k_norm = InternRMSNorm(self.embed_dim, eps=config.layer_norm_eps)
134
+
135
+ if self.use_flash_attn:
136
+ self.inner_attn = _attention.apply(attention_dropout=config.attention_dropout)
137
+ self.proj = nn.Linear(self.embed_dim, self.embed_dim)
138
+
139
+ def _naive_attn(self, x):
140
+ B, N, C = x.shape
141
+ qkv = (
142
+ self.qkv(x)
143
+ .reshape(B, N, 3, self.num_heads, C // self.num_heads)
144
+ .permute(2, 0, 3, 1, 4)
145
+ )
146
+ q, k, v = qkv.unbind(0) # make torchscript happy (cannot use tensor as tuple)
147
+
148
+ if self.qk_normalization:
149
+ B_, H_, N_, D_ = q.shape
150
+ q = (
151
+ self.q_norm(q.transpose(1, 2).flatten(-2, -1))
152
+ .view(B_, N_, H_, D_)
153
+ .transpose(1, 2)
154
+ )
155
+ k = (
156
+ self.k_norm(k.transpose(1, 2).flatten(-2, -1))
157
+ .view(B_, N_, H_, D_)
158
+ .transpose(1, 2)
159
+ )
160
+
161
+ attn = (q * self.scale) @ k.transpose(-2, -1)
162
+ attn = attn.softmax(dim=-1)
163
+ attn = self.attn_drop(attn)
164
+
165
+ x = (attn @ v).transpose(1, 2).reshape(B, N, C)
166
+ x = self.proj(x)
167
+ x = self.proj_drop(x)
168
+ return x
169
+
170
+ def _flash_attn(self, x, key_padding_mask=None, need_weights=False):
171
+ qkv = self.qkv(x)
172
+ qkv = rearrange(
173
+ qkv, "b s (three h d) -> b s three h d", three=3, h=self.num_heads
174
+ )
175
+
176
+ if self.qk_normalization:
177
+ q, k, v = qkv.unbind(2)
178
+ q = self.q_norm(q.flatten(-2, -1)).view(q.shape)
179
+ k = self.k_norm(k.flatten(-2, -1)).view(k.shape)
180
+ qkv = torch.stack([q, k, v], dim=2)
181
+
182
+ context, _ = self.inner_attn(
183
+ qkv,
184
+ key_padding_mask=key_padding_mask,
185
+ need_weights=need_weights,
186
+ causal=False,
187
+ )
188
+ outs = self.proj(rearrange(context, "b s h d -> b s (h d)"))
189
+ outs = self.proj_drop(outs)
190
+ return outs
191
+
192
+ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
193
+ x = (
194
+ self._naive_attn(hidden_states)
195
+ if not self.use_flash_attn
196
+ else self._flash_attn(hidden_states)
197
+ )
198
+ return x
199
+
200
+
201
+ class InternMLP(nn.Module):
202
+ def __init__(self, config: InternVisionConfig):
203
+ super().__init__()
204
+ self.config = config
205
+ self.act = ACT2FN[config.hidden_act]
206
+ self.fc1 = nn.Linear(config.hidden_size, config.intermediate_size)
207
+ self.fc2 = nn.Linear(config.intermediate_size, config.hidden_size)
208
+
209
+ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
210
+ hidden_states = self.fc1(hidden_states)
211
+ hidden_states = self.act(hidden_states)
212
+ hidden_states = self.fc2(hidden_states)
213
+ return hidden_states
214
+
215
+
216
+ class InternVisionEncoderLayer(nn.Module):
217
+ def __init__(self, config: InternVisionConfig, drop_path_rate: float):
218
+ super().__init__()
219
+ self.embed_dim = config.hidden_size
220
+ self.intermediate_size = config.intermediate_size
221
+
222
+ self.attn = InternAttention(config)
223
+ self.mlp = InternMLP(config)
224
+ self.norm1 = InternRMSNorm(self.embed_dim, eps=config.layer_norm_eps)
225
+ self.norm2 = InternRMSNorm(self.embed_dim, eps=config.layer_norm_eps)
226
+
227
+ self.ls1 = nn.Parameter(config.initializer_factor * torch.ones(self.embed_dim))
228
+ self.ls2 = nn.Parameter(config.initializer_factor * torch.ones(self.embed_dim))
229
+ self.drop_path1 = (
230
+ DropPath(drop_path_rate) if drop_path_rate > 0.0 else nn.Identity()
231
+ )
232
+ self.drop_path2 = (
233
+ DropPath(drop_path_rate) if drop_path_rate > 0.0 else nn.Identity()
234
+ )
235
+
236
+ def forward(
237
+ self,
238
+ hidden_states: torch.Tensor,
239
+ ) -> Tuple[
240
+ torch.FloatTensor,
241
+ Optional[torch.FloatTensor],
242
+ Optional[Tuple[torch.FloatTensor]],
243
+ ]:
244
+ """
245
+ Args:
246
+ hidden_states (`Tuple[torch.FloatTensor, Optional[torch.FloatTensor]]`): input to the layer of shape `(batch, seq_len, embed_dim)`
247
+ """
248
+ hidden_states = hidden_states + self.drop_path1(
249
+ self.attn(self.norm1(hidden_states)) * self.ls1
250
+ )
251
+
252
+ hidden_states = hidden_states + self.drop_path2(
253
+ self.mlp(self.norm2(hidden_states)) * self.ls2
254
+ )
255
+
256
+ return hidden_states
257
+
258
+
259
+ class InternVisionEncoder(nn.Module):
260
+ """
261
+ Transformer encoder consisting of `config.num_hidden_layers` self attention layers. Each layer is a
262
+ [`InternEncoderLayer`].
263
+
264
+ Args:
265
+ config (`InternConfig`):
266
+ The corresponding vision configuration for the `InternEncoder`.
267
+ """
268
+
269
+ def __init__(self, config: InternVisionConfig):
270
+ super().__init__()
271
+ self.config = config
272
+ # stochastic depth decay rule
273
+ dpr = [
274
+ x.item()
275
+ for x in torch.linspace(0, config.drop_path_rate, config.num_hidden_layers)
276
+ ]
277
+ self.layers = nn.ModuleList(
278
+ [
279
+ InternVisionEncoderLayer(config, dpr[idx])
280
+ for idx in range(config.num_hidden_layers)
281
+ ]
282
+ )
283
+ self.gradient_checkpointing = True
284
+
285
+ def forward(
286
+ self,
287
+ inputs_embeds,
288
+ output_hidden_states: Optional[bool] = None,
289
+ return_dict: Optional[bool] = None,
290
+ ) -> Union[Tuple, BaseModelOutput]:
291
+ r"""
292
+ Args:
293
+ inputs_embeds (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`):
294
+ Embedded representation of the inputs. Should be float, not int tokens.
295
+ output_hidden_states (`bool`, *optional*):
296
+ Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors
297
+ for more detail.
298
+ return_dict (`bool`, *optional*):
299
+ Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.
300
+ """
301
+ output_hidden_states = (
302
+ output_hidden_states
303
+ if output_hidden_states is not None
304
+ else self.config.output_hidden_states
305
+ )
306
+ return_dict = (
307
+ return_dict if return_dict is not None else self.config.use_return_dict
308
+ )
309
+
310
+ encoder_states = () if output_hidden_states else None
311
+ hidden_states = inputs_embeds
312
+
313
+ for idx, encoder_layer in enumerate(self.layers):
314
+ if output_hidden_states:
315
+ encoder_states = encoder_states + (hidden_states,)
316
+ if self.gradient_checkpointing and self.training:
317
+ layer_outputs = torch.utils.checkpoint.checkpoint(
318
+ encoder_layer, hidden_states
319
+ )
320
+ else:
321
+ layer_outputs = encoder_layer(
322
+ hidden_states,
323
+ )
324
+ hidden_states = layer_outputs
325
+
326
+ if output_hidden_states:
327
+ encoder_states = encoder_states + (hidden_states,)
328
+
329
+ if not return_dict:
330
+ return tuple(v for v in [hidden_states, encoder_states] if v is not None)
331
+ return BaseModelOutput(
332
+ last_hidden_state=hidden_states, hidden_states=encoder_states
333
+ )
334
+
335
+
336
+ class InternVisionModel(PreTrainedModel):
337
+ main_input_name = "pixel_values"
338
+ config_class = InternVisionConfig
339
+ _no_split_modules = ["InternVisionEncoderLayer"]
340
+
341
+ def __init__(self, config: InternVisionConfig):
342
+ super().__init__(config)
343
+ self.config = config
344
+
345
+ self.embeddings = InternVisionEmbeddings(config)
346
+ self.encoder = InternVisionEncoder(config)
347
+
348
+ def resize_pos_embeddings(self, old_size, new_size, patch_size):
349
+ pos_emb = self.embeddings.position_embedding
350
+ _, num_positions, embed_dim = pos_emb.shape
351
+ cls_emb = pos_emb[:, :1, :]
352
+ pos_emb = (
353
+ pos_emb[:, 1:, :]
354
+ .reshape(1, old_size // patch_size, old_size // patch_size, -1)
355
+ .permute(0, 3, 1, 2)
356
+ )
357
+ pos_emb = F.interpolate(
358
+ pos_emb.float(),
359
+ size=new_size // patch_size,
360
+ mode="bicubic",
361
+ align_corners=False,
362
+ )
363
+ pos_emb = pos_emb.to(cls_emb.dtype).reshape(1, embed_dim, -1).permute(0, 2, 1)
364
+ pos_emb = torch.cat([cls_emb, pos_emb], dim=1)
365
+ self.embeddings.position_embedding = nn.Parameter(pos_emb)
366
+ logger.info(
367
+ "Resized position embeddings from {} to {}".format(old_size, new_size)
368
+ )
369
+
370
+ def get_input_embeddings(self):
371
+ return self.embeddings
372
+
373
+ def forward(
374
+ self,
375
+ pixel_values: Optional[torch.FloatTensor] = None,
376
+ output_hidden_states: Optional[bool] = None,
377
+ return_dict: Optional[bool] = None,
378
+ pixel_embeds: Optional[torch.FloatTensor] = None,
379
+ ) -> Union[Tuple, BaseModelOutputWithPooling]:
380
+ output_hidden_states = (
381
+ output_hidden_states
382
+ if output_hidden_states is not None
383
+ else self.config.output_hidden_states
384
+ )
385
+ return_dict = (
386
+ return_dict if return_dict is not None else self.config.use_return_dict
387
+ )
388
+
389
+ if pixel_values is None and pixel_embeds is None:
390
+ raise ValueError("You have to specify pixel_values or pixel_embeds")
391
+
392
+ if pixel_embeds is not None:
393
+ hidden_states = pixel_embeds
394
+ else:
395
+ if len(pixel_values.shape) == 4:
396
+ hidden_states = self.embeddings(pixel_values)
397
+ else:
398
+ raise ValueError(f"wrong pixel_values size: {pixel_values.shape}")
399
+ encoder_outputs = self.encoder(
400
+ inputs_embeds=hidden_states,
401
+ output_hidden_states=output_hidden_states,
402
+ return_dict=return_dict,
403
+ )
404
+ last_hidden_state = encoder_outputs.last_hidden_state
405
+ pooled_output = last_hidden_state[:, 0, :]
406
+
407
+ if not return_dict:
408
+ return (last_hidden_state, pooled_output) + encoder_outputs[1:]
409
+
410
+ return BaseModelOutputWithPooling(
411
+ last_hidden_state=last_hidden_state,
412
+ pooler_output=pooled_output,
413
+ hidden_states=encoder_outputs.hidden_states,
414
+ attentions=encoder_outputs.attentions,
415
+ )