Yuan-lab commited on
Commit
4691c43
·
verified ·
1 Parent(s): a3f659e

Upload 7 files

Browse files
Param_Calculation.md ADDED
@@ -0,0 +1,24 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Documentation for Parameter Calculation
2
+
3
+ Thank you for your interest in our Yuan3.0 Model!
4
+ Our community is open to everyone and welcomes all kinds of comments.
5
+ This document is an explanation of Parameter Calculation script.
6
+
7
+ ## Setup for Parameter Calculation
8
+
9
+ ### 1.Download Yuan3.0 Model Locally
10
+
11
+ ### 2.Modify the "MODEL PATH" to Your Local Download Path
12
+
13
+ ```bash
14
+ vim Param_Calculation.py
15
+ # modify the above line
16
+ MODEL_PATH = "/path/to/Yuan3.0-Model"
17
+ ```
18
+
19
+ ### 3.Run Parameter Calculation script
20
+
21
+ ```bash
22
+ python Param_Calculation.py
23
+ ```
24
+
Param_Calculation.py ADDED
@@ -0,0 +1,51 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+ import torch
4
+
5
+ # ================= 配置区域 =================
6
+ # 添加Yuan3.0模型完整路径,将如下路径替换成你的路径
7
+ MODEL_PATH = "/path/to/Yuan3.0-Model"
8
+
9
+ # 将模型目录设为 Python 搜索路径的第一优先级
10
+ if MODEL_PATH not in sys.path:
11
+ sys.path.insert(0, MODEL_PATH)
12
+
13
+ # 设置环境变量,强制离线模式,禁止 HF 联网或访问远程缓存校验
14
+ os.environ["TRANSFORMERS_OFFLINE"] = "1"
15
+ os.environ["HF_DATASETS_OFFLINE"] = "1"
16
+ os.environ["HF_EVALUATE_OFFLINE"] = "1"
17
+
18
+ from transformers import AutoModel, AutoTokenizer, AutoConfig
19
+
20
+ print(f"🚀 开始从本地加载模型:{MODEL_PATH}")
21
+
22
+ # 加载模型
23
+ model = AutoModel.from_pretrained(
24
+ MODEL_PATH,
25
+ torch_dtype=torch.bfloat16,
26
+ low_cpu_mem_usage=True,
27
+ use_flash_attn=False,
28
+ device_map="cpu",
29
+ local_files_only=True,
30
+ trust_remote_code=True,
31
+ )
32
+
33
+ print("\n" + "="*30)
34
+ print("--Yuan3.0 Model Parameter--")
35
+ print("="*30)
36
+
37
+ # 统计参数
38
+ vit_params = 0
39
+ yuan_params = 0
40
+ total_params = model.num_parameters()
41
+ for n, p in model.named_parameters():
42
+ if 'vision_model' in n:
43
+ vit_params += p.numel()
44
+ else:
45
+ yuan_params += p.numel()
46
+
47
+ print(f"Vit Model Parameters: {vit_params:,}")
48
+ print(f"Yuan Model Parameters: {yuan_params:,}")
49
+ print(f"Total Parameters: {total_params:,}")
50
+ print("="*30)
51
+
generation_config.json ADDED
@@ -0,0 +1,7 @@
 
 
 
 
 
 
 
 
1
+ {
2
+ "_from_model_config": true,
3
+ "bos_token_id": 1,
4
+ "eos_token_id": 77185,
5
+ "pad_token_id": 77185,
6
+ "transformers_version": "4.55.2"
7
+ }
model.safetensors.index.json ADDED
The diff for this file is too large to render. See raw diff
 
modeling_intern_vit.py ADDED
@@ -0,0 +1,366 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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,
16
+ BaseModelOutputWithPooling)
17
+ from transformers.modeling_utils import PreTrainedModel
18
+ from transformers.utils import logging
19
+
20
+ from .configuration_intern_vit import InternVisionConfig
21
+ #try:
22
+ from .flash_attention import FlashAttention
23
+ has_flash_attn = True
24
+ #except:
25
+ # print('FlashAttention is not installed.')
26
+ # has_flash_attn = False
27
+
28
+
29
+ logger = logging.get_logger(__name__)
30
+
31
+
32
+ class InternRMSNorm(nn.Module):
33
+ def __init__(self, hidden_size, eps=1e-6):
34
+ super().__init__()
35
+ self.weight = nn.Parameter(torch.ones(hidden_size))
36
+ self.variance_epsilon = eps
37
+
38
+ def forward(self, hidden_states):
39
+ input_dtype = hidden_states.dtype
40
+ hidden_states = hidden_states.to(torch.float32)
41
+ variance = hidden_states.pow(2).mean(-1, keepdim=True)
42
+ hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)
43
+ return self.weight * hidden_states.to(input_dtype)
44
+
45
+
46
+ try:
47
+ from apex.normalization import FusedRMSNorm
48
+
49
+ InternRMSNorm = FusedRMSNorm # noqa
50
+
51
+ logger.info('Discovered apex.normalization.FusedRMSNorm - will use it instead of InternRMSNorm')
52
+ except ImportError:
53
+ # using the normal InternRMSNorm
54
+ pass
55
+ except Exception:
56
+ logger.warning('discovered apex but it failed to load, falling back to InternRMSNorm')
57
+ pass
58
+
59
+
60
+ NORM2FN = {
61
+ 'rms_norm': InternRMSNorm,
62
+ 'layer_norm': nn.LayerNorm,
63
+ }
64
+
65
+
66
+ class InternVisionEmbeddings(nn.Module):
67
+ def __init__(self, config: InternVisionConfig):
68
+ super().__init__()
69
+ self.config = config
70
+ self.embed_dim = config.hidden_size
71
+ self.image_size = config.image_size
72
+ self.patch_size = config.patch_size
73
+
74
+ self.class_embedding = nn.Parameter(
75
+ torch.randn(1, 1, self.embed_dim),
76
+ )
77
+
78
+ self.patch_embedding = nn.Conv2d(
79
+ in_channels=3, out_channels=self.embed_dim, kernel_size=self.patch_size, stride=self.patch_size
80
+ )
81
+
82
+ self.num_patches = (self.image_size // self.patch_size) ** 2
83
+ self.num_positions = self.num_patches + 1
84
+
85
+ self.position_embedding = nn.Parameter(torch.randn(1, self.num_positions, self.embed_dim))
86
+
87
+ def _get_pos_embed(self, pos_embed, H, W):
88
+ target_dtype = pos_embed.dtype
89
+ pos_embed = pos_embed.float().reshape(
90
+ 1, self.image_size // self.patch_size, self.image_size // self.patch_size, -1).permute(0, 3, 1, 2)
91
+ pos_embed = F.interpolate(pos_embed, size=(H, W), mode='bicubic', align_corners=False).\
92
+ reshape(1, -1, H * W).permute(0, 2, 1).to(target_dtype)
93
+ return pos_embed
94
+
95
+ def forward(self, pixel_values: torch.FloatTensor) -> torch.Tensor:
96
+ target_dtype = self.patch_embedding.weight.dtype
97
+ patch_embeds = self.patch_embedding(pixel_values) # shape = [*, channel, width, height]
98
+ batch_size, _, height, width = patch_embeds.shape
99
+ patch_embeds = patch_embeds.flatten(2).transpose(1, 2)
100
+ class_embeds = self.class_embedding.expand(batch_size, 1, -1).to(target_dtype)
101
+ embeddings = torch.cat([class_embeds, patch_embeds], dim=1)
102
+ position_embedding = torch.cat([
103
+ self.position_embedding[:, :1, :],
104
+ self._get_pos_embed(self.position_embedding[:, 1:, :], height, width)
105
+ ], dim=1)
106
+ embeddings = embeddings + position_embedding.to(target_dtype)
107
+ return embeddings
108
+
109
+
110
+ class InternAttention(nn.Module):
111
+ """Multi-headed attention from 'Attention Is All You Need' paper"""
112
+
113
+ def __init__(self, config: InternVisionConfig):
114
+ super().__init__()
115
+ self.config = config
116
+ self.embed_dim = config.hidden_size
117
+ self.num_heads = config.num_attention_heads
118
+ self.use_flash_attn = config.use_flash_attn and has_flash_attn
119
+ self.use_flash_attn = True # modify
120
+ if config.use_flash_attn and not has_flash_attn:
121
+ print('Warning: Flash Attention is not available, use_flash_attn is set to False.')
122
+ self.head_dim = self.embed_dim // self.num_heads
123
+ if self.head_dim * self.num_heads != self.embed_dim:
124
+ raise ValueError(
125
+ f'embed_dim must be divisible by num_heads (got `embed_dim`: {self.embed_dim} and `num_heads`:'
126
+ f' {self.num_heads}).'
127
+ )
128
+
129
+ self.scale = self.head_dim ** -0.5
130
+ self.qkv = nn.Linear(self.embed_dim, 3 * self.embed_dim, bias=config.qkv_bias)
131
+ self.attn_drop = nn.Dropout(config.attention_dropout)
132
+ self.proj_drop = nn.Dropout(config.dropout)
133
+
134
+ self.qk_normalization = config.qk_normalization
135
+
136
+ if self.qk_normalization:
137
+ self.q_norm = InternRMSNorm(self.embed_dim, eps=config.layer_norm_eps)
138
+ self.k_norm = InternRMSNorm(self.embed_dim, eps=config.layer_norm_eps)
139
+
140
+ if self.use_flash_attn:
141
+ self.inner_attn = FlashAttention(attention_dropout=config.attention_dropout)
142
+ self.proj = nn.Linear(self.embed_dim, self.embed_dim)
143
+
144
+ def _naive_attn(self, x):
145
+ B, N, C = x.shape
146
+ qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)
147
+ q, k, v = qkv.unbind(0) # make torchscript happy (cannot use tensor as tuple)
148
+
149
+ if self.qk_normalization:
150
+ B_, H_, N_, D_ = q.shape
151
+ q = self.q_norm(q.transpose(1, 2).flatten(-2, -1)).view(B_, N_, H_, D_).transpose(1, 2)
152
+ k = self.k_norm(k.transpose(1, 2).flatten(-2, -1)).view(B_, N_, H_, D_).transpose(1, 2)
153
+
154
+ attn = ((q * self.scale) @ k.transpose(-2, -1))
155
+ attn = attn.softmax(dim=-1)
156
+ attn = self.attn_drop(attn)
157
+
158
+ x = (attn @ v).transpose(1, 2).reshape(B, N, C)
159
+ x = self.proj(x)
160
+ x = self.proj_drop(x)
161
+ return x
162
+
163
+ def _flash_attn(self, x, key_padding_mask=None, need_weights=False):
164
+ qkv = self.qkv(x)
165
+ qkv = rearrange(qkv, 'b s (three h d) -> b s three h d', three=3, h=self.num_heads)
166
+
167
+ if self.qk_normalization:
168
+ q, k, v = qkv.unbind(2)
169
+ q = self.q_norm(q.flatten(-2, -1)).view(q.shape)
170
+ k = self.k_norm(k.flatten(-2, -1)).view(k.shape)
171
+ qkv = torch.stack([q, k, v], dim=2)
172
+
173
+ context, _ = self.inner_attn(
174
+ qkv, key_padding_mask=key_padding_mask, need_weights=need_weights, causal=False
175
+ )
176
+ outs = self.proj(rearrange(context, 'b s h d -> b s (h d)'))
177
+ outs = self.proj_drop(outs)
178
+ return outs
179
+
180
+ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
181
+ x = self._naive_attn(hidden_states) if not self.use_flash_attn else self._flash_attn(hidden_states)
182
+ return x
183
+
184
+
185
+ class InternMLP(nn.Module):
186
+ def __init__(self, config: InternVisionConfig):
187
+ super().__init__()
188
+ self.config = config
189
+ self.act = ACT2FN[config.hidden_act]
190
+ self.fc1 = nn.Linear(config.hidden_size, config.intermediate_size)
191
+ self.fc2 = nn.Linear(config.intermediate_size, config.hidden_size)
192
+
193
+ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
194
+ hidden_states = self.fc1(hidden_states)
195
+ hidden_states = self.act(hidden_states)
196
+ hidden_states = self.fc2(hidden_states)
197
+ return hidden_states
198
+
199
+
200
+ class InternVisionEncoderLayer(nn.Module):
201
+ def __init__(self, config: InternVisionConfig, drop_path_rate: float):
202
+ super().__init__()
203
+ self.embed_dim = config.hidden_size
204
+ self.intermediate_size = config.intermediate_size
205
+ self.norm_type = config.norm_type
206
+
207
+ self.attn = InternAttention(config)
208
+ self.mlp = InternMLP(config)
209
+ self.norm1 = NORM2FN[self.norm_type](self.embed_dim, eps=config.layer_norm_eps)
210
+ self.norm2 = NORM2FN[self.norm_type](self.embed_dim, eps=config.layer_norm_eps)
211
+
212
+ self.ls1 = nn.Parameter(config.initializer_factor * torch.ones(self.embed_dim))
213
+ self.ls2 = nn.Parameter(config.initializer_factor * torch.ones(self.embed_dim))
214
+ self.drop_path1 = DropPath(drop_path_rate) if drop_path_rate > 0. else nn.Identity()
215
+ self.drop_path2 = DropPath(drop_path_rate) if drop_path_rate > 0. else nn.Identity()
216
+
217
+ def forward(
218
+ self,
219
+ hidden_states: torch.Tensor,
220
+ ) -> Tuple[torch.FloatTensor, Optional[torch.FloatTensor], Optional[Tuple[torch.FloatTensor]]]:
221
+ """
222
+ Args:
223
+ hidden_states (`Tuple[torch.FloatTensor, Optional[torch.FloatTensor]]`): input to the layer of shape `(batch, seq_len, embed_dim)`
224
+ """
225
+
226
+ hidden_states = hidden_states + self.drop_path1(self.attn(self.norm1(hidden_states)) * self.ls1)
227
+
228
+ hidden_states = hidden_states + self.drop_path2(self.mlp(self.norm2(hidden_states)) * self.ls2)
229
+
230
+ return hidden_states
231
+
232
+
233
+ class InternVisionEncoder(nn.Module):
234
+ """
235
+ Transformer encoder consisting of `config.num_hidden_layers` self attention layers. Each layer is a
236
+ [`InternEncoderLayer`].
237
+
238
+ Args:
239
+ config (`InternConfig`):
240
+ The corresponding vision configuration for the `InternEncoder`.
241
+ """
242
+
243
+ def __init__(self, config: InternVisionConfig):
244
+ super().__init__()
245
+ self.config = config
246
+ # stochastic depth decay rule
247
+ dpr = [x.item() for x in torch.linspace(0, config.drop_path_rate, config.num_hidden_layers)]
248
+ self.layers = nn.ModuleList([
249
+ InternVisionEncoderLayer(config, dpr[idx]) for idx in range(config.num_hidden_layers)])
250
+ self.gradient_checkpointing = True
251
+
252
+ def forward(
253
+ self,
254
+ inputs_embeds,
255
+ output_hidden_states: Optional[bool] = None,
256
+ return_dict: Optional[bool] = None,
257
+ ) -> Union[Tuple, BaseModelOutput]:
258
+ r"""
259
+ Args:
260
+ inputs_embeds (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`):
261
+ Embedded representation of the inputs. Should be float, not int tokens.
262
+ output_hidden_states (`bool`, *optional*):
263
+ Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors
264
+ for more detail.
265
+ return_dict (`bool`, *optional*):
266
+ Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.
267
+ """
268
+ output_hidden_states = (
269
+ output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
270
+ )
271
+ return_dict = return_dict if return_dict is not None else self.config.use_return_dict
272
+
273
+ encoder_states = () if output_hidden_states else None
274
+ hidden_states = inputs_embeds
275
+
276
+ for idx, encoder_layer in enumerate(self.layers):
277
+ if output_hidden_states:
278
+ encoder_states = encoder_states + (hidden_states,)
279
+ if self.gradient_checkpointing and self.training:
280
+ layer_outputs = torch.utils.checkpoint.checkpoint(
281
+ encoder_layer,
282
+ hidden_states)
283
+ else:
284
+ layer_outputs = encoder_layer(
285
+ hidden_states,
286
+ )
287
+ hidden_states = layer_outputs
288
+ #import pdb
289
+ #pdb.set_trace()
290
+
291
+ if output_hidden_states:
292
+ encoder_states = encoder_states + (hidden_states,)
293
+
294
+ if not return_dict:
295
+ return tuple(v for v in [hidden_states, encoder_states] if v is not None)
296
+ return BaseModelOutput(
297
+ last_hidden_state=hidden_states, hidden_states=encoder_states
298
+ )
299
+
300
+
301
+ class InternVisionModel(PreTrainedModel):
302
+ main_input_name = 'pixel_values'
303
+ config_class = InternVisionConfig
304
+ _no_split_modules = ['InternVisionEncoderLayer']
305
+
306
+ def __init__(self, config: InternVisionConfig):
307
+ super().__init__(config)
308
+ self.config = config
309
+
310
+ self.embeddings = InternVisionEmbeddings(config)
311
+ self.encoder = InternVisionEncoder(config)
312
+
313
+ def resize_pos_embeddings(self, old_size, new_size, patch_size):
314
+ pos_emb = self.embeddings.position_embedding
315
+ _, num_positions, embed_dim = pos_emb.shape
316
+ cls_emb = pos_emb[:, :1, :]
317
+ pos_emb = pos_emb[:, 1:, :].reshape(1, old_size // patch_size, old_size // patch_size, -1).permute(0, 3, 1, 2)
318
+ pos_emb = F.interpolate(pos_emb.float(), size=new_size // patch_size, mode='bicubic', align_corners=False)
319
+ pos_emb = pos_emb.to(cls_emb.dtype).reshape(1, embed_dim, -1).permute(0, 2, 1)
320
+ pos_emb = torch.cat([cls_emb, pos_emb], dim=1)
321
+ self.embeddings.position_embedding = nn.Parameter(pos_emb)
322
+ self.embeddings.image_size = new_size
323
+ logger.info('Resized position embeddings from {} to {}'.format(old_size, new_size))
324
+
325
+ def get_input_embeddings(self):
326
+ return self.embeddings
327
+
328
+ def forward(
329
+ self,
330
+ pixel_values: Optional[torch.FloatTensor] = None,
331
+ output_hidden_states: Optional[bool] = None,
332
+ return_dict: Optional[bool] = None,
333
+ pixel_embeds: Optional[torch.FloatTensor] = None,
334
+ ) -> Union[Tuple, BaseModelOutputWithPooling]:
335
+ output_hidden_states = (
336
+ output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
337
+ )
338
+ return_dict = return_dict if return_dict is not None else self.config.use_return_dict
339
+
340
+ if pixel_values is None and pixel_embeds is None:
341
+ raise ValueError('You have to specify pixel_values or pixel_embeds')
342
+
343
+ if pixel_embeds is not None:
344
+ hidden_states = pixel_embeds
345
+ else:
346
+ if len(pixel_values.shape) == 4:
347
+ hidden_states = self.embeddings(pixel_values)
348
+ else:
349
+ raise ValueError(f'wrong pixel_values size: {pixel_values.shape}')
350
+ encoder_outputs = self.encoder(
351
+ inputs_embeds=hidden_states,
352
+ output_hidden_states=output_hidden_states,
353
+ return_dict=return_dict,
354
+ )
355
+ last_hidden_state = encoder_outputs.last_hidden_state
356
+ pooled_output = last_hidden_state[:, 0, :]
357
+
358
+ if not return_dict:
359
+ return (last_hidden_state, pooled_output) + encoder_outputs[1:]
360
+
361
+ return BaseModelOutputWithPooling(
362
+ last_hidden_state=last_hidden_state,
363
+ pooler_output=pooled_output,
364
+ hidden_states=encoder_outputs.hidden_states,
365
+ attentions=encoder_outputs.attentions,
366
+ )
modeling_yuanlm2.py ADDED
@@ -0,0 +1,1548 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # coding=utf-8
2
+ # Copyright 2022 EleutherAI and the HuggingFace Inc. team. All rights reserved.
3
+ #
4
+ # This code is based on EleutherAI's GPT-NeoX library and the GPT-NeoX
5
+ # and OPT implementations in this library. It has been modified from its
6
+ # original forms to accommodate minor architectural differences compared
7
+ # to GPT-NeoX and OPT used by the Meta AI team that trained the model.
8
+ #
9
+ # Licensed under the Apache License, Version 2.0 (the "License");
10
+ # you may not use this file except in compliance with the License.
11
+ # You may obtain a copy of the License at
12
+ #
13
+ # http://www.apache.org/licenses/LICENSE-2.0
14
+ #
15
+ # Unless required by applicable law or agreed to in writing, software
16
+ # distributed under the License is distributed on an "AS IS" BASIS,
17
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
18
+ # See the License for the specific language governing permissions and
19
+ # limitations under the License.
20
+ """ PyTorch Yuan model."""
21
+ import math
22
+ from typing import List, Optional, Tuple, Union
23
+ import torch.nn.functional as F
24
+ import torch
25
+ import torch.utils.checkpoint
26
+ from torch import einsum, nn
27
+ from torch.nn import BCEWithLogitsLoss, CrossEntropyLoss, MSELoss
28
+ from transformers.activations import ACT2FN
29
+ from transformers.generation import GenerationMixin
30
+ from transformers.modeling_outputs import BaseModelOutputWithPast, CausalLMOutputWithPast, SequenceClassifierOutputWithPast
31
+ from transformers.modeling_utils import PreTrainedModel
32
+ from transformers.utils import add_start_docstrings, add_start_docstrings_to_model_forward, logging, replace_return_docstrings
33
+ from .configuration_yuan import YuanConfig
34
+ from einops import rearrange
35
+ # from flash_attn import flash_attn_varlen_func as flash_attn_unpadded_func
36
+ #from apex.normalization import MixedFusedRMSNorm as RMSNorm
37
+ #from flash_attn import flash_attn_func
38
+ #from transformer_engine.pytorch import RMSNorm
39
+ import pdb
40
+ import copy
41
+ try:
42
+ import grouped_gemm as gg
43
+ except ImportError:
44
+ gg = None
45
+ try:
46
+ from flash_attn import flash_attn_varlen_func as flash_attn_unpadded_func
47
+ from flash_attn import flash_attn_func
48
+ except ImportError:
49
+ flash_attn_unpadded_func = None
50
+
51
+
52
+ logger = logging.get_logger(__name__)
53
+
54
+ _CONFIG_FOR_DOC = "YuanConfig"
55
+
56
+ class RMSNorm(torch.nn.Module):
57
+ def __init__(self, hidden_size, eps=1e-6):
58
+ super().__init__()
59
+ self.weight = torch.nn.Parameter(torch.ones(hidden_size))
60
+ self.variance_epsilon = eps
61
+
62
+ def forward(self, hidden_states):
63
+ variance = hidden_states.to(torch.float32).pow(2).mean(-1, keepdim=True)
64
+ hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)
65
+
66
+ # convert into half-precision if necessary
67
+ if self.weight.dtype in [torch.float16, torch.bfloat16]:
68
+ hidden_states = hidden_states.to(self.weight.dtype)
69
+
70
+ return self.weight * hidden_states
71
+
72
+
73
+ class YuanRotaryEmbedding(nn.Module):
74
+ def __init__(self, dim, base=10000, dtype=torch.float32, rotary_interleaved=False, seq_len_interpolation_factor=None):
75
+ super().__init__()
76
+ self.base = base
77
+ self.dim = dim
78
+ self.rotary_interleaved = rotary_interleaved
79
+ self.seq_len_interpolation_factor = seq_len_interpolation_factor
80
+
81
+ def get_rotary_seq_len(
82
+ self,
83
+ inference_param=None,
84
+ transformer_input: torch.Tensor=None,
85
+ transformer_config=None,
86
+ ):
87
+ if inference_param is not None:
88
+ rotary_seq_len = inference_param.max_sequence_length
89
+ else:
90
+ rotary_seq_len = transformer_input.size[0]
91
+ if transformer_config.sequence_parallel:
92
+ rotary_seq_len *= transformer_config.tensor_model_parallel_size
93
+
94
+ return rotary_seq_len
95
+
96
+ def forward(self, max_seq_len, offset=0):
97
+
98
+ """Forward pass of RoPE embedding.
99
+
100
+ Args:
101
+ max_seq_len (int): Maximum size of sequence
102
+ offset (int, optional): _description_. Defaults to 0.
103
+
104
+ Returns:
105
+ Tensor: Embeddings after applying RoPE.
106
+ """
107
+ inv_freq = (1.0 / ( self.base**(torch.arange(0, self.dim, 2, dtype=torch.float32, device=torch.cuda.current_device()) / self.dim))).to(torch.float32)
108
+
109
+ #max_seq_len_int = max_seq_len.item() if max_seq_len.numel() == 1 else max_seq_len.max().item()
110
+ seq = (
111
+ torch.arange(max_seq_len, device=inv_freq.device, dtype=inv_freq.dtype)
112
+ + offset
113
+ )
114
+
115
+ if self.seq_len_interpolation_factor is not None:
116
+ seq *= 1 / self.seq_len_interpolation_factor
117
+
118
+ freqs = torch.outer(seq, inv_freq)
119
+ # first part even vector components, second part odd vector components,
120
+ # 2 * dim in dimension size
121
+ if not self.rotary_interleaved:
122
+ emb = torch.cat((freqs, freqs), dim=-1)
123
+ else:
124
+ emb = torch.stack((freqs.view(-1, 1), freqs.view(-1, 1)), dim=-1).view(
125
+ freqs.shape[0], -1
126
+ )
127
+ # emb [seq_length, .., dim]
128
+ emb = emb[:, None, None, :]
129
+ #emb = emb[:, None, :]
130
+ return emb
131
+
132
+
133
+ def _rotate_half(x, rotary_interleaved):
134
+ """huggingface version
135
+ change sign so the last dimension becomes [-odd, +even]
136
+
137
+ x1, x2 = torch.chunk(x, 2, dim=-1)
138
+ return torch.cat((-x2, x1), dim=-1)
139
+ """
140
+ if not rotary_interleaved:
141
+ x1, x2 = torch.chunk(x, 2, dim=-1)
142
+ return torch.cat((-x2, x1), dim=-1)
143
+ else:
144
+ x1 = x[:, :, :, ::2]
145
+ x2 = x[:, :, :, 1::2]
146
+ x_new = torch.stack((-x2, x1), dim=-1)
147
+ return x_new.view(x_new.shape[0], x_new.shape[1], x_new.shape[2], -1)
148
+
149
+ def apply_rotary_pos_emb(t, freqs, position_ids, rotary_interleaved=False):
150
+
151
+ rot_dim = freqs.shape[-1]
152
+ #if position_ids.shape[1] > 1:
153
+ freqs = freqs[position_ids]
154
+ freqs = freqs.view(t.shape[1],freqs.shape[1],freqs.shape[2],freqs.shape[4]).transpose(0,1)
155
+ # ideally t_pass is empty so rotary pos embedding is applied to all tensor t
156
+ t, t_pass = t[..., :rot_dim], t[..., rot_dim:]
157
+
158
+ # first part is cosine component
159
+ # second part is sine component, need to change signs with _rotate_half method
160
+ t_type = t.dtype
161
+ cos_ = torch.cos(freqs).to(t.dtype)
162
+ sin_ = torch.sin(freqs).to(t.dtype)
163
+
164
+ t = (t * cos_) + (_rotate_half(t, rotary_interleaved) * sin_)
165
+ return torch.cat((t, t_pass), dim=-1)
166
+
167
+ return torch.cat((t, t_pass), dim=-1)
168
+
169
+ class LocalizedFiltering(torch.nn.Module):
170
+ """
171
+ Mega's Exponential Moving Average layer, largely left unmodified from the original repo with the exception of
172
+ variable names and moving away from the stateful representation of incremental decoding state. See
173
+ "https://arxiv.org/abs/2209.10655" for more details.
174
+ """
175
+
176
+ def __init__(self, hidden_size, lf_conv2d_group, lf_conv2d_num_pad, use_lfa_bias):
177
+ super().__init__()
178
+
179
+ self.embed_dim = hidden_size
180
+ self.lf_conv2d_group = lf_conv2d_group
181
+ self.lf_conv2d_num_pad = lf_conv2d_num_pad
182
+ self.use_lfa_bias = use_lfa_bias
183
+ if self.lf_conv2d_num_pad == 1:
184
+ self.training = True
185
+ self.conv1 = torch.nn.Conv2d(self.embed_dim, self.embed_dim // 2, (2, 1), stride=(1, 1), padding=(self.lf_conv2d_num_pad, 0), groups=self.lf_conv2d_group, bias=use_lfa_bias)
186
+ self.conv2 = torch.nn.Conv2d(self.embed_dim // 2, self.embed_dim, (2, 1), stride=(1, 1), padding=(self.lf_conv2d_num_pad, 0), groups=self.lf_conv2d_group, bias=use_lfa_bias)
187
+ self.output_layernorm = RMSNorm(self.embed_dim, eps=1e-6)
188
+
189
+ def _train_forward(self, inputs):
190
+ inputs = inputs.transpose(0,1)
191
+ seq_len, bsz, embed_dim = inputs.size()
192
+ if embed_dim != self.embed_dim:
193
+ raise ValueError(
194
+ f"Unexpected embedding dimension received: input is {embed_dim}, model expects {self.embed_dim}"
195
+ )
196
+ residual = inputs
197
+
198
+ inputs = inputs.view(seq_len, 1, bsz, embed_dim).permute(2, 3, 0, 1)
199
+ output1 = self.conv1(inputs)
200
+ output1 = output1[:, :, :seq_len, :]
201
+
202
+ output2 = self.conv2(output1)
203
+ output2 = output2[:, :, :seq_len, :].permute(2, 3, 0, 1).contiguous()
204
+ output2 = output2.view(seq_len, bsz, embed_dim)
205
+ assert output2.shape == residual.shape
206
+
207
+ lf_output = self.output_layernorm(output2 + residual)
208
+ lf_output = lf_output.transpose(0,1)
209
+ return lf_output
210
+
211
+ def _inference_forward(self, inputs, before_hidden_states):
212
+
213
+ if before_hidden_states is None:
214
+ residual = inputs
215
+ seq_len, bsz, embed_dim = inputs.size()
216
+
217
+ inputs = inputs.view(seq_len, 1, bsz, embed_dim).permute(2, 3, 0, 1)
218
+
219
+ pad_zero1 = torch.zeros(bsz, embed_dim, 1, 1).to(inputs)
220
+ inputs = torch.cat((pad_zero1, inputs), dim=2).contiguous()
221
+ output1 = self.conv1(inputs)
222
+
223
+ pad_zero2 = torch.zeros(bsz, embed_dim // 2, 1, 1).to(output1)
224
+ output1 = torch.cat((pad_zero2, output1), dim=2).contiguous()
225
+ output2 = self.conv2(output1)
226
+
227
+ output2 = output2.permute(2, 3, 0, 1).contiguous()
228
+
229
+ output2 = output2.view(seq_len, bsz, embed_dim)
230
+
231
+ assert output2.shape == residual.shape
232
+
233
+ lf_output = self.output_layernorm(output2 + residual)
234
+
235
+ else:
236
+ residual = inputs
237
+
238
+ seq_len, bsz, embed_dim = inputs.size()
239
+ seq_len_before, _, _ = before_hidden_states.size()
240
+
241
+ assert seq_len == 1 and seq_len_before == 2
242
+
243
+ inputs = torch.cat((before_hidden_states, inputs), dim=0)
244
+ inputs = inputs.view(3, 1, bsz, embed_dim).permute(2, 3, 0, 1)
245
+
246
+ output1 = self.conv1(inputs)
247
+ output2 = self.conv2(output1)
248
+ output2 = output2.view(1, bsz, embed_dim)
249
+
250
+ assert output2.shape == residual.shape
251
+
252
+ lf_output = self.output_layernorm(output2 + residual)
253
+
254
+ return lf_output
255
+
256
+ def forward(
257
+ self,
258
+ inputs,
259
+ before_hidden_states = None,
260
+ ) -> torch.Tensor:
261
+ # assert self.lf_conv2d_num_pad == 1
262
+ if self.training:
263
+ lf_output = self._train_forward(inputs)
264
+ else:
265
+ lf_output = self._inference_forward(inputs, before_hidden_states)
266
+
267
+ return lf_output
268
+
269
+
270
+ # Copied from transformers.models.bart.modeling_bart._make_causal_mask
271
+ def _make_causal_mask(
272
+ input_ids_shape: torch.Size, dtype: torch.dtype, device: torch.device, past_key_values_length: int = 0
273
+ ):
274
+ """
275
+ Make causal mask used for bi-directional self-attention.
276
+ """
277
+ bsz, tgt_len = input_ids_shape
278
+ mask = torch.full((tgt_len, tgt_len), torch.tensor(torch.finfo(dtype).min, device=device), device=device)
279
+ mask_cond = torch.arange(mask.size(-1), device=device)
280
+ mask.masked_fill_(mask_cond < (mask_cond + 1).view(mask.size(-1), 1), 0)
281
+ mask = mask.to(dtype)
282
+
283
+ if past_key_values_length > 0:
284
+ mask = torch.cat([torch.zeros(tgt_len, past_key_values_length, dtype=dtype, device=device), mask], dim=-1)
285
+ return mask[None, None, :, :].expand(bsz, 1, tgt_len, tgt_len + past_key_values_length)
286
+
287
+
288
+ # Copied from transformers.models.bart.modeling_bart._expand_mask
289
+ def _expand_mask(mask: torch.Tensor, dtype: torch.dtype, tgt_len: Optional[int] = None):
290
+ """
291
+ Expands attention_mask from `[bsz, seq_len]` to `[bsz, 1, tgt_seq_len, src_seq_len]`.
292
+ """
293
+ bsz, src_len = mask.size()
294
+ tgt_len = tgt_len if tgt_len is not None else src_len
295
+
296
+ expanded_mask = mask[:, None, None, :].expand(bsz, 1, tgt_len, src_len).to(dtype)
297
+
298
+ inverted_mask = 1.0 - expanded_mask
299
+
300
+ return inverted_mask.masked_fill(inverted_mask.to(torch.bool), torch.finfo(dtype).min)
301
+
302
+
303
+ class YuanRMSNorm(nn.Module):
304
+ def __init__(self, hidden_size, eps=1e-6):
305
+ """
306
+ YuanRMSNorm is equivalent to LlamaRMSNorm
307
+ """
308
+ super().__init__()
309
+ self.weight = nn.Parameter(torch.ones(hidden_size))
310
+ self.variance_epsilon = eps
311
+
312
+ def forward(self, hidden_states):
313
+ input_dtype = hidden_states.dtype
314
+ hidden_states = hidden_states.to(torch.float32)
315
+ variance = hidden_states.pow(2).mean(-1, keepdim=True)
316
+ hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)
317
+ return self.weight * hidden_states.to(input_dtype)
318
+
319
+ # flash attn
320
+ class FlashSelfAttention(torch.nn.Module):
321
+ """Implement the scaled dot product attention with softmax.
322
+ Arguments
323
+ ---------
324
+ softmax_scale: The temperature to use for the softmax attention.
325
+ (default: 1/sqrt(d_keys) where d_keys is computed at
326
+ runtime)
327
+ attention_dropout: The dropout rate to apply to the attention
328
+ (default: 0.0)
329
+ """
330
+ def __init__(self, causal=False, softmax_scale=None, attention_dropout=0.0,
331
+ device=None, dtype=None):
332
+ super().__init__()
333
+ assert flash_attn_unpadded_func is not None, ('Please install FlashAttention first, '
334
+ 'e.g., with pip install flash-attn')
335
+ assert rearrange is not None, 'Please install einops first, e.g., with pip install einops'
336
+ self.causal = causal
337
+ self.softmax_scale = softmax_scale
338
+ self.dropout_p = attention_dropout
339
+
340
+ def forward(self, q, k, v):
341
+ """Implements the multihead softmax attention.
342
+ Arguments
343
+ ---------
344
+ q, k, v: The tensor containing the query, key, and value. (B, S, H, D)
345
+ """
346
+
347
+ assert all((i.dtype in [torch.float16, torch.bfloat16] for i in (q,k,v)))
348
+ assert all((i.is_cuda for i in (q,k,v)))
349
+
350
+ batch_size, seqlen_q = q.shape[1], q.shape[0]
351
+ seqlen_k = k.shape[0]
352
+ q, k, v = [rearrange(x, 'b s ... -> (b s) ...') for x in [q, k, v]]
353
+ cu_seqlens_q = torch.arange(0, (batch_size + 1) * seqlen_q, step=seqlen_q, dtype=torch.int32, device=q.device)
354
+ if self.training:
355
+ # during training q,k,v always have same seqlen
356
+ assert seqlen_k == seqlen_q
357
+ is_causal = self.causal
358
+ cu_seqlens_k = cu_seqlens_q
359
+ dropout_p = self.dropout_p
360
+ else:
361
+ # turn off FA causal mask after first inference autoregressive iteration
362
+ # only on first autoregressive step q,k,v have same seqlen
363
+ is_causal = seqlen_q == seqlen_k
364
+ cu_seqlens_k = torch.arange(0, (batch_size + 1) * seqlen_k, step=seqlen_k, dtype=torch.int32, device=q.device)
365
+ #cu_seqlens_q = [cu_seqlens_q[0], cu_seqlens_q[-1]]
366
+ #cu_seqlens_k = [cu_seqlens_k[0], cu_seqlens_k[-1]]
367
+ dropout_p = 0
368
+
369
+ output = flash_attn_unpadded_func(q, k, v, cu_seqlens_q, cu_seqlens_k, seqlen_q, seqlen_k, dropout_p, softmax_scale=self.softmax_scale, causal=is_causal)
370
+
371
+ output = rearrange(output, '(b s) ... -> b s ...', b=batch_size)
372
+ return output
373
+
374
+ class ParallelAttention_router(nn.Module):
375
+ def __init__(self, config, num_experts):
376
+ super(ParallelAttention_router, self).__init__()
377
+ layer_number=0
378
+ self.layer_number = max(1, layer_number)
379
+
380
+ self.hidden_size = config.hidden_size
381
+ self.projection_size = num_experts
382
+
383
+ self.num_attention_router_heads = config.moe_config['num_attention_router_heads']
384
+ self.hidden_size_per_attention_head = config.max_position_embeddings // self.num_attention_router_heads
385
+ self.query_key_value = nn.Linear(self.hidden_size, self.projection_size*3, bias=False)
386
+
387
+ def forward(self, hidden_states, attention_mask=None, enc_position_ids=None,
388
+ encoder_output=None, inference_params=None,
389
+ rotary_pos_emb=None):
390
+ is_first_step = False
391
+ before_hidden_states = None
392
+
393
+ #mixed_x_layer = torch.matmul(hidden_states, self.query_key_value)
394
+ mixed_x_layer = self.query_key_value(hidden_states)
395
+ (query_layer, key_layer, value_layer) = torch.split(mixed_x_layer, self.projection_size, -1)
396
+ b, s, z = query_layer.shape
397
+
398
+ # use fp32 router
399
+ query_layer = query_layer.float().view(b,s,z,1)
400
+ key_layer = key_layer.float().view(b,s,z,1)
401
+ value_layer = value_layer.float().view(b,s,z,1)
402
+
403
+ attn_weights = torch.matmul(query_layer, key_layer.transpose(2, 3))
404
+ attn_weights = nn.functional.softmax(attn_weights, dim=-1)
405
+ attn_output = torch.matmul(attn_weights, value_layer)
406
+ router_output = attn_output.view(-1, z)
407
+ return router_output
408
+
409
+ class YuanExpertMLP(nn.Module):
410
+ def __init__(self, config):
411
+ super(YuanExpertMLP, self).__init__()
412
+ self.gated_linear_unit = config.moe_config['gated_linear_unit']
413
+ #self.ffn_hidden_size = config.moe_config['ffn_hidden_size']
414
+ self.ffn_hidden_size = config.ffn_hidden_size
415
+
416
+
417
+ if self.gated_linear_unit:
418
+ self.w1 = nn.Linear(config.hidden_size, self.ffn_hidden_size*2, bias=False)
419
+
420
+ else:
421
+ self.w1 = nn.Linear(config.hidden_size, self.ffn_hidden_size, bias=False)
422
+
423
+ self.act_fn = ACT2FN[config.hidden_act]
424
+ self.w2 = nn.Linear(self.ffn_hidden_size, config.hidden_size, bias=False)
425
+
426
+
427
+ def forward(self, x):
428
+ x = self.w1(x)
429
+ if self.gated_linear_unit:
430
+ x = torch.chunk(x, 2, dim=-1)
431
+ x = self.act_fn(x[0]) * x[1]
432
+ else:
433
+ x = self.act_fn(x)
434
+ x = self.w2(x)
435
+ return x
436
+
437
+
438
+
439
+ class YuanMLP(nn.Module):
440
+ def __init__(
441
+ self,
442
+ hidden_size: int,
443
+ intermediate_size: int,
444
+ hidden_act: str
445
+ ):
446
+ super().__init__()
447
+ self.up_proj = nn.Linear(hidden_size, intermediate_size, bias=False)
448
+ self.gate_proj = nn.Linear(hidden_size, intermediate_size, bias=False)
449
+ self.down_proj = nn.Linear(intermediate_size, hidden_size, bias=False)
450
+ self.act_fn = ACT2FN[hidden_act]
451
+
452
+ def forward(self, x):
453
+ return self.down_proj(self.gate_proj(x) * self.act_fn(self.up_proj(x)))
454
+
455
+
456
+ class YuanAttention(nn.Module):
457
+ """Localized Filtering-based Attention 'YUAN 2.0: A Large Language Model with Localized Filtering-based Attention' paper"""
458
+
459
+ def __init__(self, config: YuanConfig):
460
+ super().__init__()
461
+ self.config = config
462
+ self.hidden_size = config.hidden_size
463
+ self.num_heads = config.num_attention_heads
464
+ self.lf_conv2d_group = config.lf_conv2d_group
465
+ self.lf_conv2d_num_pad = config.lf_conv2d_num_pad
466
+ self.use_lfa_bias = config.use_lfa_bias
467
+ try:
468
+ self.attention_projection_size = config.attention_projection_size
469
+ except:
470
+ self.attention_projection_size = None
471
+
472
+ if self.attention_projection_size is None:
473
+ self.head_dim = self.hidden_size // self.num_heads
474
+ else:
475
+ self.head_dim = self.attention_projection_size // self.num_heads
476
+
477
+ self.max_position_embeddings = config.max_position_embeddings
478
+ self.causal_mask = config.causal_mask
479
+ self.attn_mask_type = config.attn_mask_type
480
+ self.softmax_scale = 1.0 / math.sqrt(self.head_dim)
481
+ self.use_flash_attention = config.use_flash_attention
482
+ try:
483
+ self.use_shareqk = config.use_shareqk
484
+ except Exception as e:
485
+ self.use_shareqk=False
486
+ self.dropout = 0.0
487
+ self.attention_projection_size = config.attention_projection_size
488
+ self.v_proj = nn.Linear(self.hidden_size, self.num_heads * self.head_dim, bias=False)
489
+ self.o_proj = nn.Linear(self.num_heads * self.head_dim, self.hidden_size, bias=False)
490
+
491
+ if self.use_shareqk:
492
+ self.qk_proj = nn.Linear(self.hidden_size, self.num_heads * self.head_dim, bias=False)
493
+ self.qk_weight = nn.Parameter(torch.Tensor(2, self.hidden_size))
494
+ self.qk_bias = nn.Parameter(torch.Tensor(2, self.hidden_size))
495
+ else:
496
+ self.lf_gate = LocalizedFiltering(self.hidden_size, self.lf_conv2d_group, self.lf_conv2d_num_pad, self.use_lfa_bias)
497
+ self.get_query_key = nn.Linear(self.hidden_size, 2 * self.attention_projection_size, bias=False)
498
+ self.core_attention = FlashSelfAttention(causal=True, attention_dropout=config.attn_dropout, softmax_scale=self.softmax_scale)
499
+ #self.core_attention_flash = DotProductAttention(num_attention_heads=self.num_heads,
500
+ # kv_channels=self.head_dim)
501
+
502
+ def _shape(self, tensor: torch.Tensor, seq_len: int, bsz: int):
503
+ return tensor.view(bsz, seq_len, self.num_heads, self.head_dim).transpose(1, 2).contiguous()
504
+
505
+ def forward(
506
+ self,
507
+ hidden_states: torch.Tensor,
508
+ attention_mask: Optional[torch.Tensor] = None,
509
+ position_ids: Optional[torch.LongTensor] = None,
510
+ position_ids_k: Optional[torch.LongTensor] = None,
511
+ past_key_value: Optional[Tuple[torch.Tensor]] = None,
512
+ rotary_pos_emb: Optional[Tuple[torch.Tensor]] = None,
513
+ output_attentions: bool = False,
514
+ use_cache: bool = False,
515
+ ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]:
516
+ q_len, bsz, _ = hidden_states.size()
517
+ hidden_states = hidden_states#.to('cuda:1')
518
+ is_first_step = False
519
+ if use_cache:
520
+ if past_key_value is None:
521
+ before_hidden_states = None
522
+ is_first_step = True
523
+ if q_len > 1:
524
+ inference_hidden_states_memory = hidden_states[-2:, :, :]
525
+ else:
526
+ inference_hidden_states_memory = torch.cat((torch.zeros_like(hidden_states), hidden_states), dim=0)
527
+ else:
528
+ before_hidden_states = past_key_value[2]
529
+ inference_hidden_states_memory = torch.cat((before_hidden_states[-1:, :, :], hidden_states), dim=0)
530
+ value_states = self.v_proj(hidden_states).view(q_len, bsz, self.num_heads, self.head_dim)
531
+ if self.use_shareqk:
532
+ qk_states = self.qk_proj(hidden_states).view(q_len, bsz, self.num_heads*self.head_dim)
533
+ query_key = qk_states.unsqueeze(2) * self.qk_weight + self.qk_bias
534
+ query_states, key_states = torch.unbind(query_key, dim=2)
535
+
536
+ query_states = query_states.view(q_len, bsz, self.num_heads, self.head_dim).transpose(1, 2)
537
+ key_states = key_states.view(q_len, bsz, self.num_heads, self.head_dim).transpose(1, 2)
538
+ else:
539
+ hidden_states = self.lf_gate(hidden_states, before_hidden_states)
540
+ mixed_qk_layer = self.get_query_key(hidden_states)
541
+ #mixed_qk_layer = torch.matmul(hidden_states, qk_tensor)
542
+ new_tensor_shape = mixed_qk_layer.size()[:-1] + (self.num_heads, 2 * self.head_dim)
543
+ mixed_qk_layer = mixed_qk_layer.view(*new_tensor_shape)
544
+ (query_states, key_states) = torch.split(mixed_qk_layer, self.head_dim, dim=-1)
545
+
546
+
547
+ kv_seq_len = key_states.shape[1]
548
+ if past_key_value is not None:
549
+ kv_seq_len += past_key_value[0].shape[1]
550
+
551
+ # duplicate the pos_emb for self attention
552
+ if rotary_pos_emb is not None:
553
+ if position_ids.shape[1] == 1:
554
+ q_seq_start = position_ids[0,-1]
555
+ #seq_start = past_key_value[0].shape[0]
556
+ q_seq_end = q_seq_start + 1
557
+ k_seq_end = q_seq_end
558
+ else:
559
+ q_seq_start = 0
560
+ q_seq_end = q_seq_start+key_states.shape[0]
561
+ k_seq_end = q_seq_end
562
+
563
+ rotary_pos_shape = rotary_pos_emb.shape
564
+ if isinstance(rotary_pos_emb, tuple):
565
+ rotary_pos_emb = rotary_pos_emb
566
+ else:
567
+ rotary_pos_emb = ((rotary_pos_emb,) * 2)
568
+ q_pos_emb, k_pos_emb = rotary_pos_emb
569
+ if past_key_value is not None:
570
+ # reuse k, v, self_attention
571
+ key_states = torch.cat([past_key_value[0], key_states], dim=0)
572
+ value_states = torch.cat([past_key_value[1], value_states], dim=0)
573
+ past_key_value = (key_states, value_states, inference_hidden_states_memory) if use_cache else None
574
+ #query_states = apply_rotary_pos_emb(query_states.permute(1, 0, 2, 3), q_pos_emb, position_ids)
575
+ #key_states = apply_rotary_pos_emb(key_states.permute(1, 0, 2, 3), k_pos_emb, position_ids)
576
+ query_states = apply_rotary_pos_emb(query_states, q_pos_emb, position_ids)
577
+ key_states = apply_rotary_pos_emb(key_states, k_pos_emb, position_ids_k)
578
+
579
+ attn_weights = None
580
+ #query_states = query_states.transpose(0,1)
581
+ #key_states = key_states.transpose(0,1)
582
+ #value_states = value_states
583
+ attn_output = self.core_attention(query_states, key_states, value_states)
584
+ #attn_output = self.core_attention(query_states, key_states, value_states, attention_mask)
585
+ q_len, bsz, _, _ = attn_output.shape
586
+ attn_output = attn_output.reshape(q_len, bsz, -1)
587
+
588
+ attn_output = self.o_proj(attn_output)
589
+
590
+ return attn_output, attn_weights, past_key_value
591
+
592
+ class MoEDroplessTokenDispatcher:
593
+ def __init__(self, num_experts: int, config: YuanConfig) -> None:
594
+ self.num_experts = num_experts
595
+ assert self.num_experts > 0, "Expected at least one expert"
596
+ self.router_topk = config.moe_config['moe_top_k']
597
+
598
+ def token_permutation(
599
+ self, hidden_states: torch.Tensor, max_prob: torch.Tensor, max_ind: torch.Tensor
600
+ ):
601
+ self.hidden_shape = hidden_states.shape
602
+ hidden_states = hidden_states.view(-1, self.hidden_shape[-1])
603
+
604
+ if self.router_topk > 1:
605
+ global_local_map = torch.ones_like(max_ind).bool()
606
+ local_indices = max_ind.masked_select(global_local_map)
607
+ local_probs = max_prob.masked_select(global_local_map)
608
+ global_local_map = global_local_map.nonzero()[:, 0]
609
+ global_local_map = global_local_map.view(-1, 1).expand(-1, hidden_states.shape[-1])
610
+ local_hidden_states = torch.gather(hidden_states, 0, global_local_map)
611
+
612
+ indices = torch.argsort(local_indices, dim=0)
613
+ tokens_per_expert = torch.histc(
614
+ local_indices,
615
+ bins=self.num_experts,
616
+ min=0,
617
+ max=self.num_experts - 1,
618
+ )
619
+ tokens_per_expert = tokens_per_expert.cpu().to(torch.long)
620
+
621
+ indices = indices.view(-1, 1).expand(-1, hidden_states.shape[-1])
622
+ permuted_local_hidden_states = torch.gather(local_hidden_states, 0, indices)
623
+ return (permuted_local_hidden_states, tokens_per_expert, local_probs, indices, global_local_map)
624
+
625
+ def token_unpermutation(
626
+ self,
627
+ hidden_states: torch.Tensor,
628
+ scores: torch.Tensor,
629
+ indices: torch.Tensor,
630
+ global_local_map: torch.Tensor = None,
631
+ ):
632
+ scores = scores.to(dtype=hidden_states.dtype)
633
+ unpermuted_local_hidden = torch.zeros_like(hidden_states)
634
+ #assert indices.shape == hidden_states.shape, f'{indices.shape}, {hidden_states.shape}'
635
+ unpermuted_local_hidden = unpermuted_local_hidden.scatter(0, indices, hidden_states)
636
+
637
+ if self.router_topk > 1:
638
+ unpermuted_local_hidden = unpermuted_local_hidden * scores.view(-1, 1)
639
+ unpermuted_local_bias = None
640
+ output_total = unpermuted_local_hidden
641
+ output_bias_total = unpermuted_local_bias
642
+
643
+ if self.router_topk > 1:
644
+ global_num_tokens = self.hidden_shape[0] * self.hidden_shape[1]
645
+ global_hidden_shape = [global_num_tokens, hidden_states.shape[-1]]
646
+ unpermuted_global_hidden = torch.zeros(
647
+ global_hidden_shape,
648
+ dtype=hidden_states.dtype,
649
+ device=hidden_states.device,
650
+ )
651
+ output_total = unpermuted_global_hidden.scatter_add(
652
+ 0, global_local_map, unpermuted_local_hidden
653
+ )
654
+
655
+ output_total = output_total.view(self.hidden_shape)
656
+
657
+ return output_total
658
+
659
+ class YuanExpertMLP(nn.Module):
660
+ def __init__(self, hidden_size, ffn_hidden_size):
661
+ super().__init__()
662
+ def glu(x):
663
+ x = torch.chunk(x, 2, dim=-1)
664
+ return torch.nn.functional.silu(x[0]) * x[1]
665
+
666
+ self.activation_func = glu
667
+
668
+ self.w1 = nn.Linear(hidden_size, ffn_hidden_size * 2, bias=False)
669
+ self.w2 = nn.Linear(ffn_hidden_size, hidden_size, bias=False)
670
+
671
+ def forward(self, x):
672
+ # GLU activation: split into two halves
673
+ w1_out = self.w1(x) # [..., ffn_hidden_size * 2]
674
+ intermediate = self.activation_func(w1_out)
675
+ return self.w2(intermediate)
676
+
677
+
678
+ class GroupedMLP(nn.Module):
679
+ """An efficient implementation of the Experts layer using CUTLASS GroupedGEMM.
680
+
681
+ This class is designed to execute multiple experts in parallel, thereby maximizing computational efficiency.
682
+ """
683
+
684
+ def __init__(self, num_experts: int, config: YuanConfig):
685
+ super().__init__()
686
+ self.num_experts = num_experts
687
+ self.config = config
688
+ self.hidden_size = self.config.hidden_size
689
+ self.ffn_hidden_size = config.ffn_hidden_size
690
+
691
+ self.experts = nn.ModuleList([
692
+ YuanExpertMLP(self.hidden_size, self.ffn_hidden_size)
693
+ for _ in range(num_experts)
694
+ ])
695
+ #self.w1 = nn.ModuleList([nn.Linear(self.config.hidden_size, self.ffn_hidden_size * 2, bias=False) for _ in range(num_experts)])
696
+ #self.w2 = nn.ModuleList([nn.Linear(self.ffn_hidden_size, self.config.hidden_size, bias=False) for _ in range(num_experts)])
697
+ def forward(self, permuted_hidden_states, tokens_per_expert):
698
+ fc2_outputs = []
699
+ start_idx = 0
700
+ for i in range(self.num_experts):
701
+ #if tokens_per_expert[i] == 0:
702
+ # continue
703
+ end_idx = start_idx + tokens_per_expert[i]
704
+ fc2_output = self.experts[i](permuted_hidden_states[start_idx:end_idx])
705
+ fc2_outputs.append(fc2_output)
706
+ start_idx = end_idx
707
+ fc2_output = torch.cat(fc2_outputs, dim=0)
708
+ return fc2_output#.to('cuda:1')
709
+
710
+ class YuanMoeLayer(nn.Module):
711
+ def __init__(self, config:YuanConfig, num_experts):
712
+ super().__init__()
713
+ self.config = config
714
+ self.num_experts = num_experts
715
+ self.top_k = config.moe_config['moe_top_k']
716
+ self.norm_topk_prob = config.moe_config['norm_topk_prob']
717
+ self.hidden_size = config.hidden_size
718
+
719
+ expert_indices_offset = (0)
720
+
721
+ #self.router = ParallelAttention_router(config, num_experts)
722
+ if config.moe_config['router_type'] == 'attn_router':
723
+ self.router = ParallelAttention_router(config, num_experts)
724
+ else:
725
+ self.router = nn.Linear(config.hidden_size, self.num_experts, bias=False)
726
+ self.token_dispatcher = MoEDroplessTokenDispatcher(self.num_experts, config=self.config)
727
+ self.experts = GroupedMLP(self.num_experts, self.config)
728
+
729
+ def routing(self, logits: torch.Tensor) -> torch.Tensor:
730
+ top_logits, indices = torch.topk(logits, k=self.top_k, dim=1)
731
+ scores = torch.softmax(top_logits, dim=-1, dtype=torch.float32).type_as(logits)
732
+ return scores, indices
733
+
734
+ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
735
+ batch_size, sequence_length, hidden_dim = hidden_states.shape
736
+ #logits = self.gate(hidden_states)
737
+ logits = self.router(hidden_states)
738
+ if len(logits.shape) > 2:
739
+ logits = logits.view(-1, self.num_experts)
740
+ scores, indices = self.routing(logits)
741
+ scores = scores.to(hidden_states.dtype)
742
+ (dispatched_input, tokens_per_expert, scores, indices, global_local_map, ) = self.token_dispatcher.token_permutation(hidden_states, scores, indices)
743
+ expert_output = self.experts(dispatched_input, tokens_per_expert)
744
+ output = self.token_dispatcher.token_unpermutation(expert_output, scores, indices, global_local_map)
745
+ return output
746
+
747
+ class YuanDecoderLayer(nn.Module):
748
+ def __init__(self, config: YuanConfig, num_layer):
749
+ super().__init__()
750
+ self.hidden_size = config.hidden_size
751
+ self.self_attn = YuanAttention(config=config)
752
+ self.num_layer = num_layer
753
+
754
+ if config.use_moe:
755
+ assert config.moe_config['per_layer_experts_blocks'] != None or config.moe_config['num_experts'] != None
756
+ if config.moe_config['per_layer_experts_blocks'] != None:
757
+ self.num_experts = config.moe_config['per_layer_experts_blocks'][num_layer]
758
+ else:
759
+ self.num_experts = config.moe_config['num_experts']
760
+ self.mlp = YuanMoeLayer(config, self.num_experts)
761
+
762
+ else:
763
+ self.mlp = YuanMLP(
764
+ hidden_size=self.hidden_size,
765
+ intermediate_size=config.intermediate_size,
766
+ hidden_act=config.hidden_act,
767
+ )
768
+
769
+ self.input_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
770
+ self.post_attention_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
771
+
772
+ def forward(
773
+ self,
774
+ hidden_states: torch.Tensor,
775
+ attention_mask: Optional[torch.Tensor] = None,
776
+ position_ids: Optional[torch.LongTensor] = None,
777
+ position_ids_k: Optional[torch.LongTensor] = None,
778
+ past_key_value: Optional[Tuple[torch.Tensor]] = None,
779
+ rotary_pos_emb: Optional[Tuple[torch.Tensor]] = None,
780
+ output_attentions: Optional[bool] = False,
781
+ use_cache: Optional[bool] = False,
782
+ ) -> Tuple[torch.FloatTensor, Optional[Tuple[torch.FloatTensor, torch.FloatTensor]]]:
783
+ """
784
+ Args:
785
+ hidden_states (`torch.FloatTensor`): input to the layer of shape `(batch, seq_len, embed_dim)`
786
+ attention_mask (`torch.FloatTensor`, *optional*): attention mask of size
787
+ `(batch, 1, tgt_len, src_len)` where padding elements are indicated by very large negative values.
788
+ output_attentions (`bool`, *optional*):
789
+ Whether or not to return the attentions tensors of all attention layers. See `attentions` under
790
+ returned tensors for more detail.
791
+ use_cache (`bool`, *optional*):
792
+ If set to `True`, `past_key_values` key value states are returned and can be used to speed up decoding
793
+ (see `past_key_values`).
794
+ past_key_value (`Tuple(torch.FloatTensor)`, *optional*): cached past key and value projection states
795
+ """
796
+ residual = hidden_states#.to('cuda:1')
797
+ hidden_states = self.input_layernorm(hidden_states) #.to('cuda:0')).to('cuda:1')
798
+ # Self Attention
799
+ hidden_states, self_attn_weights, present_key_value = self.self_attn(
800
+ hidden_states=hidden_states,
801
+ attention_mask=attention_mask,
802
+ position_ids=position_ids,
803
+ position_ids_k=position_ids_k,
804
+ past_key_value=past_key_value,
805
+ rotary_pos_emb=rotary_pos_emb,
806
+ output_attentions=output_attentions,
807
+ use_cache=use_cache,
808
+ )
809
+ hidden_states = residual + hidden_states.permute(1, 0, 2)
810
+ # Fully Connected
811
+ residual = hidden_states#.to('cuda:1')
812
+ hidden_states = self.post_attention_layernorm(hidden_states) #.to('cuda:0')).to('cuda:1')
813
+ hidden_states = self.mlp(hidden_states)# .to('cuda:1')
814
+ hidden_states = residual + hidden_states
815
+ outputs = (hidden_states,)
816
+
817
+ if output_attentions:
818
+ outputs += (self_attn_weights,)
819
+
820
+ if use_cache:
821
+ outputs += (present_key_value,)
822
+
823
+ return outputs
824
+
825
+
826
+ YUAN_START_DOCSTRING = r"""
827
+ This model inherits from [`PreTrainedModel`]. Check the superclass documentation for the generic methods the
828
+ library implements for all its model (such as downloading or saving, resizing the input embeddings, pruning heads
829
+ etc.)
830
+
831
+ This model is also a PyTorch [torch.nn.Module](https://pytorch.org/docs/stable/nn.html#torch.nn.Module) subclass.
832
+ Use it as a regular PyTorch Module and refer to the PyTorch documentation for all matter related to general usage
833
+ and behavior.
834
+
835
+ Parameters:
836
+ config ([`YuanConfig`]):
837
+ Model configuration class with all the parameters of the model. Initializing with a config file does not
838
+ load the weights associated with the model, only the configuration. Check out the
839
+ [`~PreTrainedModel.from_pretrained`] method to load the model weights.
840
+ """
841
+
842
+
843
+ @add_start_docstrings(
844
+ "The bare Yuan Model outputting raw hidden-states without any specific head on top.",
845
+ YUAN_START_DOCSTRING,
846
+ )
847
+ class YuanPreTrainedModel(PreTrainedModel):
848
+ config_class = YuanConfig
849
+ base_model_prefix = "model"
850
+ supports_gradient_checkpointing = True
851
+ _no_split_modules = ["YuanDecoderLayer"]
852
+ _skip_keys_device_placement = "past_key_values"
853
+ _keys_to_ignore_on_load_unexpected = [r"decoder\.version"]
854
+
855
+ def _init_weights(self, module):
856
+ std = self.config.initializer_range
857
+ if isinstance(module, nn.Linear):
858
+ module.weight.data.normal_(mean=0.0, std=std)
859
+ if module.bias is not None:
860
+ module.bias.data.zero_()
861
+ elif isinstance(module, nn.Embedding):
862
+ module.weight.data.normal_(mean=0.0, std=std)
863
+ if module.padding_idx is not None:
864
+ module.weight.data[module.padding_idx].zero_()
865
+
866
+ def _set_gradient_checkpointing(self, module, value=False):
867
+ if isinstance(module, YuanModel):
868
+ module.gradient_checkpointing = value
869
+
870
+
871
+ YUAN_INPUTS_DOCSTRING = r"""
872
+ Args:
873
+ input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`):
874
+ Indices of input sequence tokens in the vocabulary. Padding will be ignored by default should you provide
875
+ it.
876
+
877
+ Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and
878
+ [`PreTrainedTokenizer.__call__`] for details.
879
+
880
+ [What are input IDs?](../glossary#input-ids)
881
+ attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):
882
+ Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:
883
+
884
+ - 1 for tokens that are **not masked**,
885
+ - 0 for tokens that are **masked**.
886
+
887
+ [What are attention masks?](../glossary#attention-mask)
888
+
889
+ Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and
890
+ [`PreTrainedTokenizer.__call__`] for details.
891
+
892
+ If `past_key_values` is used, optionally only the last `decoder_input_ids` have to be input (see
893
+ `past_key_values`).
894
+
895
+ If you want to change padding behavior, you should read [`modeling_opt._prepare_decoder_attention_mask`]
896
+ and modify to your needs. See diagram 1 in [the paper](https://arxiv.org/abs/1910.13461) for more
897
+ information on the default strategy.
898
+
899
+ - 1 indicates the head is **not masked**,
900
+ - 0 indicates the head is **masked**.
901
+ position_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
902
+ Indices of positions of each input sequence tokens in the position embeddings. Selected in the range `[0,
903
+ config.n_positions - 1]`.
904
+
905
+ [What are position IDs?](../glossary#position-ids)
906
+ past_key_values (`tuple(tuple(torch.FloatTensor))`, *optional*, returned when `use_cache=True` is passed or when `config.use_cache=True`):
907
+ Tuple of `tuple(torch.FloatTensor)` of length `config.n_layers`, with each tuple having 2 tensors of shape
908
+ `(batch_size, num_heads, sequence_length, embed_size_per_head)`) and 2 additional tensors of shape
909
+ `(batch_size, num_heads, encoder_sequence_length, embed_size_per_head)`.
910
+
911
+ Contains pre-computed hidden-states (key and values in the self-attention blocks and in the cross-attention
912
+ blocks) that can be used (see `past_key_values` input) to speed up sequential decoding.
913
+
914
+ If `past_key_values` are used, the user can optionally input only the last `decoder_input_ids` (those that
915
+ don't have their past key value states given to this model) of shape `(batch_size, 1)` instead of all
916
+ `decoder_input_ids` of shape `(batch_size, sequence_length)`.
917
+ inputs_embeds (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`, *optional*):
918
+ Optionally, instead of passing `input_ids` you can choose to directly pass an embedded representation. This
919
+ is useful if you want more control over how to convert `input_ids` indices into associated vectors than the
920
+ model's internal embedding lookup matrix.
921
+ use_cache (`bool`, *optional*):
922
+ If set to `True`, `past_key_values` key value states are returned and can be used to speed up decoding (see
923
+ `past_key_values`).
924
+ output_attentions (`bool`, *optional*):
925
+ Whether or not to return the attentions tensors of all attention layers. See `attentions` under returned
926
+ tensors for more detail.
927
+ output_hidden_states (`bool`, *optional*):
928
+ Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors for
929
+ more detail.
930
+ return_dict (`bool`, *optional*):
931
+ Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.
932
+ """
933
+
934
+
935
+ @add_start_docstrings(
936
+ "The bare Yuan Model outputting raw hidden-states without any specific head on top.",
937
+ YUAN_START_DOCSTRING,
938
+ )
939
+ class YuanModel(YuanPreTrainedModel):
940
+ """
941
+ Transformer decoder consisting of *config.num_hidden_layers* layers. Each layer is a [`YuanDecoderLayer`]
942
+
943
+ Args:
944
+ config: YuanConfig
945
+ """
946
+
947
+ def __init__(self, config: YuanConfig):
948
+ super().__init__(config)
949
+ self.padding_idx = config.pad_token_id
950
+ self.vocab_size = config.vocab_size
951
+
952
+ #TODO: control it by config
953
+ self.eod_token = config.eod_token
954
+ self.reset_attention_mask = config.reset_attention_mask
955
+ self.reset_position_ids = config.reset_position_ids
956
+ self.max_position_embeddings = config.max_position_embeddings
957
+ self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx)
958
+ self.layers = nn.ModuleList([YuanDecoderLayer(config, i) for i in range(config.num_hidden_layers)])
959
+ self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
960
+ self.gradient_checkpointing = False
961
+ # Initialize weights and apply final processing
962
+ self.post_init()
963
+
964
+ self.seq_length = config.max_position_embeddings
965
+ rotary_dim = config.hidden_size // config.num_attention_heads
966
+ if config.rotary_percent < 1.0:
967
+ rotary_dim = int(rotary_dim * config.rotary_percent)
968
+ self.rotary_pos_emb = YuanRotaryEmbedding(rotary_dim, base=config.rotary_base, dtype=config.torch_dtype)
969
+
970
+
971
+ def get_input_embeddings(self, input_ids: torch.Tensor) -> torch.Tensor:
972
+ return self.embed_tokens(input_ids)
973
+
974
+ def set_input_embeddings(self, value):
975
+ self.embed_tokens = value
976
+
977
+ # Copied from transformers.models.bart.modeling_bart.BartDecoder._prepare_decoder_attention_mask
978
+ def _prepare_decoder_attention_mask(self, attention_mask, input_shape, inputs_embeds, past_key_values_length):
979
+ # create causal mask
980
+ # [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]
981
+ combined_attention_mask = None
982
+ if input_shape[-1] > 1:
983
+ combined_attention_mask = _make_causal_mask(
984
+ input_shape,
985
+ inputs_embeds.dtype,
986
+ device=inputs_embeds.device,
987
+ past_key_values_length=past_key_values_length,
988
+ )
989
+
990
+ if attention_mask is not None:
991
+ # [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]
992
+ expanded_attn_mask = _expand_mask(attention_mask, inputs_embeds.dtype, tgt_len=input_shape[-1]).to(
993
+ inputs_embeds.device
994
+ )
995
+ combined_attention_mask = (
996
+ expanded_attn_mask if combined_attention_mask is None else expanded_attn_mask + combined_attention_mask
997
+ )
998
+
999
+ return combined_attention_mask
1000
+
1001
+ def _prepare_decoder_attention_mask_training(self, input_id, inputs_embeds, eod_token, reset_mask_flag ,reset_attention_mask=True, reset_position_ids=True):
1002
+
1003
+ micro_batch_size, seq_length = input_id.size()
1004
+
1005
+ attention_mask = torch.tril(torch.ones(
1006
+ (micro_batch_size, seq_length, seq_length), device=inputs_embeds.device)).view(
1007
+ micro_batch_size, 1, seq_length, seq_length)
1008
+
1009
+ position_ids = torch.arange(seq_length, dtype=torch.long,
1010
+ device=inputs_embeds.device)
1011
+ position_ids = position_ids.unsqueeze(0).expand_as(input_id)
1012
+
1013
+ if reset_position_ids:
1014
+ position_ids = position_ids.clone()
1015
+
1016
+ if reset_position_ids or reset_attention_mask:
1017
+ # Loop through the batches:
1018
+ for b in range(micro_batch_size):
1019
+
1020
+ # Find indecies where EOD token is.
1021
+ eod_index = position_ids[b, input_id[b] == eod_token]
1022
+
1023
+ # Detach indecies from positions if going to modify positions.
1024
+ if reset_position_ids:
1025
+ eod_index = eod_index.clone()
1026
+ # Loop through EOD indecies:
1027
+ prev_index = 0
1028
+ for j in range(eod_index.size()[0]):
1029
+ i = eod_index[j]
1030
+ # Mask attention loss.
1031
+ if reset_attention_mask:
1032
+ attention_mask[b, 0, (i + 1):, :(i + 1)] = 0
1033
+ # Reset positions.
1034
+ if reset_position_ids:
1035
+ position_ids[b, (i + 1):] -= (i + 1 - prev_index)
1036
+ prev_index = i + 1
1037
+
1038
+ inverted_mask = 1 - attention_mask
1039
+ output_attn_mask = inverted_mask.masked_fill(inverted_mask.to(torch.bool), torch.finfo(inputs_embeds.dtype).min)
1040
+ if reset_mask_flag:
1041
+ output_attn_mask = output_attn_mask[:,:,-1:,:]
1042
+ return output_attn_mask, position_ids
1043
+
1044
+ @add_start_docstrings_to_model_forward(YUAN_INPUTS_DOCSTRING)
1045
+ def forward(
1046
+ self,
1047
+ input_ids: torch.LongTensor = None,
1048
+ attention_mask: Optional[torch.Tensor] = None,
1049
+ position_ids: Optional[torch.LongTensor] = None,
1050
+ past_key_values: Optional[List[torch.FloatTensor]] = None,
1051
+ inputs_embeds: Optional[torch.FloatTensor] = None,
1052
+ use_cache: Optional[bool] = None,
1053
+ output_attentions: Optional[bool] = None,
1054
+ output_hidden_states: Optional[bool] = None,
1055
+ output_router_logits: Optional[bool] = None,
1056
+ return_dict: Optional[bool] = None,
1057
+ ) -> Union[Tuple, BaseModelOutputWithPast, torch.Tensor]:
1058
+ output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
1059
+ output_router_logits = (
1060
+ output_router_logits if output_router_logits is not None else self.config.output_router_logits
1061
+ )
1062
+ output_hidden_states = (
1063
+ output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
1064
+ )
1065
+ use_cache = use_cache if use_cache is not None else self.config.use_cache
1066
+ return_dict = return_dict if return_dict is not None else self.config.use_return_dict
1067
+ input_ids1 = copy.deepcopy(input_ids)
1068
+ reset_mask_flag = False
1069
+ if past_key_values:
1070
+ input_ids = input_ids
1071
+ input_ids = input_ids[:,-1:]
1072
+ if use_cache:
1073
+ reset_mask_flag = True
1074
+ # retrieve input_ids and inputs_embeds
1075
+ if input_ids is not None and inputs_embeds is not None:
1076
+ raise ValueError("You cannot specify both decoder_input_ids and decoder_inputs_embeds at the same time")
1077
+ elif input_ids is not None:
1078
+ input_ids = input_ids
1079
+ batch_size, seq_length = input_ids.shape
1080
+ elif inputs_embeds is not None:
1081
+ inputs_embeds = inputs_embeds.transpose(0,1)
1082
+ batch_size, seq_length, _ = inputs_embeds.shape
1083
+ else:
1084
+ raise ValueError("You have to specify either decoder_input_ids or decoder_inputs_embeds")
1085
+
1086
+ seq_length_with_past = seq_length
1087
+ past_key_values_length = 0
1088
+ if past_key_values is not None:
1089
+ #past_key_values_length = past_key_values[0][0].shape[2]
1090
+ #modify
1091
+ past_key_values_length = past_key_values[0][0].shape[0]
1092
+ seq_length_with_past = seq_length_with_past + past_key_values_length
1093
+ else:
1094
+ print('1111')
1095
+
1096
+ # modify to reset position ids
1097
+ if past_key_values is not None:
1098
+ pos_start = position_ids[:,-1]+1
1099
+ pos_end = pos_start+past_key_values[0][0].shape[0]-position_ids.shape[1]+1
1100
+ position_ids_k = torch.arange(pos_start.item(), pos_end.item()).to(position_ids.device)
1101
+ position_ids_k = position_ids_k.unsqueeze(0)
1102
+ position_ids_k = torch.cat((position_ids, position_ids_k), dim=1)
1103
+ position_ids = position_ids[:,-1]+past_key_values[0][0].shape[0]-position_ids.shape[1]+1
1104
+ position_ids = position_ids.unsqueeze(0)
1105
+ else:
1106
+ position_ids_k = position_ids
1107
+
1108
+ if position_ids is None:
1109
+ device = input_ids.device if input_ids is not None else inputs_embeds.device
1110
+ position_ids = torch.arange(
1111
+ past_key_values_length, seq_length + past_key_values_length, dtype=torch.long, device=device
1112
+ )
1113
+ position_ids = position_ids.unsqueeze(0).view(-1, seq_length)
1114
+ else:
1115
+ pass
1116
+
1117
+ if inputs_embeds is None:
1118
+ inputs_embeds = self.embed_tokens(input_ids).transpose(0,1)
1119
+
1120
+ if self.training or self.reset_position_ids:
1121
+ attention_mask, _ = self._prepare_decoder_attention_mask_training(input_ids1, inputs_embeds, self.eod_token, reset_mask_flag, self.reset_attention_mask, self.reset_position_ids)
1122
+ else:
1123
+ if attention_mask is None:
1124
+ attention_mask = torch.ones(
1125
+ (batch_size, seq_length_with_past), dtype=torch.bool, device=inputs_embeds.device
1126
+ )
1127
+ attention_mask = self._prepare_decoder_attention_mask(
1128
+ attention_mask, (batch_size, seq_length), inputs_embeds, past_key_values_length
1129
+ )
1130
+
1131
+ #rotary_pos_emb = self.rotary_pos_emb(self.max_position_embeddings)
1132
+ # Rotary positional embeddings (embedding is None for PP intermediate devices)
1133
+ rotary_pos_emb = None
1134
+ rotary_pos_emb = self.rotary_pos_emb(self.max_position_embeddings)
1135
+
1136
+ hidden_states = inputs_embeds
1137
+ if self.gradient_checkpointing and self.training:
1138
+ if use_cache:
1139
+ logger.warning_once(
1140
+ "`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`..."
1141
+ )
1142
+ use_cache = False
1143
+
1144
+ # decoder layers
1145
+ all_hidden_states = () if output_hidden_states else None
1146
+ all_self_attns = () if output_attentions else None
1147
+ next_decoder_cache = () if use_cache else None
1148
+ #position_ids = position_ids.cpu()
1149
+ #position_ids_k = position_ids_k.cpu()
1150
+ for idx, decoder_layer in enumerate(self.layers):
1151
+ if output_hidden_states:
1152
+ all_hidden_states += (hidden_states,)
1153
+
1154
+ past_key_value = past_key_values[idx] if past_key_values is not None else None
1155
+
1156
+ if self.gradient_checkpointing and self.training:
1157
+ def create_custom_forward(module):
1158
+ def custom_forward(*inputs):
1159
+ # None for past_key_value
1160
+ return module(*inputs, output_attentions, None)
1161
+
1162
+ return custom_forward
1163
+
1164
+ layer_outputs = torch.utils.checkpoint.checkpoint(
1165
+ create_custom_forward(decoder_layer),
1166
+ hidden_states,
1167
+ attention_mask,
1168
+ position_ids,
1169
+ None,
1170
+ )
1171
+ else:
1172
+ layer_outputs = decoder_layer(
1173
+ hidden_states,
1174
+ attention_mask=attention_mask,
1175
+ position_ids=position_ids,
1176
+ position_ids_k=position_ids_k,
1177
+ past_key_value=past_key_value,
1178
+ rotary_pos_emb=rotary_pos_emb,
1179
+ output_attentions=output_attentions,
1180
+ use_cache=use_cache,
1181
+ )
1182
+ hidden_states = layer_outputs[0]
1183
+
1184
+ if use_cache:
1185
+ next_decoder_cache += (layer_outputs[2 if output_attentions else 1],)
1186
+
1187
+ if output_attentions:
1188
+ all_self_attns += (layer_outputs[1],)
1189
+ hidden_states = hidden_states#.to('cuda:0')
1190
+ hidden_states = self.norm(hidden_states)
1191
+ #print(hidden_states)
1192
+ # add hidden states from the last decoder layer
1193
+ if output_hidden_states:
1194
+ all_hidden_states += (hidden_states,)
1195
+ next_cache = next_decoder_cache if use_cache else None
1196
+ if not return_dict:
1197
+ return tuple(v for v in [hidden_states, next_cache, all_hidden_states, all_self_attns] if v is not None)
1198
+ return BaseModelOutputWithPast(
1199
+ last_hidden_state=hidden_states,
1200
+ past_key_values=next_cache,
1201
+ hidden_states=all_hidden_states,
1202
+ attentions=all_self_attns,
1203
+ )
1204
+
1205
+
1206
+ class YuanForCausalLM(YuanPreTrainedModel, GenerationMixin):
1207
+ def __init__(self, config):
1208
+ super().__init__(config)
1209
+ self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
1210
+ self.model = YuanModel(config)
1211
+ #self.output = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
1212
+ self.post_init()
1213
+
1214
+ def get_input_embeddings(self):
1215
+ return self.model.embed_tokens
1216
+
1217
+ def set_input_embeddings(self, value):
1218
+ self.model.embed_tokens = value
1219
+
1220
+ def get_output_embeddings(self):
1221
+ return self.lm_head
1222
+
1223
+ def set_output_embeddings(self, new_embeddings):
1224
+ self.lm_head = new_embeddings
1225
+
1226
+ def set_decoder(self, decoder):
1227
+ self.model = decoder
1228
+
1229
+ def get_decoder(self):
1230
+ return self.model
1231
+
1232
+ def get_loss_mask(self, input_ids, labels, eod_token, sep_token):
1233
+ micro_batch_size, seq_length = input_ids.size()
1234
+ loss_mask = torch.ones(input_ids.size(), dtype=torch.float, device=input_ids.device)
1235
+ position_ids = torch.arange(seq_length, dtype=torch.long, device=input_ids.device)
1236
+ position_ids = position_ids.unsqueeze(0).expand_as(input_ids)
1237
+
1238
+
1239
+ """modify loss_mask to only calculate the loss of the answer (separated with [SEP])"""
1240
+
1241
+ for b in range(micro_batch_size):
1242
+ eod_indexs = position_ids[b, input_ids[b] == eod_token]
1243
+ sep_indexs = position_ids[b, input_ids[b] == sep_token]
1244
+
1245
+ if len(eod_indexs) == 0 or len(sep_indexs) == 0:
1246
+ loss_mask[b] = 1.0
1247
+ else:
1248
+ if eod_indexs[0] > sep_indexs[0]:
1249
+ loss_mask[b, 0:sep_indexs[0]] = 0
1250
+
1251
+ if len(eod_indexs) == len(sep_indexs):
1252
+ for ii, eod_index in enumerate(eod_indexs):
1253
+ start_index = eod_index
1254
+ if ii == (len(sep_indexs) - 1):
1255
+ stop_index = seq_length
1256
+ else:
1257
+ stop_index = sep_indexs[ii + 1]
1258
+ loss_mask[b, start_index:stop_index] = 0.0
1259
+ else:
1260
+ if len(eod_indexs) > len(sep_indexs):
1261
+ loss_mask[b,:] = 1.0
1262
+ else:
1263
+ for ii, eod_index in enumerate(eod_indexs):
1264
+ start_index = eod_index
1265
+ stop_index = sep_indexs[ii + 1]
1266
+
1267
+ loss_mask[b, start_index:stop_index] = 0.0
1268
+
1269
+ elif eod_indexs[0] < sep_indexs[0]:
1270
+
1271
+ if len(eod_indexs) == len(sep_indexs):
1272
+ for ii, eod_index in enumerate(eod_indexs):
1273
+ start_index = eod_index
1274
+ stop_index = sep_indexs[ii]
1275
+ loss_mask[b, start_index:stop_index] = 0.0
1276
+
1277
+ else:
1278
+ if len(eod_indexs) < len(sep_indexs):
1279
+ loss_mask[b,:] = 1.0
1280
+ else:
1281
+ for ii, eod_index in enumerate(eod_indexs):
1282
+ start_index = eod_index
1283
+ if ii >= len(sep_indexs):
1284
+ stop_index = seq_length
1285
+ else:
1286
+ stop_index = sep_indexs[ii]
1287
+ loss_mask[b, start_index:stop_index] = 0.0
1288
+
1289
+ loss_mask[input_ids == eod_token] = 1.0
1290
+ return loss_mask
1291
+ @add_start_docstrings_to_model_forward(YUAN_INPUTS_DOCSTRING)
1292
+ @replace_return_docstrings(output_type=CausalLMOutputWithPast, config_class=_CONFIG_FOR_DOC)
1293
+ def forward(
1294
+ self,
1295
+ input_ids: torch.LongTensor = None,
1296
+ attention_mask: Optional[torch.Tensor] = None,
1297
+ position_ids: Optional[torch.LongTensor] = None,
1298
+ past_key_values: Optional[List[torch.FloatTensor]] = None,
1299
+ inputs_embeds: Optional[torch.FloatTensor] = None,
1300
+ labels: Optional[torch.LongTensor] = None,
1301
+ use_cache: Optional[bool] = None,
1302
+ output_attentions: Optional[bool] = None,
1303
+ output_hidden_states: Optional[bool] = None,
1304
+ return_dict: Optional[bool] = None,
1305
+ ) -> Union[Tuple, CausalLMOutputWithPast]:
1306
+ """
1307
+ ## modify delete routers
1308
+ Args:
1309
+ labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
1310
+ Labels for computing the masked language modeling loss. Indices should either be in `[0, ...,
1311
+ config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored
1312
+ (masked), the loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`.
1313
+
1314
+ Returns:
1315
+
1316
+ Example:
1317
+
1318
+ ```python
1319
+ >>> from transformers import AutoTokenizer, YuanForCausalLM
1320
+
1321
+ >>> model = YuanForCausalLM.from_pretrained(PATH_TO_CONVERTED_WEIGHTS)
1322
+ >>> tokenizer = AutoTokenizer.from_pretrained(PATH_TO_CONVERTED_TOKENIZER)
1323
+
1324
+ >>> prompt = "Hey, are you consciours? Can you talk to me?"
1325
+ >>> inputs = tokenizer(prompt, return_tensors="pt")
1326
+
1327
+ >>> # Generate
1328
+ >>> generate_ids = model.generate(inputs.input_ids, max_length=30)
1329
+ >>> tokenizer.batch_decode(generate_ids, skip_special_tokens=True, clean_up_tokenization_spaces=False)[0]
1330
+ "Hey, are you consciours? Can you talk to me?\nI'm not consciours, but I can talk to you."
1331
+ ```"""
1332
+
1333
+ output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
1334
+
1335
+ output_hidden_states = (
1336
+ output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
1337
+ )
1338
+
1339
+ return_dict = return_dict if return_dict is not None else self.config.use_return_dict
1340
+
1341
+ outputs = self.model(
1342
+ input_ids=input_ids,
1343
+ attention_mask=attention_mask,
1344
+ position_ids=position_ids,
1345
+ past_key_values=past_key_values,
1346
+ inputs_embeds=inputs_embeds,
1347
+ use_cache=use_cache,
1348
+ output_attentions=output_attentions,
1349
+ output_hidden_states=output_hidden_states,
1350
+ return_dict=return_dict,
1351
+ )
1352
+ hidden_states = outputs[0].transpose(0,1)
1353
+ #print(hidden_states)
1354
+ logits = self.lm_head(hidden_states)
1355
+
1356
+ loss = None
1357
+ if labels is not None:
1358
+ if self.use_loss_mask:
1359
+ loss_mask = self.get_loss_mask(input_ids, labels, self.eod_token, self.sep_token)
1360
+ # Shift so that tokens < n predict n
1361
+ shift_logits = logits[..., :-1, :].contiguous()
1362
+ shift_labels = labels[..., 1:].contiguous()
1363
+ # Flatten the tokens
1364
+ if self.use_loss_mask:
1365
+ loss_fct = CrossEntropyLoss(reduction='none')
1366
+ shift_logits = shift_logits.view(-1, self.config.vocab_size)
1367
+ shift_labels = shift_labels.view(-1)
1368
+ # Enable model parallelism
1369
+ shift_labels = shift_labels.to(shift_logits.device)
1370
+ loss = loss_fct(shift_logits, shift_labels)
1371
+ loss = torch.sum(loss * loss_mask) / loss_mask.sum()
1372
+ else:
1373
+ loss_fct = CrossEntropyLoss()
1374
+ shift_logits = shift_logits.view(-1, self.config.vocab_size)
1375
+ shift_labels = shift_labels.view(-1)
1376
+ # Enable model parallelism
1377
+ shift_labels = shift_labels.to(shift_logits.device)
1378
+ loss = loss_fct(shift_logits, shift_labels)
1379
+ if not return_dict:
1380
+ output = (logits,) + outputs[1:]
1381
+ return (loss,) + output if loss is not None else output
1382
+
1383
+ return CausalLMOutputWithPast(
1384
+ loss=loss,
1385
+ logits=logits,
1386
+ past_key_values=outputs.past_key_values,
1387
+ hidden_states=hidden_states,
1388
+ attentions=outputs.attentions,
1389
+ )
1390
+
1391
+ def prepare_inputs_for_generation(
1392
+ self, input_ids, past_key_values=None, attention_mask=None, inputs_embeds=None, **kwargs
1393
+ ):
1394
+
1395
+ position_ids = kwargs.get("position_ids", None)
1396
+ if attention_mask is not None and position_ids is None:
1397
+ # create position_ids on the fly for batch generation
1398
+ position_ids = attention_mask.long().cumsum(-1) - 1
1399
+ position_ids.masked_fill_(attention_mask == 0, 1)
1400
+ if past_key_values:
1401
+ position_ids = position_ids[:, -1].unsqueeze(-1)
1402
+
1403
+ # if `inputs_embeds` are passed, we only want to use them in the 1st generation step
1404
+ if inputs_embeds is not None and past_key_values is None:
1405
+ model_inputs = {"inputs_embeds": inputs_embeds}
1406
+ else:
1407
+ model_inputs = {"input_ids": input_ids}
1408
+
1409
+ model_inputs.update(
1410
+ {
1411
+ "position_ids": position_ids,
1412
+ "past_key_values": past_key_values,
1413
+ "use_cache": kwargs.get("use_cache"),
1414
+ "attention_mask": attention_mask,
1415
+ }
1416
+ )
1417
+ return model_inputs
1418
+
1419
+ @staticmethod
1420
+ def _reorder_cache(past_key_values, beam_idx):
1421
+ reordered_past = ()
1422
+ for layer_past in past_key_values:
1423
+ reordered_past += (tuple(past_state.index_select(0, beam_idx) for past_state in layer_past),)
1424
+ return reordered_past
1425
+
1426
+
1427
+ @add_start_docstrings(
1428
+ """
1429
+ The Yuan Model transformer with a sequence classification head on top (linear layer).
1430
+
1431
+ [`YuanForSequenceClassification`] uses the last token in order to do the classification, as other causal models
1432
+ (e.g. GPT-2) do.
1433
+
1434
+ Since it does classification on the last token, it requires to know the position of the last token. If a
1435
+ `pad_token_id` is defined in the configuration, it finds the last token that is not a padding token in each row. If
1436
+ no `pad_token_id` is defined, it simply takes the last value in each row of the batch. Since it cannot guess the
1437
+ padding tokens when `inputs_embeds` are passed instead of `input_ids`, it does the same (take the last value in
1438
+ each row of the batch).
1439
+ """,
1440
+ YUAN_START_DOCSTRING,
1441
+ )
1442
+ class YuanForSequenceClassification(YuanPreTrainedModel):
1443
+ #_keys_to_ignore_on_load_missing = [r"lm_head.weight"]
1444
+
1445
+ def __init__(self, config):
1446
+ super().__init__(config)
1447
+ self.num_labels = config.num_labels
1448
+ self.model = YuanModel(config)
1449
+ self.score = nn.Linear(config.hidden_size, self.num_labels, bias=False)
1450
+
1451
+ # Initialize weights and apply final processing
1452
+ self.post_init()
1453
+
1454
+ def get_input_embeddings(self):
1455
+ return self.model.embed_tokens
1456
+
1457
+ def set_input_embeddings(self, value):
1458
+ self.model.embed_tokens = value
1459
+
1460
+ @add_start_docstrings_to_model_forward(YUAN_INPUTS_DOCSTRING)
1461
+ def forward(
1462
+ self,
1463
+ input_ids: torch.LongTensor = None,
1464
+ attention_mask: Optional[torch.Tensor] = None,
1465
+ position_ids: Optional[torch.LongTensor] = None,
1466
+ past_key_values: Optional[List[torch.FloatTensor]] = None,
1467
+ inputs_embeds: Optional[torch.FloatTensor] = None,
1468
+ labels: Optional[torch.LongTensor] = None,
1469
+ use_cache: Optional[bool] = None,
1470
+ output_attentions: Optional[bool] = None,
1471
+ output_hidden_states: Optional[bool] = None,
1472
+ return_dict: Optional[bool] = None,
1473
+ ) -> Union[Tuple, SequenceClassifierOutputWithPast]:
1474
+ r"""
1475
+ labels (`torch.LongTensor` of shape `(batch_size,)`, *optional*):
1476
+ Labels for computing the sequence classification/regression loss. Indices should be in `[0, ...,
1477
+ config.num_labels - 1]`. If `config.num_labels == 1` a regression loss is computed (Mean-Square loss), If
1478
+ `config.num_labels > 1` a classification loss is computed (Cross-Entropy).
1479
+ """
1480
+ return_dict = return_dict if return_dict is not None else self.config.use_return_dict
1481
+ transformer_outputs = self.model(
1482
+ input_ids,
1483
+ attention_mask=attention_mask,
1484
+ position_ids=position_ids,
1485
+ past_key_values=past_key_values,
1486
+ inputs_embeds=inputs_embeds,
1487
+ use_cache=use_cache,
1488
+ output_attentions=output_attentions,
1489
+ output_hidden_states=output_hidden_states,
1490
+ return_dict=return_dict,
1491
+ )
1492
+ hidden_states = transformer_outputs[0]
1493
+ logits = self.score(hidden_states)
1494
+
1495
+ if input_ids is not None:
1496
+ batch_size = input_ids.shape[0]
1497
+ else:
1498
+ batch_size = inputs_embeds.shape[0]
1499
+
1500
+ if self.config.pad_token_id is None and batch_size != 1:
1501
+ raise ValueError("Cannot handle batch sizes > 1 if no padding token is defined.")
1502
+ if self.config.pad_token_id is None:
1503
+ sequence_lengths = -1
1504
+ else:
1505
+ if input_ids is not None:
1506
+ sequence_lengths = (torch.ne(input_ids, self.config.pad_token_id).sum(-1) - 1).to(logits.device)
1507
+ else:
1508
+ sequence_lengths = -1
1509
+
1510
+ pooled_logits = logits[torch.arange(batch_size, device=logits.device), sequence_lengths]
1511
+
1512
+ loss = None
1513
+ if labels is not None:
1514
+ labels = labels.to(logits.device)
1515
+ if self.config.problem_type is None:
1516
+ if self.num_labels == 1:
1517
+ self.config.problem_type = "regression"
1518
+ elif self.num_labels > 1 and (labels.dtype == torch.long or labels.dtype == torch.int):
1519
+ self.config.problem_type = "single_label_classification"
1520
+ else:
1521
+ self.config.problem_type = "multi_label_classification"
1522
+
1523
+ if self.config.problem_type == "regression":
1524
+ loss_fct = MSELoss()
1525
+ if self.num_labels == 1:
1526
+ loss = loss_fct(pooled_logits.squeeze(), labels.squeeze())
1527
+ else:
1528
+ loss = loss_fct(pooled_logits, labels)
1529
+ elif self.config.problem_type == "single_label_classification":
1530
+ loss_fct = CrossEntropyLoss()
1531
+ loss = loss_fct(pooled_logits.view(-1, self.num_labels), labels.view(-1))
1532
+ elif self.config.problem_type == "multi_label_classification":
1533
+ loss_fct = BCEWithLogitsLoss()
1534
+ loss = loss_fct(pooled_logits, labels)
1535
+ if not return_dict:
1536
+ output = (pooled_logits,) + transformer_outputs[1:]
1537
+ return ((loss,) + output) if loss is not None else output
1538
+
1539
+ return SequenceClassifierOutputWithPast(
1540
+ loss=loss,
1541
+ logits=pooled_logits,
1542
+ past_key_values=transformer_outputs.past_key_values,
1543
+ hidden_states=transformer_outputs.hidden_states,
1544
+ attentions=transformer_outputs.attentions,
1545
+ )
1546
+
1547
+
1548
+
modeling_yuanvl_chat.py ADDED
@@ -0,0 +1,400 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # --------------------------------------------------------
2
+ # YuanVL
3
+ # Copyright (c) 2024 OpenGVLab
4
+ # Licensed under The MIT License [see LICENSE for details]
5
+ # --------------------------------------------------------
6
+
7
+ import warnings
8
+ from typing import (Any, Callable, Iterable, List, Literal, Mapping, Optional,
9
+ Set, Tuple, Type, TypedDict, Union)
10
+
11
+ import torch.utils.checkpoint
12
+ import transformers
13
+ import torch
14
+ from torch import nn
15
+ from torch.nn import CrossEntropyLoss
16
+ from transformers import (AutoModel, GenerationConfig, LlamaForCausalLM,
17
+ LlamaTokenizer)
18
+ from transformers.modeling_outputs import CausalLMOutputWithPast
19
+ from transformers.modeling_utils import PreTrainedModel
20
+ from transformers.generation import GenerationMixin
21
+ from transformers.utils import ModelOutput, logging
22
+
23
+ #from transformer_engine.pytorch import RMSNorm
24
+ from transformers.activations import ACT2FN
25
+
26
+ from .configuration_yuanvl import YuanVLChatConfig
27
+ from .conversation import get_conv_template
28
+ from .modeling_intern_vit import InternVisionModel, has_flash_attn
29
+ from .modeling_yuanlm2 import YuanForCausalLM
30
+ from .utils import flatten_bn, merge_multimodal_embeddings
31
+
32
+ logger = logging.get_logger(__name__)
33
+
34
+ class RMSNorm(torch.nn.Module):
35
+ def __init__(self, hidden_size, eps=1e-6):
36
+ super().__init__()
37
+ self.weight = torch.nn.Parameter(torch.ones(hidden_size))
38
+ self.variance_epsilon = eps
39
+
40
+ def forward(self, hidden_states):
41
+ variance = hidden_states.to(torch.float32).pow(2).mean(-1, keepdim=True)
42
+ hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)
43
+
44
+ # convert into half-precision if necessary
45
+ if self.weight.dtype in [torch.float16, torch.bfloat16]:
46
+ hidden_states = hidden_states.to(self.weight.dtype)
47
+
48
+ return self.weight * hidden_states
49
+
50
+ class InternVLImagePixelInputs(TypedDict):
51
+ type: Literal["pixel_values"]
52
+ data: Union[torch.Tensor, List[torch.Tensor]]
53
+ """
54
+ Shape: `(batch_size, 1 + num_patches, num_channels, height, width)`
55
+
56
+ Note that `num_patches` may be different for each batch, in which case
57
+ the data is passed as a list instead of a batched tensor.
58
+ """
59
+ patches_per_image: List[int]
60
+ """
61
+ List of number of total patches for each image in the batch.
62
+ """
63
+
64
+
65
+ class InternVLImageEmbeddingInputs(TypedDict):
66
+ type: Literal["image_embeds"]
67
+ data: Any # in vllm vision this is a NestedTensors
68
+ """
69
+ A tensor of shape `(num_images, total_image_feature_size, hidden_size)`
70
+ or a list of tensors of shape `(total_image_feature_size, hidden_size)`
71
+
72
+ `hidden_size` must match the hidden size of language model backbone.
73
+ """
74
+
75
+
76
+ InternVLImageInputs = Union[InternVLImagePixelInputs,
77
+ InternVLImageEmbeddingInputs]
78
+
79
+
80
+ def version_cmp(v1, v2, op='eq'):
81
+ import operator
82
+
83
+ from packaging import version
84
+ op_func = getattr(operator, op)
85
+ return op_func(version.parse(v1), version.parse(v2))
86
+
87
+ class YuanImageMLP(nn.Module):
88
+
89
+ def __init__(
90
+ self,
91
+ hidden_size: int,
92
+ intermediate_size: int,
93
+ output_size: int,
94
+ hidden_act: str,
95
+ ) -> None:
96
+ super().__init__()
97
+ self.up_proj = nn.Linear(hidden_size, intermediate_size, bias=False)
98
+ self.gate_proj = nn.Linear(hidden_size, intermediate_size, bias=False)
99
+ self.down_proj = nn.Linear(intermediate_size, output_size, bias=False)
100
+
101
+ if hidden_act != "silu":
102
+ raise ValueError(f"Unsupported activation: {hidden_act}. Only silu is supported for now.")
103
+
104
+ self.act_fn = ACT2FN[hidden_act]
105
+
106
+ @torch.compile
107
+ def swiglu(self, y_1, y_2):
108
+ return self.act_fn(y_1) * y_2
109
+
110
+ def forward(self, x):
111
+ x1 = self.up_proj(x)
112
+ x2 = self.gate_proj(x)
113
+ x3 = self.swiglu(x1, x2)
114
+ x = self.down_proj(x3)
115
+ return x
116
+
117
+ class YuanVLChatModel(PreTrainedModel, GenerationMixin):
118
+ config_class = YuanVLChatConfig
119
+ main_input_name = 'pixel_values'
120
+ base_model_prefix = 'language_model'
121
+ _supports_flash_attn_2 = True
122
+ _no_split_modules = ['InternVisionModel', 'YuanDeocderLayer']
123
+
124
+ def __init__(self, config: YuanVLChatConfig, vision_model=None, language_model=None, use_flash_attn=True):
125
+ super().__init__(config)
126
+
127
+ assert version_cmp(transformers.__version__, '4.37.0', 'ge')
128
+ image_size = config.force_image_size or config.vision_config.image_size
129
+ patch_size = config.vision_config.patch_size
130
+ self.patch_size = patch_size
131
+ self.select_layer = config.select_layer
132
+ self.template = config.template
133
+ self.num_image_token = int((image_size // patch_size) ** 2 * (config.downsample_ratio ** 2))
134
+ self.downsample_ratio = config.downsample_ratio
135
+ self.ps_version = config.ps_version
136
+ use_flash_attn = use_flash_attn if has_flash_attn else False
137
+ config.vision_config.use_flash_attn = True if use_flash_attn else False
138
+ config.llm_config._attn_implementation = 'flash_attention_2' if use_flash_attn else 'eager'
139
+
140
+ logger.info(f'num_image_token: {self.num_image_token}')
141
+ logger.info(f'ps_version: {self.ps_version}')
142
+ if vision_model is not None:
143
+ self.vision_model = vision_model
144
+ else:
145
+ self.vision_model = InternVisionModel(config.vision_config)
146
+ if language_model is not None:
147
+ self.language_model = language_model
148
+ else:
149
+ if config.llm_config.architectures[0] == 'YuanForCausalLM':
150
+ self.language_model = YuanForCausalLM(config.llm_config)
151
+ else:
152
+ raise NotImplementedError(f'{config.llm_config.architectures[0]} is not implemented.')
153
+
154
+ self.pixel_unshuffle = torch.nn.PixelUnshuffle(downscale_factor=2)
155
+ layernorm_epsilon = config.llm_config.rms_norm_eps
156
+
157
+ self.imagemlp_input_hiddensize = int(config.vision_config.hidden_size / self.downsample_ratio ** 2)
158
+ self.imagemlp_ffn_hidden_size = config.llm_config.ffn_hidden_size
159
+
160
+ self.imagemlp = YuanImageMLP(self.imagemlp_input_hiddensize, self.imagemlp_ffn_hidden_size,
161
+ output_size=config.llm_config.hidden_size, hidden_act="silu")
162
+ self.imagemlp_layernorm = RMSNorm(config.llm_config.hidden_size, eps=layernorm_epsilon)
163
+
164
+ self.img_context_token_id = config.img_context_token_id
165
+ self.conv_template = get_conv_template(self.template)
166
+ self.system_message = self.conv_template.system_message
167
+
168
+ def _validate_pixel_values(self,
169
+ data: Union[torch.Tensor, List[torch.Tensor]]
170
+ ) -> Union[torch.Tensor, List[torch.Tensor]]:
171
+
172
+ h = w = self.config.vision_config.image_size
173
+ expected_dims = (3, h, w)
174
+
175
+ def _validate_shape(d: torch.Tensor):
176
+ actual_dims = tuple(d.shape)
177
+ if actual_dims != expected_dims:
178
+ # expected_expr = ("num_patches", *map(str, expected_dims))
179
+ expected_expr = (expected_dims)
180
+ raise ValueError("The expected shape of pixel values in each batch element "
181
+ f" is {expected_expr}. You supplied {tuple(d.shape)}.")
182
+ # data的数据类型可以是tensor,也可以是List[tensor]
183
+ # 从这一段上来看,image tensor的个数为 imbs*num_images
184
+ for d in data:
185
+ _validate_shape(d)
186
+ return data
187
+
188
+
189
+
190
+ def _parse_and_validate_image_input(self,
191
+ pixel_values: List[torch.Tensor] = None,
192
+ image_token_id: torch.Tensor = None,
193
+ image_embeds: torch.Tensor = None,
194
+ ) -> Optional[InternVLImagePixelInputs]:
195
+ # 没有图像数据
196
+ if pixel_values is None and image_embeds is None:
197
+ return None
198
+
199
+ # 传入数据有image_embeds
200
+ if image_embeds is not None:
201
+ if not isinstance(image_embeds, torch.Tensor):
202
+ raise ValueError("Incorrect type of image embeddings. "
203
+ f"Got type: {type(image_embeds)}")
204
+ return InternVLImageEmbeddingInputs(
205
+ type="image_embeds",
206
+ data=flatten_bn(image_embeds),
207
+ )
208
+
209
+ #self.img_context_token_id = image_token_id[0]
210
+ if pixel_values is not None:
211
+ if not isinstance(pixel_values, (torch.Tensor, list)):
212
+ raise ValueError("Incorrect type of pixel values. "
213
+ f"Got type: {type(pixel_values)}")
214
+ patches_per_image = []
215
+ # bsz/request循环
216
+ for request_pixel_values in pixel_values:
217
+ # 每个request的images循环
218
+ patches_per_image.append(request_pixel_values.shape[0])
219
+
220
+ # We need to flatten (B, N, P) to (B*N*P)
221
+ # so we call flatten_bn twice.
222
+ # (total_patches, 3, h, w)
223
+ return InternVLImagePixelInputs(
224
+ type="pixel_values",
225
+ data=self._validate_pixel_values(flatten_bn(pixel_values)),
226
+ patches_per_image=patches_per_image)
227
+ raise AssertionError("This line should be unreachable")
228
+
229
+ def _process_image_input(
230
+ self,
231
+ image_input: InternVLImageInputs,
232
+ ) -> Tuple[torch.Tensor] :
233
+ if image_input["type"] == "image_embeds":
234
+ return image_input["data"]
235
+ assert self.vision_model is not None
236
+ # (total_patches, tokens_per_image, llm_config.hidden_size)
237
+ image_embeds = self.extract_feature(image_input["data"])
238
+ patches_per_image = image_input["patches_per_image"]
239
+
240
+ # Only one image in the current batch
241
+ # bsz=1的情况,直接返回image_embeds
242
+ if len(patches_per_image) == 1:
243
+ # 返回一个tensor,[1, num_patches*256, text_config.hidden_size]
244
+ image_embeds = image_embeds.view(-1, self.config.llm_config.hidden_size).unsqueeze(1)
245
+ return image_embeds
246
+ # NOTE: Image embeddings are split into separate tensors for each image
247
+ # by the size of each embedding.
248
+ # feature_size 每个patch 256个token位置
249
+ feature_size = image_embeds.shape[1]
250
+ # (total_image_tokens, llm_config.hidden_size)
251
+ image_embeds = image_embeds.view(-1, self.config.llm_config.hidden_size)
252
+ image_feature_sizes = [num_patches * feature_size for num_patches in patches_per_image]
253
+ # 切分后得到一个Tuple,元组每个元胞表示一个image的image_embed, [num_patches * 256, llm_config.hidden_size]
254
+ image_embeds = image_embeds.split(image_feature_sizes)
255
+
256
+ return image_embeds
257
+
258
+
259
+
260
+ def get_multimodal_embeddings(self,
261
+ pixel_values: Optional[List[torch.Tensor]] = None,
262
+ image_token_id: Optional[List[torch.Tensor]] = None,
263
+ image_embeds: Optional[List[torch.Tensor]] = None,
264
+ image_input: InternVLImageInputs = None,
265
+ ):
266
+ image_input = self._parse_and_validate_image_input(pixel_values, image_token_id, image_embeds)
267
+ if image_input is None:
268
+ return None
269
+
270
+ # image_input: (total_patches, 3, h, w)
271
+ vision_embeddings = self._process_image_input(image_input)
272
+ return vision_embeddings
273
+
274
+ def get_input_embeddings(
275
+ self,
276
+ input_ids: torch.Tensor,
277
+ multimodal_embeddings: Optional[torch.Tensor]
278
+ ) -> torch.Tensor:
279
+ # 生成 token_embeddings
280
+ inputs_embeds = self.language_model.model.get_input_embeddings(input_ids)
281
+ # 将image embed放到img_context_token_id的位置
282
+ if multimodal_embeddings is not None:
283
+ assert self.img_context_token_id is not None
284
+ # input_ids: torch.Tensor
285
+ # inputs_embeds: torch.Tensor
286
+ # multimodal_embeddings: torch.Tensor
287
+ # placeholder_token_id: img_context_token_id
288
+ inputs_embeds = merge_multimodal_embeddings(
289
+ input_ids, inputs_embeds, multimodal_embeddings,
290
+ self.img_context_token_id)
291
+ return inputs_embeds
292
+
293
+ def forward(
294
+ self,
295
+ input_ids: torch.LongTensor = None,
296
+ attention_mask: torch.Tensor = None,
297
+ position_ids: torch.LongTensor = None,
298
+ past_key_values: List[torch.FloatTensor] = None,
299
+ inputs_embeds: Optional[torch.FloatTensor] = None,
300
+ labels: Optional[torch.LongTensor] = None,
301
+ use_cache: Optional[bool] = None,
302
+ output_attentions: Optional[bool] = None,
303
+ output_hidden_states: Optional[bool] = None,
304
+ return_dict: Optional[bool] = None,
305
+ pixel_values: Optional[List[torch.Tensor]] = None,
306
+ image_token_id: Optional[List[torch.Tensor]] = None,
307
+ image_embeds: Optional[List[torch.Tensor]] = None,
308
+ ) -> Union[Tuple, CausalLMOutputWithPast]:
309
+
310
+ if inputs_embeds is None:
311
+ # (images, patches * token_per_image)
312
+ vision_embeddings = self.get_multimodal_embeddings(pixel_values, image_token_id, image_embeds)
313
+ # (tokens, hidden_size)
314
+ if input_ids is not None:
315
+ vision_embeddings = vision_embeddings.to(input_ids.device)
316
+ inputs_embeds = self.get_input_embeddings(input_ids, vision_embeddings) #.permute(1, 0, 2)
317
+ input_ids = None
318
+
319
+ hidden_states = self.language_model.model(input_ids, attention_mask, position_ids, past_key_values,
320
+ inputs_embeds, labels, use_cache, output_attentions,
321
+ output_hidden_states, return_dict)
322
+ return hidden_states
323
+
324
+ def pixel_shuffle(self, x, scale_factor=0.5):
325
+ n, w, h, c = x.size()
326
+ # N, W, H, C --> N, W, H * scale, C // scale
327
+ x = x.view(n, w, int(h * scale_factor), int(c / scale_factor))
328
+ # N, W, H * scale, C // scale --> N, H * scale, W, C // scale
329
+ x = x.permute(0, 2, 1, 3).contiguous()
330
+ # N, H * scale, W, C // scale --> N, H * scale, W * scale, C // (scale ** 2)
331
+ x = x.view(n, int(h * scale_factor), int(w * scale_factor),
332
+ int(c / (scale_factor * scale_factor)))
333
+ if self.ps_version == 'v1':
334
+ warnings.warn("In ps_version 'v1', the height and width have not been swapped back, "
335
+ 'which results in a transposed image.')
336
+ else:
337
+ x = x.permute(0, 2, 1, 3).contiguous()
338
+ return x
339
+
340
+ # Internvl vision
341
+ def extract_feature(self, pixel_values):
342
+ # pixel_values: (imbs * num_image, ic, ih, iw)
343
+ pixel_values = pixel_values.to(torch.bfloat16)
344
+ output = self.vision_model(pixel_values=pixel_values)
345
+ vit_embeds=output[0]
346
+ # vit_embeds: (imbs * num_images, h*w, vit_dim)
347
+ vit_embeds = vit_embeds[:, 1:, :]
348
+
349
+ pn, phw, pc = vit_embeds.shape
350
+ ph = pw = int(phw**0.5)
351
+ vit_embeds = vit_embeds.view(pn, ph, pw, pc).permute(0, 3, 1, 2)
352
+ vit_embeds = self.pixel_unshuffle(vit_embeds)
353
+ pn, pc, ph, pw = vit_embeds.shape
354
+ vit_embeds = vit_embeds.view(pn, pc, ph * pw).permute(0, 2, 1)
355
+ num_images, cvs, chs = vit_embeds.shape
356
+ #_, cvs, chs = vit_embeds.shape
357
+ #assert self.imagemlp_ffn_hidden_size == chs
358
+ #vit_embeds = vit_embeds.contiguous().view(imbs, num_image * cvs, chs).permute(1, 0, 2).contiguous()
359
+ vit_embeds = vit_embeds.reshape(1, -1, vit_embeds.shape[-1]).permute(1, 0, 2)
360
+ vit_embeds = self.imagemlp(vit_embeds)
361
+ vit_embeds = self.imagemlp_layernorm(vit_embeds)
362
+ vit_embeds = vit_embeds.view(num_images, cvs, -1)
363
+ return vit_embeds
364
+
365
+ @torch.no_grad()
366
+ def generate(
367
+ self,
368
+ pixel_values: Optional[torch.FloatTensor] = None,
369
+ input_ids: Optional[torch.FloatTensor] = None,
370
+ attention_mask: Optional[torch.LongTensor] = None,
371
+ visual_features: Optional[torch.FloatTensor] = None,
372
+ generation_config: Optional[GenerationConfig] = None,
373
+ position_ids: Optional[torch.Tensor] = None,
374
+ output_hidden_states: Optional[bool] = None,
375
+ ) -> torch.LongTensor:
376
+
377
+
378
+ if pixel_values is not None:
379
+ if visual_features is not None:
380
+ vit_embeds = visual_features
381
+ else:
382
+ vit_embeds = self.get_multimodal_embeddings(pixel_values)
383
+ if input_ids is not None:
384
+ vit_embeds = vit_embeds.to(input_ids.device)
385
+ inputs_embeds = self.get_input_embeddings(input_ids, vit_embeds)
386
+ input_ids = None
387
+
388
+
389
+ outputs = self.language_model.generate(
390
+ inputs_embeds=inputs_embeds,
391
+ attention_mask=attention_mask,
392
+ generation_config=generation_config,
393
+ output_hidden_states=output_hidden_states,
394
+ position_ids=position_ids,
395
+ max_length=8192,
396
+ use_cache=True,
397
+ )
398
+
399
+
400
+ return outputs