x54-729 commited on
Commit
a80bd54
·
verified ·
1 Parent(s): 8f98ba2

Update modeling_internlm2.py

Browse files
Files changed (1) hide show
  1. modeling_internlm2.py +761 -352
modeling_internlm2.py CHANGED
@@ -13,11 +13,10 @@
13
  # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
14
  # See the License for the specific language governing permissions and
15
  # limitations under the License.
16
- """ PyTorch InternLM2 model."""
17
  import math
18
  import queue
19
  import threading
20
- import warnings
21
  from typing import List, Optional, Tuple, Union
22
 
23
  import torch
@@ -27,49 +26,50 @@ from einops import rearrange
27
  from torch import nn
28
  from torch.nn import BCEWithLogitsLoss, CrossEntropyLoss, MSELoss
29
  from transformers.activations import ACT2FN
 
 
30
  from transformers.modeling_outputs import (
31
  BaseModelOutputWithPast,
32
  CausalLMOutputWithPast,
 
33
  SequenceClassifierOutputWithPast,
 
34
  )
35
  from transformers.modeling_utils import PreTrainedModel
 
36
  from transformers.utils import (
37
  add_start_docstrings,
38
  add_start_docstrings_to_model_forward,
 
39
  logging,
40
  replace_return_docstrings,
41
  )
42
 
43
  try:
44
  from transformers.generation.streamers import BaseStreamer
45
- except: # noqa # pylint: disable=bare-except
46
  BaseStreamer = None
47
 
48
  from .configuration_internlm2 import InternLM2Config
49
 
 
 
 
 
 
 
 
 
50
  logger = logging.get_logger(__name__)
51
 
52
  _CONFIG_FOR_DOC = "InternLM2Config"
53
 
54
- flash_attn_func, flash_attn_varlen_func = None, None
55
- pad_input, index_first_axis, unpad_input = None, None, None
56
- def _import_flash_attn():
57
- global flash_attn_func, flash_attn_varlen_func
58
- global pad_input, index_first_axis, unpad_input
59
- try:
60
- from flash_attn import flash_attn_func as _flash_attn_func, flash_attn_varlen_func as _flash_attn_varlen_func
61
- from flash_attn.bert_padding import pad_input as _pad_input, index_first_axis as _index_first_axis, unpad_input as _unpad_input
62
- flash_attn_func, flash_attn_varlen_func = _flash_attn_func, _flash_attn_varlen_func
63
- pad_input, index_first_axis, unpad_input = _pad_input, _index_first_axis, _unpad_input
64
- except ImportError:
65
- raise ImportError("flash_attn is not installed.")
66
-
67
- # Copied from transformers.models.llama.modeling_llama._get_unpad_data
68
  def _get_unpad_data(attention_mask):
69
  seqlens_in_batch = attention_mask.sum(dim=-1, dtype=torch.int32)
70
  indices = torch.nonzero(attention_mask.flatten(), as_tuple=False).flatten()
71
  max_seqlen_in_batch = seqlens_in_batch.max().item()
72
- cu_seqlens = F.pad(torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.torch.int32), (1, 0))
73
  return (
74
  indices,
75
  cu_seqlens,
@@ -77,45 +77,10 @@ def _get_unpad_data(attention_mask):
77
  )
78
 
79
 
80
- # Copied from transformers.models.bart.modeling_bart._make_causal_mask
81
- def _make_causal_mask(
82
- input_ids_shape: torch.Size, dtype: torch.dtype, device: torch.device, past_key_values_length: int = 0
83
- ):
84
- """
85
- Make causal mask used for bi-directional self-attention.
86
- """
87
- bsz, tgt_len = input_ids_shape
88
- mask = torch.full((tgt_len, tgt_len), torch.tensor(torch.finfo(dtype).min, device=device), device=device)
89
- mask_cond = torch.arange(mask.size(-1), device=device)
90
- mask.masked_fill_(mask_cond < (mask_cond + 1).view(mask.size(-1), 1), 0)
91
- mask = mask.to(dtype)
92
-
93
- if past_key_values_length > 0:
94
- mask = torch.cat([torch.zeros(tgt_len, past_key_values_length, dtype=dtype, device=device), mask], dim=-1)
95
- return mask[None, None, :, :].expand(bsz, 1, tgt_len, tgt_len + past_key_values_length)
96
-
97
-
98
- # Copied from transformers.models.bart.modeling_bart._expand_mask
99
- def _expand_mask(mask: torch.Tensor, dtype: torch.dtype, tgt_len: Optional[int] = None):
100
- """
101
- Expands attention_mask from `[bsz, seq_len]` to `[bsz, 1, tgt_seq_len, src_seq_len]`.
102
- """
103
- bsz, src_len = mask.size()
104
- tgt_len = tgt_len if tgt_len is not None else src_len
105
-
106
- expanded_mask = mask[:, None, None, :].expand(bsz, 1, tgt_len, src_len).to(dtype)
107
-
108
- inverted_mask = 1.0 - expanded_mask
109
-
110
- return inverted_mask.masked_fill(inverted_mask.to(torch.bool), torch.finfo(dtype).min)
111
-
112
-
113
- # Copied from transformers.models.llama.modeling_llama.LlamaRMSNorm with Llama->InternLM2
114
  class InternLM2RMSNorm(nn.Module):
 
 
115
  def __init__(self, hidden_size, eps=1e-6):
116
- """
117
- InternLM2RMSNorm is equivalent to T5LayerNorm
118
- """
119
  super().__init__()
120
  self.weight = nn.Parameter(torch.ones(hidden_size))
121
  self.variance_epsilon = eps
@@ -128,93 +93,68 @@ class InternLM2RMSNorm(nn.Module):
128
  return self.weight * hidden_states.to(input_dtype)
129
 
130
 
131
- # Copied from transformers.model.llama.modeling_llama.LlamaRotaryEmbedding with Llama->InternLM2
 
 
132
  class InternLM2RotaryEmbedding(nn.Module):
133
- def __init__(self, dim, max_position_embeddings=2048, base=10000, device=None):
134
- super().__init__()
135
 
 
 
 
136
  self.dim = dim
137
  self.max_position_embeddings = max_position_embeddings
138
  self.base = base
139
- inv_freq = 1.0 / (self.base ** (torch.arange(0, self.dim, 2).float().to(device) / self.dim))
140
  self.register_buffer("inv_freq", inv_freq, persistent=False)
 
 
141
 
142
- # Build here to make `torch.jit.trace` work.
143
- self._set_cos_sin_cache(
144
- seq_len=max_position_embeddings, device=self.inv_freq.device, dtype=torch.get_default_dtype()
145
- )
146
-
147
- def _set_cos_sin_cache(self, seq_len, device, dtype):
148
- self.max_seq_len_cached = seq_len
149
- t = torch.arange(self.max_seq_len_cached, device=device, dtype=self.inv_freq.dtype)
150
-
151
- freqs = torch.einsum("i,j->ij", t, self.inv_freq)
152
- # Different from paper, but it uses a different permutation in order to obtain the same calculation
153
- emb = torch.cat((freqs, freqs), dim=-1)
154
- self.register_buffer("cos_cached", emb.cos().to(dtype), persistent=False)
155
- self.register_buffer("sin_cached", emb.sin().to(dtype), persistent=False)
156
-
157
- def forward(self, x, seq_len=None):
158
  # x: [bs, num_attention_heads, seq_len, head_size]
159
- if seq_len > self.max_seq_len_cached:
160
- self._set_cos_sin_cache(seq_len=seq_len, device=x.device, dtype=torch.float32)
 
 
 
 
 
 
 
 
 
 
161
 
162
- return (
163
- self.cos_cached[:seq_len].to(dtype=x.dtype),
164
- self.sin_cached[:seq_len].to(dtype=x.dtype),
165
- )
166
 
167
-
168
- # Copied from transformers.model.llama.modeling_llama.LlamaLinearScalingRotaryEmbedding with Llama->InternLM2
169
  class InternLM2LinearScalingRotaryEmbedding(InternLM2RotaryEmbedding):
170
  """InternLM2RotaryEmbedding extended with linear scaling. Credits to the Reddit user /u/kaiokendev"""
171
 
172
- def __init__(self, dim, max_position_embeddings=2048, base=10000, device=None, scaling_factor=1.0):
173
- self.scaling_factor = scaling_factor
174
- super().__init__(dim, max_position_embeddings, base, device)
175
-
176
- def _set_cos_sin_cache(self, seq_len, device, dtype):
177
- self.max_seq_len_cached = seq_len
178
- t = torch.arange(self.max_seq_len_cached, device=device, dtype=self.inv_freq.dtype)
179
- t = t / self.scaling_factor
180
-
181
- freqs = torch.einsum("i,j->ij", t, self.inv_freq)
182
- # Different from paper, but it uses a different permutation in order to obtain the same calculation
183
- emb = torch.cat((freqs, freqs), dim=-1)
184
- self.register_buffer("cos_cached", emb.cos().to(dtype), persistent=False)
185
- self.register_buffer("sin_cached", emb.sin().to(dtype), persistent=False)
186
 
187
 
188
- # Copied from transformers.model.llama.modeling_llama.LlamaDynamicNTKScalingRotaryEmbedding with Llama->InternLM2
189
  class InternLM2DynamicNTKScalingRotaryEmbedding(InternLM2RotaryEmbedding):
190
  """InternLM2RotaryEmbedding extended with Dynamic NTK scaling.
191
- Credits to the Reddit users /u/bloc97 and /u/emozilla.
192
- """
193
-
194
- def __init__(self, dim, max_position_embeddings=2048, base=10000, device=None, scaling_factor=1.0):
195
- self.scaling_factor = scaling_factor
196
- super().__init__(dim, max_position_embeddings, base, device)
197
-
198
- def _set_cos_sin_cache(self, seq_len, device, dtype):
199
- self.max_seq_len_cached = seq_len
200
 
 
 
 
201
  if seq_len > self.max_position_embeddings:
202
  base = self.base * (
203
  (self.scaling_factor * seq_len / self.max_position_embeddings) - (self.scaling_factor - 1)
204
  ) ** (self.dim / (self.dim - 2))
205
- inv_freq = 1.0 / (base ** (torch.arange(0, self.dim, 2).float().to(device) / self.dim))
206
- self.register_buffer("inv_freq", inv_freq, persistent=False)
207
 
208
- t = torch.arange(self.max_seq_len_cached, device=device, dtype=self.inv_freq.dtype)
 
209
 
210
- freqs = torch.einsum("i,j->ij", t, self.inv_freq)
211
- # Different from paper, but it uses a different permutation in order to obtain the same calculation
212
- emb = torch.cat((freqs, freqs), dim=-1)
213
- self.register_buffer("cos_cached", emb.cos().to(dtype), persistent=False)
214
- self.register_buffer("sin_cached", emb.sin().to(dtype), persistent=False)
215
 
216
-
217
- # Copied from transformers.model.llama.modeling_llama.rotate_half
218
  def rotate_half(x):
219
  """Rotates half the hidden dims of the input."""
220
  x1 = x[..., : x.shape[-1] // 2]
@@ -222,17 +162,36 @@ def rotate_half(x):
222
  return torch.cat((-x2, x1), dim=-1)
223
 
224
 
225
- # Copied from transformers.model.llama.modeling_llama.apply_rotary_pos_emb
226
- def apply_rotary_pos_emb(q, k, cos, sin, position_ids, unsqueeze_dim=1):
227
- """Applies Rotary Position Embedding to the query and key tensors."""
228
- cos = cos[position_ids].unsqueeze(unsqueeze_dim)
229
- sin = sin[position_ids].unsqueeze(unsqueeze_dim)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
230
  q_embed = (q * cos) + (rotate_half(q) * sin)
231
  k_embed = (k * cos) + (rotate_half(k) * sin)
232
  return q_embed, k_embed
233
 
234
 
235
  class InternLM2MLP(nn.Module):
 
 
236
  def __init__(self, config):
237
  super().__init__()
238
  self.config = config
@@ -249,7 +208,6 @@ class InternLM2MLP(nn.Module):
249
  return down_proj
250
 
251
 
252
- # Copied from transformers.model.llama.modeling_llama.repeat_kv
253
  def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor:
254
  """
255
  This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch,
@@ -262,19 +220,27 @@ def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor:
262
  return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim)
263
 
264
 
265
- # Modified from transformers.model.llama.modeling_llama.LlamaAttention
266
  class InternLM2Attention(nn.Module):
267
  """Multi-headed attention from 'Attention Is All You Need' paper"""
268
 
269
- def __init__(self, config: InternLM2Config):
270
  super().__init__()
271
  self.config = config
 
 
 
 
 
 
 
 
272
  self.hidden_size = config.hidden_size
273
  self.num_heads = config.num_attention_heads
274
  self.head_dim = self.hidden_size // self.num_heads
275
  self.num_key_value_heads = config.num_key_value_heads
276
  self.num_key_value_groups = self.num_heads // self.num_key_value_heads
277
  self.max_position_embeddings = config.max_position_embeddings
 
278
  self.is_causal = True
279
 
280
  if (self.head_dim * self.num_heads) != self.hidden_size:
@@ -288,8 +254,8 @@ class InternLM2Attention(nn.Module):
288
  (self.num_heads + 2 * self.num_key_value_heads) * self.head_dim,
289
  bias=config.bias,
290
  )
291
-
292
  self.wo = nn.Linear(self.num_heads * self.head_dim, self.hidden_size, bias=config.bias)
 
293
  self._init_rope()
294
 
295
  def _init_rope(self):
@@ -297,51 +263,49 @@ class InternLM2Attention(nn.Module):
297
  self.rotary_emb = InternLM2RotaryEmbedding(
298
  self.head_dim,
299
  max_position_embeddings=self.max_position_embeddings,
300
- base=self.config.rope_theta,
301
  )
302
  else:
303
  scaling_type = self.config.rope_scaling["type"]
304
  scaling_factor = self.config.rope_scaling["factor"]
305
- if scaling_type == "dynamic":
306
- self.rotary_emb = InternLM2DynamicNTKScalingRotaryEmbedding(
307
  self.head_dim,
308
  max_position_embeddings=self.max_position_embeddings,
309
- base=self.config.rope_theta,
310
  scaling_factor=scaling_factor,
 
311
  )
312
- elif scaling_type == "linear":
313
- self.rotary_emb = InternLM2LinearScalingRotaryEmbedding(
314
  self.head_dim,
315
  max_position_embeddings=self.max_position_embeddings,
316
- base=self.config.rope_theta,
317
  scaling_factor=scaling_factor,
 
318
  )
319
  else:
320
- raise ValueError("Currently we only support rotary embedding's type being 'dynamic' or 'linear'.")
321
- return self.rotary_emb
322
-
323
- def _shape(self, tensor: torch.Tensor, seq_len: int, bsz: int):
324
- return tensor.view(bsz, seq_len, self.num_heads, self.head_dim).transpose(1, 2).contiguous()
325
 
326
  def forward(
327
  self,
328
  hidden_states: torch.Tensor,
329
  attention_mask: Optional[torch.Tensor] = None,
330
  position_ids: Optional[torch.LongTensor] = None,
331
- past_key_value: Optional[Tuple[torch.Tensor]] = None,
332
  output_attentions: bool = False,
333
- use_cache: bool = False,
334
- **kwargs,
335
  ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]:
336
- if "padding_mask" in kwargs:
337
- warnings.warn(
338
- "Passing `padding_mask` is deprecated and will be removed in v4.37. "
339
- "Please make sure use `attention_mask` instead.`"
340
- )
341
-
342
  bsz, q_len, _ = hidden_states.size()
343
 
344
- qkv_states = self.wqkv(hidden_states)
 
 
 
 
 
 
 
 
345
 
346
  qkv_states = rearrange(
347
  qkv_states,
@@ -351,44 +315,26 @@ class InternLM2Attention(nn.Module):
351
  )
352
 
353
  query_states = qkv_states[..., : self.num_key_value_groups, :]
354
- query_states = rearrange(query_states, "b q h gs d -> b q (h gs) d")
355
- key_states = qkv_states[..., -2, :]
356
- value_states = qkv_states[..., -1, :]
357
 
358
- query_states = query_states.transpose(1, 2)
359
- key_states = key_states.transpose(1, 2)
360
- value_states = value_states.transpose(1, 2)
361
-
362
- kv_seq_len = key_states.shape[-2]
363
- if past_key_value is not None:
364
- kv_seq_len += past_key_value[0].shape[-2]
365
- cos, sin = self.rotary_emb(value_states, seq_len=kv_seq_len)
366
  query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin, position_ids)
367
 
368
  if past_key_value is not None:
369
- # reuse k, v, self_attention
370
- key_states = torch.cat([past_key_value[0], key_states], dim=2)
371
- value_states = torch.cat([past_key_value[1], value_states], dim=2)
372
-
373
- past_key_value = (key_states, value_states) if use_cache else None
374
 
375
  key_states = repeat_kv(key_states, self.num_key_value_groups)
376
  value_states = repeat_kv(value_states, self.num_key_value_groups)
377
 
378
  attn_weights = torch.matmul(query_states, key_states.transpose(2, 3)) / math.sqrt(self.head_dim)
379
 
380
- if attn_weights.size() != (bsz, self.num_heads, q_len, kv_seq_len):
381
- raise ValueError(
382
- f"Attention weights should be of size {(bsz, self.num_heads, q_len, kv_seq_len)}, but is"
383
- f" {attn_weights.size()}"
384
- )
385
-
386
- if attention_mask is not None:
387
- if attention_mask.size() != (bsz, 1, q_len, kv_seq_len):
388
- raise ValueError(
389
- f"Attention mask should be of size {(bsz, 1, q_len, kv_seq_len)}, but is {attention_mask.size()}"
390
- )
391
- attn_weights = attn_weights + attention_mask
392
 
393
  # upcast attention to fp32
394
  attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query_states.dtype)
@@ -401,9 +347,20 @@ class InternLM2Attention(nn.Module):
401
  )
402
 
403
  attn_output = attn_output.transpose(1, 2).contiguous()
 
404
  attn_output = attn_output.reshape(bsz, q_len, self.hidden_size)
405
 
406
- attn_output = self.wo(attn_output)
 
 
 
 
 
 
 
 
 
 
407
 
408
  if not output_attentions:
409
  attn_weights = None
@@ -411,7 +368,6 @@ class InternLM2Attention(nn.Module):
411
  return attn_output, attn_weights, past_key_value
412
 
413
 
414
- # Modified from transformers.model.llama.modeling_llama.InternLM2FlashAttention2
415
  class InternLM2FlashAttention2(InternLM2Attention):
416
  """
417
  InternLM2 flash attention module. This module inherits from `InternLM2Attention` as the weights of the module stays
@@ -419,26 +375,34 @@ class InternLM2FlashAttention2(InternLM2Attention):
419
  flash attention and deal with padding tokens in case the input contains any of them.
420
  """
421
 
 
 
 
 
 
 
 
 
 
 
 
422
  def forward(
423
  self,
424
  hidden_states: torch.Tensor,
425
  attention_mask: Optional[torch.LongTensor] = None,
426
  position_ids: Optional[torch.LongTensor] = None,
427
- past_key_value: Optional[Tuple[torch.Tensor]] = None,
428
  output_attentions: bool = False,
429
  use_cache: bool = False,
430
- **kwargs,
431
  ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]:
432
- # InternLM2FlashAttention2 attention does not support output_attentions
433
- if "padding_mask" in kwargs:
434
- warnings.warn(
435
- "Passing `padding_mask` is deprecated and will be removed in v4.37. "
436
- "Please make sure use `attention_mask` instead.`"
437
  )
438
 
439
- # overwrite attention_mask with padding_mask
440
- attention_mask = kwargs.pop("padding_mask")
441
-
442
  output_attentions = False
443
 
444
  bsz, q_len, _ = hidden_states.size()
@@ -461,35 +425,61 @@ class InternLM2FlashAttention2(InternLM2Attention):
461
  key_states = key_states.transpose(1, 2)
462
  value_states = value_states.transpose(1, 2)
463
 
464
- kv_seq_len = key_states.shape[-2]
465
- if past_key_value is not None:
466
- kv_seq_len += past_key_value[0].shape[-2]
467
-
468
- cos, sin = self.rotary_emb(value_states, seq_len=kv_seq_len)
469
-
470
- query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin, position_ids)
471
 
472
  if past_key_value is not None:
473
- # reuse k, v, self_attention
474
- key_states = torch.cat([past_key_value[0], key_states], dim=2)
475
- value_states = torch.cat([past_key_value[1], value_states], dim=2)
476
-
477
- past_key_value = (key_states, value_states) if use_cache else None
478
 
 
 
 
479
  query_states = query_states.transpose(1, 2)
480
  key_states = key_states.transpose(1, 2)
481
  value_states = value_states.transpose(1, 2)
482
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
483
  attn_output = self._flash_attention_forward(
484
- query_states, key_states, value_states, attention_mask, q_len
485
  )
 
486
  attn_output = attn_output.reshape(bsz, q_len, self.hidden_size).contiguous()
487
  attn_output = self.wo(attn_output)
488
 
489
  if not output_attentions:
490
  attn_weights = None
491
 
492
- return attn_output, attn_weights, past_key_value
493
 
494
  def _flash_attention_forward(
495
  self, query_states, key_states, value_states, attention_mask, query_length, dropout=0.0, softmax_scale=None
@@ -508,23 +498,29 @@ class InternLM2FlashAttention2(InternLM2Attention):
508
  attention_mask (`torch.Tensor`):
509
  The padding mask - corresponds to a tensor of size `(batch_size, seq_len)` where 0 stands for the
510
  position of padding tokens and 1 for the position of non-padding tokens.
511
- dropout (`int`, *optional*):
512
  Attention dropout
513
  softmax_scale (`float`, *optional*):
514
  The scaling of QK^T before applying softmax. Default to 1 / sqrt(head_dim)
515
  """
 
 
 
 
 
 
 
516
  # Contains at least one padding token in the sequence
517
- causal = self.is_causal and query_length != 1
518
  if attention_mask is not None:
519
  batch_size = query_states.shape[0]
520
- query_states, key_states, value_states, indices_q, cu_seq_lens, max_seq_lens = self._unpad_input(
521
  query_states, key_states, value_states, attention_mask, query_length
522
  )
523
 
524
  cu_seqlens_q, cu_seqlens_k = cu_seq_lens
525
  max_seqlen_in_batch_q, max_seqlen_in_batch_k = max_seq_lens
526
 
527
- attn_output_unpad = flash_attn_varlen_func(
528
  query_states,
529
  key_states,
530
  value_states,
@@ -537,27 +533,26 @@ class InternLM2FlashAttention2(InternLM2Attention):
537
  causal=causal,
538
  )
539
 
540
- attn_output = pad_input(attn_output_unpad, indices_q, batch_size, query_length)
541
  else:
542
- attn_output = flash_attn_func(
543
  query_states, key_states, value_states, dropout, softmax_scale=softmax_scale, causal=causal
544
  )
545
 
546
  return attn_output
547
 
548
- def _unpad_input(self, query_layer, key_layer, value_layer, attention_mask, query_length):
549
  indices_k, cu_seqlens_k, max_seqlen_in_batch_k = _get_unpad_data(attention_mask)
550
  batch_size, kv_seq_len, num_key_value_heads, head_dim = key_layer.shape
551
 
552
- key_layer = index_first_axis(
553
  key_layer.reshape(batch_size * kv_seq_len, num_key_value_heads, head_dim), indices_k
554
  )
555
- value_layer = index_first_axis(
556
  value_layer.reshape(batch_size * kv_seq_len, num_key_value_heads, head_dim), indices_k
557
  )
558
-
559
  if query_length == kv_seq_len:
560
- query_layer = index_first_axis(
561
  query_layer.reshape(batch_size * kv_seq_len, self.num_heads, head_dim), indices_k
562
  )
563
  cu_seqlens_q = cu_seqlens_k
@@ -573,29 +568,139 @@ class InternLM2FlashAttention2(InternLM2Attention):
573
  else:
574
  # The -q_len: slice assumes left padding.
575
  attention_mask = attention_mask[:, -query_length:]
576
- query_layer, indices_q, cu_seqlens_q, max_seqlen_in_batch_q = unpad_input(query_layer, attention_mask)
 
 
577
 
578
  return (
579
  query_layer,
580
  key_layer,
581
  value_layer,
582
- indices_q.to(torch.int64),
583
  (cu_seqlens_q, cu_seqlens_k),
584
  (max_seqlen_in_batch_q, max_seqlen_in_batch_k),
585
  )
586
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
587
  INTERNLM2_ATTENTION_CLASSES = {
588
  "eager": InternLM2Attention,
589
  "flash_attention_2": InternLM2FlashAttention2,
 
590
  }
591
 
592
- # Modified from transformers.model.llama.modeling_llama.LlamaDecoderLayer
 
593
  class InternLM2DecoderLayer(nn.Module):
594
- def __init__(self, config: InternLM2Config):
 
 
595
  super().__init__()
596
  self.hidden_size = config.hidden_size
 
597
 
598
- self.attention = INTERNLM2_ATTENTION_CLASSES[config.attn_implementation](config=config)
599
 
600
  self.feed_forward = InternLM2MLP(config)
601
  self.attention_norm = InternLM2RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
@@ -606,10 +711,10 @@ class InternLM2DecoderLayer(nn.Module):
606
  hidden_states: torch.Tensor,
607
  attention_mask: Optional[torch.Tensor] = None,
608
  position_ids: Optional[torch.LongTensor] = None,
609
- past_key_value: Optional[Tuple[torch.Tensor]] = None,
610
  output_attentions: Optional[bool] = False,
611
  use_cache: Optional[bool] = False,
612
- **kwargs,
613
  ) -> Tuple[torch.FloatTensor, Optional[Tuple[torch.FloatTensor, torch.FloatTensor]]]:
614
  """
615
  Args:
@@ -625,12 +730,6 @@ class InternLM2DecoderLayer(nn.Module):
625
  (see `past_key_values`).
626
  past_key_value (`Tuple(torch.FloatTensor)`, *optional*): cached past key and value projection states
627
  """
628
- if "padding_mask" in kwargs:
629
- warnings.warn(
630
- "Passing `padding_mask` is deprecated and will be removed in v4.37. "
631
- "Please make sure use `attention_mask` instead.`"
632
- )
633
-
634
  residual = hidden_states
635
 
636
  hidden_states = self.attention_norm(hidden_states)
@@ -643,7 +742,7 @@ class InternLM2DecoderLayer(nn.Module):
643
  past_key_value=past_key_value,
644
  output_attentions=output_attentions,
645
  use_cache=use_cache,
646
- **kwargs,
647
  )
648
  hidden_states = residual + hidden_states
649
 
@@ -687,11 +786,20 @@ InternLM2_START_DOCSTRING = r"""
687
  InternLM2_START_DOCSTRING,
688
  )
689
  class InternLM2PreTrainedModel(PreTrainedModel):
 
 
 
 
690
  config_class = InternLM2Config
691
  base_model_prefix = "model"
692
  supports_gradient_checkpointing = True
693
  _no_split_modules = ["InternLM2DecoderLayer"]
694
- _skip_keys_device_placement = "past_key_values"
 
 
 
 
 
695
 
696
  def _init_weights(self, module):
697
  std = self.config.initializer_range
@@ -740,14 +848,19 @@ InternLM2_INPUTS_DOCSTRING = r"""
740
  config.n_positions - 1]`.
741
 
742
  [What are position IDs?](../glossary#position-ids)
743
- past_key_values (`tuple(tuple(torch.FloatTensor))`, *optional*, returned when `use_cache=True` is passed or
744
- when `config.use_cache=True`):
745
- Tuple of `tuple(torch.FloatTensor)` of length `config.n_layers`, with each tuple having 2 tensors of shape
746
- `(batch_size, num_heads, sequence_length, embed_size_per_head)`) and 2 additional tensors of shape
747
- `(batch_size, num_heads, decoder_sequence_length, embed_size_per_head)`.
 
 
 
 
 
748
 
749
- Contains pre-computed hidden-states (key and values in the self-attention blocks and in the cross-attention
750
- blocks) that can be used (see `past_key_values` input) to speed up sequential decoding.
751
 
752
  If `past_key_values` are used, the user can optionally input only the last `input_ids` (those that don't
753
  have their past key value states given to this model) of shape `(batch_size, 1)` instead of all `input_ids`
@@ -767,10 +880,14 @@ InternLM2_INPUTS_DOCSTRING = r"""
767
  more detail.
768
  return_dict (`bool`, *optional*):
769
  Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.
 
 
 
 
770
  """
771
 
772
 
773
- # Modified from transformers.model.llama.modeling_llama.LlamaModel
774
  @add_start_docstrings(
775
  "The bare InternLM2 Model outputting raw hidden-states without any specific head on top.",
776
  InternLM2_START_DOCSTRING,
@@ -793,7 +910,9 @@ class InternLM2Model(InternLM2PreTrainedModel):
793
 
794
  self.tok_embeddings = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx)
795
 
796
- self.layers = nn.ModuleList([InternLM2DecoderLayer(config) for _ in range(config.num_hidden_layers)])
 
 
797
  self.norm = InternLM2RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
798
 
799
  self.gradient_checkpointing = False
@@ -806,142 +925,96 @@ class InternLM2Model(InternLM2PreTrainedModel):
806
  def set_input_embeddings(self, value):
807
  self.tok_embeddings = value
808
 
809
- def _prepare_decoder_attention_mask(self, attention_mask, input_shape, inputs_embeds, past_key_values_length):
810
- # create causal mask
811
- # [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]
812
- combined_attention_mask = None
813
- if input_shape[-1] > 1:
814
- combined_attention_mask = _make_causal_mask(
815
- input_shape,
816
- inputs_embeds.dtype,
817
- device=inputs_embeds.device,
818
- past_key_values_length=past_key_values_length,
819
- )
820
-
821
- if attention_mask is not None:
822
- # [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]
823
- expanded_attn_mask = _expand_mask(attention_mask, inputs_embeds.dtype, tgt_len=input_shape[-1]).to(
824
- inputs_embeds.device
825
- )
826
- combined_attention_mask = (
827
- expanded_attn_mask if combined_attention_mask is None else expanded_attn_mask + combined_attention_mask
828
- )
829
-
830
- return combined_attention_mask
831
-
832
  @add_start_docstrings_to_model_forward(InternLM2_INPUTS_DOCSTRING)
833
  def forward(
834
  self,
835
  input_ids: torch.LongTensor = None,
836
  attention_mask: Optional[torch.Tensor] = None,
837
  position_ids: Optional[torch.LongTensor] = None,
838
- past_key_values: Optional[List[torch.FloatTensor]] = None,
839
  inputs_embeds: Optional[torch.FloatTensor] = None,
840
  use_cache: Optional[bool] = None,
841
  output_attentions: Optional[bool] = None,
842
  output_hidden_states: Optional[bool] = None,
843
  return_dict: Optional[bool] = None,
 
844
  ) -> Union[Tuple, BaseModelOutputWithPast]:
845
  output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
846
  output_hidden_states = (
847
  output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
848
  )
849
  use_cache = use_cache if use_cache is not None else self.config.use_cache
850
-
851
  return_dict = return_dict if return_dict is not None else self.config.use_return_dict
852
 
853
- if self.config.attn_implementation == "flash_attention_2":
854
- _import_flash_attn()
855
-
856
- # retrieve input_ids and inputs_embeds
857
- if input_ids is not None and inputs_embeds is not None:
858
- raise ValueError("You cannot specify both input_ids and inputs_embeds at the same time")
859
- elif input_ids is not None:
860
- batch_size, seq_length = input_ids.shape[:2]
861
- elif inputs_embeds is not None:
862
- batch_size, seq_length = inputs_embeds.shape[:2]
863
- else:
864
- raise ValueError("You have to specify either input_ids or inputs_embeds")
865
-
866
- seq_length_with_past = seq_length
867
- past_key_values_length = 0
868
- if past_key_values is not None:
869
- past_key_values_length = past_key_values[0][0].shape[2]
870
- seq_length_with_past = seq_length_with_past + past_key_values_length
871
 
872
- if position_ids is None:
873
- device = input_ids.device if input_ids is not None else inputs_embeds.device
874
- position_ids = torch.arange(
875
- past_key_values_length, seq_length + past_key_values_length, dtype=torch.long, device=device
876
  )
877
- position_ids = position_ids.unsqueeze(0)
878
 
879
  if inputs_embeds is None:
880
  inputs_embeds = self.tok_embeddings(input_ids)
881
 
882
- if self.config.attn_implementation == "flash_attention_2":
883
- # 2d mask is passed through the layers
884
- attention_mask = attention_mask if (attention_mask is not None and 0 in attention_mask) else None
885
- else:
886
- if attention_mask is None:
887
- attention_mask = torch.ones(
888
- (batch_size, seq_length_with_past), dtype=torch.bool, device=inputs_embeds.device
889
- )
890
- attention_mask = self._prepare_decoder_attention_mask(
891
- attention_mask, (batch_size, seq_length), inputs_embeds, past_key_values_length
892
  )
 
 
 
 
 
 
893
 
894
  # embed positions
895
  hidden_states = inputs_embeds
896
 
897
- if self.gradient_checkpointing and self.training:
898
- if use_cache:
899
- logger.warning_once(
900
- "`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`..."
901
- )
902
- use_cache = False
903
-
904
  # decoder layers
905
  all_hidden_states = () if output_hidden_states else None
906
  all_self_attns = () if output_attentions else None
907
- next_decoder_cache = () if use_cache else None
908
 
909
- for idx, decoder_layer in enumerate(self.layers):
910
  if output_hidden_states:
911
  all_hidden_states += (hidden_states,)
912
 
913
- past_key_value = past_key_values[idx] if past_key_values is not None else None
914
-
915
  if self.gradient_checkpointing and self.training:
916
-
917
- def create_custom_forward(module):
918
- def custom_forward(*inputs):
919
- # None for past_key_value
920
- return module(*inputs, output_attentions, None)
921
-
922
- return custom_forward
923
-
924
- layer_outputs = torch.utils.checkpoint.checkpoint(
925
- create_custom_forward(decoder_layer),
926
  hidden_states,
927
- attention_mask,
928
  position_ids,
929
- None,
 
 
 
930
  )
931
  else:
932
  layer_outputs = decoder_layer(
933
  hidden_states,
934
- attention_mask=attention_mask,
935
  position_ids=position_ids,
936
- past_key_value=past_key_value,
937
  output_attentions=output_attentions,
938
  use_cache=use_cache,
 
939
  )
940
 
941
  hidden_states = layer_outputs[0]
942
 
943
  if use_cache:
944
- next_decoder_cache += (layer_outputs[2 if output_attentions else 1],)
945
 
946
  if output_attentions:
947
  all_self_attns += (layer_outputs[1],)
@@ -953,6 +1026,9 @@ class InternLM2Model(InternLM2PreTrainedModel):
953
  all_hidden_states += (hidden_states,)
954
 
955
  next_cache = next_decoder_cache if use_cache else None
 
 
 
956
  if not return_dict:
957
  return tuple(v for v in [hidden_states, next_cache, all_hidden_states, all_self_attns] if v is not None)
958
  return BaseModelOutputWithPast(
@@ -962,11 +1038,91 @@ class InternLM2Model(InternLM2PreTrainedModel):
962
  attentions=all_self_attns,
963
  )
964
 
 
 
 
 
 
 
 
 
 
 
 
 
 
965
 
966
- # Modified from transformers.model.llama.modeling_llama.LlamaForCausalLM
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
967
  class InternLM2ForCausalLM(InternLM2PreTrainedModel):
968
- _auto_class = "AutoModelForCausalLM"
969
 
 
970
  _tied_weights_keys = ["output.weight"]
971
 
972
  def __init__(self, config):
@@ -1003,13 +1159,14 @@ class InternLM2ForCausalLM(InternLM2PreTrainedModel):
1003
  input_ids: torch.LongTensor = None,
1004
  attention_mask: Optional[torch.Tensor] = None,
1005
  position_ids: Optional[torch.LongTensor] = None,
1006
- past_key_values: Optional[List[torch.FloatTensor]] = None,
1007
  inputs_embeds: Optional[torch.FloatTensor] = None,
1008
  labels: Optional[torch.LongTensor] = None,
1009
  use_cache: Optional[bool] = None,
1010
  output_attentions: Optional[bool] = None,
1011
  output_hidden_states: Optional[bool] = None,
1012
  return_dict: Optional[bool] = None,
 
1013
  ) -> Union[Tuple, CausalLMOutputWithPast]:
1014
  r"""
1015
  Args:
@@ -1025,8 +1182,8 @@ class InternLM2ForCausalLM(InternLM2PreTrainedModel):
1025
  ```python
1026
  >>> from transformers import AutoTokenizer, InternLM2ForCausalLM
1027
 
1028
- >>> model = InternLM2ForCausalLM.from_pretrained(PATH_TO_CONVERTED_WEIGHTS)
1029
- >>> tokenizer = AutoTokenizer.from_pretrained(PATH_TO_CONVERTED_TOKENIZER)
1030
 
1031
  >>> prompt = "Hey, are you conscious? Can you talk to me?"
1032
  >>> inputs = tokenizer(prompt, return_tensors="pt")
@@ -1054,10 +1211,19 @@ class InternLM2ForCausalLM(InternLM2PreTrainedModel):
1054
  output_attentions=output_attentions,
1055
  output_hidden_states=output_hidden_states,
1056
  return_dict=return_dict,
 
1057
  )
1058
 
1059
  hidden_states = outputs[0]
1060
- logits = self.output(hidden_states)
 
 
 
 
 
 
 
 
1061
  logits = logits.float()
1062
 
1063
  loss = None
@@ -1086,19 +1252,48 @@ class InternLM2ForCausalLM(InternLM2PreTrainedModel):
1086
  )
1087
 
1088
  def prepare_inputs_for_generation(
1089
- self, input_ids, past_key_values=None, attention_mask=None, inputs_embeds=None, **kwargs
 
 
 
 
 
 
 
1090
  ):
 
1091
  if past_key_values is not None:
1092
- past_length = past_key_values[0][0].shape[2]
1093
-
1094
- # Some generation methods already pass only the last input ID
1095
- if input_ids.shape[1] > past_length:
1096
- remove_prefix_length = past_length
 
 
 
 
1097
  else:
1098
- # Default to old behavior: keep only final ID
1099
- remove_prefix_length = input_ids.shape[1] - 1
1100
-
1101
- input_ids = input_ids[:, remove_prefix_length:]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1102
 
1103
  position_ids = kwargs.get("position_ids", None)
1104
  if attention_mask is not None and position_ids is None:
@@ -1112,13 +1307,24 @@ class InternLM2ForCausalLM(InternLM2PreTrainedModel):
1112
  if inputs_embeds is not None and past_key_values is None:
1113
  model_inputs = {"inputs_embeds": inputs_embeds}
1114
  else:
1115
- model_inputs = {"input_ids": input_ids}
 
 
 
 
 
 
 
 
 
 
1116
 
1117
  model_inputs.update(
1118
  {
1119
  "position_ids": position_ids,
 
1120
  "past_key_values": past_key_values,
1121
- "use_cache": kwargs.get("use_cache"),
1122
  "attention_mask": attention_mask,
1123
  }
1124
  )
@@ -1133,7 +1339,9 @@ class InternLM2ForCausalLM(InternLM2PreTrainedModel):
1133
  )
1134
  return reordered_past
1135
 
1136
- def build_inputs(self, tokenizer, query: str, history: List[Tuple[str, str]] = [], meta_instruction=""):
 
 
1137
  if tokenizer.add_bos_token:
1138
  prompt = ""
1139
  else:
@@ -1150,17 +1358,21 @@ class InternLM2ForCausalLM(InternLM2PreTrainedModel):
1150
  self,
1151
  tokenizer,
1152
  query: str,
1153
- history: List[Tuple[str, str]] = [],
1154
  streamer: Optional[BaseStreamer] = None,
1155
  max_new_tokens: int = 1024,
1156
  do_sample: bool = True,
1157
  temperature: float = 0.8,
1158
  top_p: float = 0.8,
1159
  meta_instruction: str = "You are an AI assistant whose name is InternLM (书生·浦语).\n"
1160
- "- InternLM (书生·浦语) is a conversational language model that is developed by Shanghai AI Laboratory (上海人工智能实验室). It is designed to be helpful, honest, and harmless.\n"
1161
- "- InternLM (书生·浦语) can understand and communicate fluently in the language chosen by the user such as English and 中文.",
 
 
1162
  **kwargs,
1163
  ):
 
 
1164
  inputs = self.build_inputs(tokenizer, query, history, meta_instruction)
1165
  inputs = {k: v.to(self.device) for k, v in inputs.items() if torch.is_tensor(v)}
1166
  # also add end-of-assistant token in eos token id to avoid unnecessary generation
@@ -1186,13 +1398,15 @@ class InternLM2ForCausalLM(InternLM2PreTrainedModel):
1186
  self,
1187
  tokenizer,
1188
  query: str,
1189
- history: List[Tuple[str, str]] = [],
1190
  max_new_tokens: int = 1024,
1191
  do_sample: bool = True,
1192
  temperature: float = 0.8,
1193
  top_p: float = 0.8,
1194
  **kwargs,
1195
  ):
 
 
1196
  """
1197
  Return a generator in format: (response, history)
1198
  Eg.
@@ -1208,6 +1422,10 @@ class InternLM2ForCausalLM(InternLM2PreTrainedModel):
1208
  response_queue = queue.Queue(maxsize=20)
1209
 
1210
  class ChatStreamer(BaseStreamer):
 
 
 
 
1211
  def __init__(self, tokenizer) -> None:
1212
  super().__init__()
1213
  self.tokenizer = tokenizer
@@ -1268,13 +1486,13 @@ class InternLM2ForCausalLM(InternLM2PreTrainedModel):
1268
  return consumer()
1269
 
1270
 
1271
- # Copied from transformers.model.llama.modeling_llama.LlamaForSequenceClassification with Llama->InternLM2
1272
  @add_start_docstrings(
1273
  """
1274
  The InternLM2 Model transformer with a sequence classification head on top (linear layer).
1275
 
1276
- [`InternLM2ForSequenceClassification`] uses the last token in order to do the classification,
1277
- as other causal models (e.g. GPT-2) do.
1278
 
1279
  Since it does classification on the last token, it requires to know the position of the last token. If a
1280
  `pad_token_id` is defined in the configuration, it finds the last token that is not a padding token in each row. If
@@ -1285,6 +1503,8 @@ class InternLM2ForCausalLM(InternLM2PreTrainedModel):
1285
  InternLM2_START_DOCSTRING,
1286
  )
1287
  class InternLM2ForSequenceClassification(InternLM2PreTrainedModel):
 
 
1288
  def __init__(self, config):
1289
  super().__init__(config)
1290
  self.num_labels = config.num_labels
@@ -1306,7 +1526,7 @@ class InternLM2ForSequenceClassification(InternLM2PreTrainedModel):
1306
  input_ids: torch.LongTensor = None,
1307
  attention_mask: Optional[torch.Tensor] = None,
1308
  position_ids: Optional[torch.LongTensor] = None,
1309
- past_key_values: Optional[List[torch.FloatTensor]] = None,
1310
  inputs_embeds: Optional[torch.FloatTensor] = None,
1311
  labels: Optional[torch.LongTensor] = None,
1312
  use_cache: Optional[bool] = None,
@@ -1347,9 +1567,10 @@ class InternLM2ForSequenceClassification(InternLM2PreTrainedModel):
1347
  sequence_lengths = -1
1348
  else:
1349
  if input_ids is not None:
1350
- sequence_lengths = (torch.eq(input_ids, self.config.pad_token_id).int().argmax(-1) - 1).to(
1351
- logits.device
1352
- )
 
1353
  else:
1354
  sequence_lengths = -1
1355
 
@@ -1361,7 +1582,7 @@ class InternLM2ForSequenceClassification(InternLM2PreTrainedModel):
1361
  if self.config.problem_type is None:
1362
  if self.num_labels == 1:
1363
  self.config.problem_type = "regression"
1364
- elif self.num_labels > 1 and (labels.dtype == torch.long or labels.dtype == torch.int):
1365
  self.config.problem_type = "single_label_classification"
1366
  else:
1367
  self.config.problem_type = "multi_label_classification"
@@ -1389,3 +1610,191 @@ class InternLM2ForSequenceClassification(InternLM2PreTrainedModel):
1389
  hidden_states=transformer_outputs.hidden_states,
1390
  attentions=transformer_outputs.attentions,
1391
  )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
13
  # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
14
  # See the License for the specific language governing permissions and
15
  # limitations under the License.
16
+ """PyTorch InternLM2.5 model."""
17
  import math
18
  import queue
19
  import threading
 
20
  from typing import List, Optional, Tuple, Union
21
 
22
  import torch
 
26
  from torch import nn
27
  from torch.nn import BCEWithLogitsLoss, CrossEntropyLoss, MSELoss
28
  from transformers.activations import ACT2FN
29
+ from transformers.cache_utils import Cache, DynamicCache, StaticCache
30
+ from transformers.modeling_attn_mask_utils import AttentionMaskConverter
31
  from transformers.modeling_outputs import (
32
  BaseModelOutputWithPast,
33
  CausalLMOutputWithPast,
34
+ QuestionAnsweringModelOutput,
35
  SequenceClassifierOutputWithPast,
36
+ TokenClassifierOutput,
37
  )
38
  from transformers.modeling_utils import PreTrainedModel
39
+ from transformers.pytorch_utils import ALL_LAYERNORM_LAYERS
40
  from transformers.utils import (
41
  add_start_docstrings,
42
  add_start_docstrings_to_model_forward,
43
+ is_flash_attn_greater_or_equal_2_10,
44
  logging,
45
  replace_return_docstrings,
46
  )
47
 
48
  try:
49
  from transformers.generation.streamers import BaseStreamer
50
+ except Exception:
51
  BaseStreamer = None
52
 
53
  from .configuration_internlm2 import InternLM2Config
54
 
55
+
56
+ try:
57
+ from flash_attn import flash_attn_func, flash_attn_varlen_func
58
+ from flash_attn.bert_padding import index_first_axis, pad_input, unpad_input
59
+ except:
60
+ pass
61
+
62
+
63
  logger = logging.get_logger(__name__)
64
 
65
  _CONFIG_FOR_DOC = "InternLM2Config"
66
 
67
+
 
 
 
 
 
 
 
 
 
 
 
 
 
68
  def _get_unpad_data(attention_mask):
69
  seqlens_in_batch = attention_mask.sum(dim=-1, dtype=torch.int32)
70
  indices = torch.nonzero(attention_mask.flatten(), as_tuple=False).flatten()
71
  max_seqlen_in_batch = seqlens_in_batch.max().item()
72
+ cu_seqlens = F.pad(torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0)) # pylint: disable=E1102
73
  return (
74
  indices,
75
  cu_seqlens,
 
77
  )
78
 
79
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
80
  class InternLM2RMSNorm(nn.Module):
81
+ """InternLM2RMSNorm is equivalent to T5LayerNorm."""
82
+
83
  def __init__(self, hidden_size, eps=1e-6):
 
 
 
84
  super().__init__()
85
  self.weight = nn.Parameter(torch.ones(hidden_size))
86
  self.variance_epsilon = eps
 
93
  return self.weight * hidden_states.to(input_dtype)
94
 
95
 
96
+ ALL_LAYERNORM_LAYERS.append(InternLM2RMSNorm)
97
+
98
+
99
  class InternLM2RotaryEmbedding(nn.Module):
100
+ """Rotary Position Embedding for the InternLM2 model. Credits to the Reddit user /u/lucidrains."""
 
101
 
102
+ def __init__(self, dim, max_position_embeddings=2048, base=10000, device=None, scaling_factor=1.0):
103
+ super().__init__()
104
+ self.scaling_factor = scaling_factor
105
  self.dim = dim
106
  self.max_position_embeddings = max_position_embeddings
107
  self.base = base
108
+ inv_freq = 1.0 / (self.base ** (torch.arange(0, self.dim, 2, dtype=torch.int64).float().to(device) / self.dim))
109
  self.register_buffer("inv_freq", inv_freq, persistent=False)
110
+ # For BC we register cos and sin cached
111
+ self.max_seq_len_cached = max_position_embeddings
112
 
113
+ @torch.no_grad()
114
+ def forward(self, x, position_ids):
 
 
 
 
 
 
 
 
 
 
 
 
 
 
115
  # x: [bs, num_attention_heads, seq_len, head_size]
116
+ inv_freq_expanded = self.inv_freq[None, :, None].float().expand(position_ids.shape[0], -1, 1)
117
+ position_ids_expanded = position_ids[:, None, :].float()
118
+ # Force float32 since bfloat16 loses precision on long contexts
119
+ # See https://github.com/huggingface/transformers/pull/29285
120
+ device_type = x.device.type
121
+ device_type = device_type if isinstance(device_type, str) and device_type != "mps" else "cpu"
122
+ with torch.autocast(device_type=device_type, enabled=False):
123
+ freqs = (inv_freq_expanded.float() @ position_ids_expanded.float()).transpose(1, 2)
124
+ emb = torch.cat((freqs, freqs), dim=-1)
125
+ cos = emb.cos()
126
+ sin = emb.sin()
127
+ return cos.to(dtype=x.dtype), sin.to(dtype=x.dtype)
128
 
 
 
 
 
129
 
 
 
130
  class InternLM2LinearScalingRotaryEmbedding(InternLM2RotaryEmbedding):
131
  """InternLM2RotaryEmbedding extended with linear scaling. Credits to the Reddit user /u/kaiokendev"""
132
 
133
+ def forward(self, x, position_ids):
134
+ # difference to the original RoPE: a scaling factor is aplied to the position ids
135
+ position_ids = position_ids.float() / self.scaling_factor
136
+ cos, sin = super().forward(x, position_ids)
137
+ return cos, sin
 
 
 
 
 
 
 
 
 
138
 
139
 
 
140
  class InternLM2DynamicNTKScalingRotaryEmbedding(InternLM2RotaryEmbedding):
141
  """InternLM2RotaryEmbedding extended with Dynamic NTK scaling.
142
+ Credits to the Reddit users /u/bloc97 and /u/emozilla"""
 
 
 
 
 
 
 
 
143
 
144
+ def forward(self, x, position_ids):
145
+ # difference to the original RoPE: inv_freq is recomputed when the sequence length > original length
146
+ seq_len = torch.max(position_ids) + 1
147
  if seq_len > self.max_position_embeddings:
148
  base = self.base * (
149
  (self.scaling_factor * seq_len / self.max_position_embeddings) - (self.scaling_factor - 1)
150
  ) ** (self.dim / (self.dim - 2))
151
+ inv_freq = 1.0 / (base ** (torch.arange(0, self.dim, 2, dtype=torch.int64).float().to(x.device) / self.dim))
152
+ self.register_buffer("inv_freq", inv_freq, persistent=False) # TODO joao: this may break with compilation
153
 
154
+ cos, sin = super().forward(x, position_ids)
155
+ return cos, sin
156
 
 
 
 
 
 
157
 
 
 
158
  def rotate_half(x):
159
  """Rotates half the hidden dims of the input."""
160
  x1 = x[..., : x.shape[-1] // 2]
 
162
  return torch.cat((-x2, x1), dim=-1)
163
 
164
 
165
+ def apply_rotary_pos_emb(q, k, cos, sin, position_ids=None, unsqueeze_dim=1): # pylint: disable=unused-argument
166
+ """Applies Rotary Position Embedding to the query and key tensors.
167
+
168
+ Args:
169
+ q (`torch.Tensor`): The query tensor.
170
+ k (`torch.Tensor`): The key tensor.
171
+ cos (`torch.Tensor`): The cosine part of the rotary embedding.
172
+ sin (`torch.Tensor`): The sine part of the rotary embedding.
173
+ position_ids (`torch.Tensor`, *optional*):
174
+ Deprecated and unused.
175
+ unsqueeze_dim (`int`, *optional*, defaults to 1):
176
+ The 'unsqueeze_dim' argument specifies the dimension along which to unsqueeze cos[position_ids] and
177
+ sin[position_ids] so that they can be properly broadcasted to the dimensions of q and k. For example, note
178
+ that cos[position_ids] and sin[position_ids] have the shape [batch_size, seq_len, head_dim]. Then, if q and
179
+ k have the shape [batch_size, heads, seq_len, head_dim], then setting unsqueeze_dim=1 makes
180
+ cos[position_ids] and sin[position_ids] broadcastable to the shapes of q and k. Similarly, if q and k have
181
+ the shape [batch_size, seq_len, heads, head_dim], then set unsqueeze_dim=2.
182
+ Returns:
183
+ `tuple(torch.Tensor)` comprising of the query and key tensors rotated using the Rotary Position Embedding.
184
+ """
185
+ cos = cos.unsqueeze(unsqueeze_dim)
186
+ sin = sin.unsqueeze(unsqueeze_dim)
187
  q_embed = (q * cos) + (rotate_half(q) * sin)
188
  k_embed = (k * cos) + (rotate_half(k) * sin)
189
  return q_embed, k_embed
190
 
191
 
192
  class InternLM2MLP(nn.Module):
193
+ """MLP for InternLM2 model."""
194
+
195
  def __init__(self, config):
196
  super().__init__()
197
  self.config = config
 
208
  return down_proj
209
 
210
 
 
211
  def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor:
212
  """
213
  This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch,
 
220
  return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim)
221
 
222
 
 
223
  class InternLM2Attention(nn.Module):
224
  """Multi-headed attention from 'Attention Is All You Need' paper"""
225
 
226
+ def __init__(self, config: InternLM2Config, layer_idx: Optional[int] = None):
227
  super().__init__()
228
  self.config = config
229
+ self.layer_idx = layer_idx
230
+ if layer_idx is None:
231
+ logger.warning_once(
232
+ f"Instantiating {self.__class__.__name__} without passing a `layer_idx` is not recommended and will "
233
+ "lead to errors during the forward call if caching is used. Please make sure to provide a `layer_idx` "
234
+ "when creating this class."
235
+ )
236
+
237
  self.hidden_size = config.hidden_size
238
  self.num_heads = config.num_attention_heads
239
  self.head_dim = self.hidden_size // self.num_heads
240
  self.num_key_value_heads = config.num_key_value_heads
241
  self.num_key_value_groups = self.num_heads // self.num_key_value_heads
242
  self.max_position_embeddings = config.max_position_embeddings
243
+ self.rope_theta = config.rope_theta
244
  self.is_causal = True
245
 
246
  if (self.head_dim * self.num_heads) != self.hidden_size:
 
254
  (self.num_heads + 2 * self.num_key_value_heads) * self.head_dim,
255
  bias=config.bias,
256
  )
 
257
  self.wo = nn.Linear(self.num_heads * self.head_dim, self.hidden_size, bias=config.bias)
258
+
259
  self._init_rope()
260
 
261
  def _init_rope(self):
 
263
  self.rotary_emb = InternLM2RotaryEmbedding(
264
  self.head_dim,
265
  max_position_embeddings=self.max_position_embeddings,
266
+ base=self.rope_theta,
267
  )
268
  else:
269
  scaling_type = self.config.rope_scaling["type"]
270
  scaling_factor = self.config.rope_scaling["factor"]
271
+ if scaling_type == "linear":
272
+ self.rotary_emb = InternLM2LinearScalingRotaryEmbedding(
273
  self.head_dim,
274
  max_position_embeddings=self.max_position_embeddings,
 
275
  scaling_factor=scaling_factor,
276
+ base=self.rope_theta,
277
  )
278
+ elif scaling_type == "dynamic":
279
+ self.rotary_emb = InternLM2DynamicNTKScalingRotaryEmbedding(
280
  self.head_dim,
281
  max_position_embeddings=self.max_position_embeddings,
 
282
  scaling_factor=scaling_factor,
283
+ base=self.rope_theta,
284
  )
285
  else:
286
+ raise ValueError(f"Unknown RoPE scaling type {scaling_type}")
 
 
 
 
287
 
288
  def forward(
289
  self,
290
  hidden_states: torch.Tensor,
291
  attention_mask: Optional[torch.Tensor] = None,
292
  position_ids: Optional[torch.LongTensor] = None,
293
+ past_key_value: Optional[Cache] = None,
294
  output_attentions: bool = False,
295
+ use_cache: bool = False, # pylint: disable=unused-argument
296
+ cache_position: Optional[torch.LongTensor] = None,
297
  ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]:
 
 
 
 
 
 
298
  bsz, q_len, _ = hidden_states.size()
299
 
300
+ if self.config.pretraining_tp > 1:
301
+ # split qkv_states by tp size
302
+ key_value_slicing = (self.num_key_value_heads * self.head_dim) // self.config.pretraining_tp
303
+ qkv_slices = self.wqkv.weight.split(key_value_slicing, dim=0)
304
+ qkv_states = torch.cat(
305
+ [F.linear(hidden_states, qkv_slice) for qkv_slice in qkv_slices], dim=-1 # pylint: disable=E1102
306
+ )
307
+ else:
308
+ qkv_states = self.wqkv(hidden_states)
309
 
310
  qkv_states = rearrange(
311
  qkv_states,
 
315
  )
316
 
317
  query_states = qkv_states[..., : self.num_key_value_groups, :]
318
+ query_states = rearrange(query_states, "b q h gs d -> b q (h gs) d").transpose(1, 2)
319
+ key_states = qkv_states[..., -2, :].transpose(1, 2)
320
+ value_states = qkv_states[..., -1, :].transpose(1, 2)
321
 
322
+ cos, sin = self.rotary_emb(value_states, position_ids)
 
 
 
 
 
 
 
323
  query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin, position_ids)
324
 
325
  if past_key_value is not None:
326
+ # sin and cos are specific to RoPE models; cache_position needed for the static cache
327
+ cache_kwargs = {"sin": sin, "cos": cos, "cache_position": cache_position}
328
+ key_states, value_states = past_key_value.update(key_states, value_states, self.layer_idx, cache_kwargs)
 
 
329
 
330
  key_states = repeat_kv(key_states, self.num_key_value_groups)
331
  value_states = repeat_kv(value_states, self.num_key_value_groups)
332
 
333
  attn_weights = torch.matmul(query_states, key_states.transpose(2, 3)) / math.sqrt(self.head_dim)
334
 
335
+ if attention_mask is not None: # no matter the length, we just slice it
336
+ causal_mask = attention_mask[:, :, :, : key_states.shape[-2]]
337
+ attn_weights = attn_weights + causal_mask
 
 
 
 
 
 
 
 
 
338
 
339
  # upcast attention to fp32
340
  attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query_states.dtype)
 
347
  )
348
 
349
  attn_output = attn_output.transpose(1, 2).contiguous()
350
+
351
  attn_output = attn_output.reshape(bsz, q_len, self.hidden_size)
352
 
353
+ if self.config.pretraining_tp > 1:
354
+ attn_output = attn_output.split(self.hidden_size // self.config.pretraining_tp, dim=2)
355
+ o_proj_slices = self.wo.weight.split(self.hidden_size // self.config.pretraining_tp, dim=1)
356
+ attn_output = sum(
357
+ [
358
+ F.linear(attn_output[i], o_proj_slices[i]) # pylint: disable=E1102
359
+ for i in range(self.config.pretraining_tp)
360
+ ]
361
+ )
362
+ else:
363
+ attn_output = self.wo(attn_output)
364
 
365
  if not output_attentions:
366
  attn_weights = None
 
368
  return attn_output, attn_weights, past_key_value
369
 
370
 
 
371
  class InternLM2FlashAttention2(InternLM2Attention):
372
  """
373
  InternLM2 flash attention module. This module inherits from `InternLM2Attention` as the weights of the module stays
 
375
  flash attention and deal with padding tokens in case the input contains any of them.
376
  """
377
 
378
+ def __init__(self, *args, **kwargs):
379
+ super().__init__(*args, **kwargs)
380
+
381
+ # TODO: Should be removed once Flash Attention for RoCm is bumped to 2.1.
382
+ # flash_attn<2.1 generates top-left aligned causal mask, while what is needed here is bottom-right alignement,
383
+ # that was made default for flash_attn>=2.1. This attribute is used to handle this difference.
384
+ # Reference: https://github.com/Dao-AILab/flash-attention/releases/tag/v2.1.0.
385
+ # Beware that with flash_attn<2.1, using q_seqlen != k_seqlen (except for the case q_seqlen == 1)
386
+ # produces a wrong mask (top-left).
387
+ self._flash_attn_uses_top_left_mask = not is_flash_attn_greater_or_equal_2_10()
388
+
389
  def forward(
390
  self,
391
  hidden_states: torch.Tensor,
392
  attention_mask: Optional[torch.LongTensor] = None,
393
  position_ids: Optional[torch.LongTensor] = None,
394
+ past_key_value: Optional[Cache] = None,
395
  output_attentions: bool = False,
396
  use_cache: bool = False,
397
+ cache_position: Optional[torch.LongTensor] = None,
398
  ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]:
399
+ if isinstance(past_key_value, StaticCache):
400
+ raise ValueError(
401
+ "`static` cache implementation is not compatible with `attn_implementation==flash_attention_2` "
402
+ "make sure to use `sdpa` in the mean time, and open an issue at "
403
+ "https://github.com/huggingface/transformers"
404
  )
405
 
 
 
 
406
  output_attentions = False
407
 
408
  bsz, q_len, _ = hidden_states.size()
 
425
  key_states = key_states.transpose(1, 2)
426
  value_states = value_states.transpose(1, 2)
427
 
428
+ cos, sin = self.rotary_emb(value_states, position_ids)
429
+ query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin)
 
 
 
 
 
430
 
431
  if past_key_value is not None:
432
+ # sin and cos are specific to RoPE models; cache_position needed for the static cache
433
+ cache_kwargs = {"sin": sin, "cos": cos, "cache_position": cache_position}
434
+ key_states, value_states = past_key_value.update(key_states, value_states, self.layer_idx, cache_kwargs)
 
 
435
 
436
+ # TODO: These transpose are quite inefficient but Flash Attention requires the layout
437
+ # [batch_size, sequence_length, num_heads, head_dim]. We would need to refactor the KV cache
438
+ # to be able to avoid many of these transpose/reshape/view.
439
  query_states = query_states.transpose(1, 2)
440
  key_states = key_states.transpose(1, 2)
441
  value_states = value_states.transpose(1, 2)
442
 
443
+ # dropout_rate = self.attention_dropout if self.training else 0.0
444
+ dropout_rate = 0.0
445
+
446
+ # In PEFT, usually we cast the layer norms in float32 for training stability reasons
447
+ # therefore the input hidden states gets silently casted in float32. Hence, we need
448
+ # cast them back in the correct dtype just to be sure everything works as expected.
449
+ # This might slowdown training & inference so it is recommended to not cast the LayerNorms
450
+ # in fp32. (InternLM2RMSNorm handles it correctly)
451
+
452
+ input_dtype = query_states.dtype
453
+ if input_dtype == torch.float32:
454
+ if torch.is_autocast_enabled():
455
+ target_dtype = torch.get_autocast_gpu_dtype()
456
+ # Handle the case where the model is quantized
457
+ elif hasattr(self.config, "_pre_quantization_dtype"):
458
+ target_dtype = self.config._pre_quantization_dtype
459
+ else:
460
+ target_dtype = self.wqkv.weight.dtype
461
+
462
+ logger.warning_once(
463
+ f"The input hidden states seems to be silently casted in float32, this might be related to"
464
+ f" the fact you have upcasted embedding or layer norm layers in float32. We will cast back the input in"
465
+ f" {target_dtype}."
466
+ )
467
+
468
+ query_states = query_states.to(target_dtype)
469
+ key_states = key_states.to(target_dtype)
470
+ value_states = value_states.to(target_dtype)
471
+
472
  attn_output = self._flash_attention_forward(
473
+ query_states, key_states, value_states, attention_mask, q_len, dropout=dropout_rate
474
  )
475
+
476
  attn_output = attn_output.reshape(bsz, q_len, self.hidden_size).contiguous()
477
  attn_output = self.wo(attn_output)
478
 
479
  if not output_attentions:
480
  attn_weights = None
481
 
482
+ return attn_output, attn_weights, past_key_value # pylint: disable=E0606
483
 
484
  def _flash_attention_forward(
485
  self, query_states, key_states, value_states, attention_mask, query_length, dropout=0.0, softmax_scale=None
 
498
  attention_mask (`torch.Tensor`):
499
  The padding mask - corresponds to a tensor of size `(batch_size, seq_len)` where 0 stands for the
500
  position of padding tokens and 1 for the position of non-padding tokens.
501
+ dropout (`float`):
502
  Attention dropout
503
  softmax_scale (`float`, *optional*):
504
  The scaling of QK^T before applying softmax. Default to 1 / sqrt(head_dim)
505
  """
506
+ if not self._flash_attn_uses_top_left_mask:
507
+ causal = self.is_causal
508
+ else:
509
+ # TODO: Remove the `query_length != 1` check once Flash Attention for RoCm is bumped to 2.1.
510
+ # For details, please see the comment in InternLM2FlashAttention2 __init__.
511
+ causal = self.is_causal and query_length != 1
512
+
513
  # Contains at least one padding token in the sequence
 
514
  if attention_mask is not None:
515
  batch_size = query_states.shape[0]
516
+ query_states, key_states, value_states, indices_q, cu_seq_lens, max_seq_lens = self._upad_input(
517
  query_states, key_states, value_states, attention_mask, query_length
518
  )
519
 
520
  cu_seqlens_q, cu_seqlens_k = cu_seq_lens
521
  max_seqlen_in_batch_q, max_seqlen_in_batch_k = max_seq_lens
522
 
523
+ attn_output_unpad = flash_attn_varlen_func( # pylint: disable=E0606
524
  query_states,
525
  key_states,
526
  value_states,
 
533
  causal=causal,
534
  )
535
 
536
+ attn_output = pad_input(attn_output_unpad, indices_q, batch_size, query_length) # pylint: disable=E0606
537
  else:
538
+ attn_output = flash_attn_func( # pylint: disable=E0606
539
  query_states, key_states, value_states, dropout, softmax_scale=softmax_scale, causal=causal
540
  )
541
 
542
  return attn_output
543
 
544
+ def _upad_input(self, query_layer, key_layer, value_layer, attention_mask, query_length):
545
  indices_k, cu_seqlens_k, max_seqlen_in_batch_k = _get_unpad_data(attention_mask)
546
  batch_size, kv_seq_len, num_key_value_heads, head_dim = key_layer.shape
547
 
548
+ key_layer = index_first_axis( # pylint: disable=E0606
549
  key_layer.reshape(batch_size * kv_seq_len, num_key_value_heads, head_dim), indices_k
550
  )
551
+ value_layer = index_first_axis( # pylint: disable=E0606
552
  value_layer.reshape(batch_size * kv_seq_len, num_key_value_heads, head_dim), indices_k
553
  )
 
554
  if query_length == kv_seq_len:
555
+ query_layer = index_first_axis( # pylint: disable=E0606
556
  query_layer.reshape(batch_size * kv_seq_len, self.num_heads, head_dim), indices_k
557
  )
558
  cu_seqlens_q = cu_seqlens_k
 
568
  else:
569
  # The -q_len: slice assumes left padding.
570
  attention_mask = attention_mask[:, -query_length:]
571
+ query_layer, indices_q, cu_seqlens_q, max_seqlen_in_batch_q = unpad_input( # pylint: disable=E0606
572
+ query_layer, attention_mask
573
+ )
574
 
575
  return (
576
  query_layer,
577
  key_layer,
578
  value_layer,
579
+ indices_q,
580
  (cu_seqlens_q, cu_seqlens_k),
581
  (max_seqlen_in_batch_q, max_seqlen_in_batch_k),
582
  )
583
 
584
+
585
+ # Copied from transformers.models.llama.modeling_llama.LllamaSdpaAttention with Llama->InternLM2
586
+ class InternLM2SdpaAttention(InternLM2Attention):
587
+ """
588
+ InternLM2 attention module using torch.nn.functional.scaled_dot_product_attention. This module inherits from
589
+ `InternLM2Attention` as the weights of the module stays untouched. The only changes are on the forward pass
590
+ to adapt to SDPA API.
591
+ """
592
+
593
+ # Adapted from InternLM2Attention.forward
594
+ def forward(
595
+ self,
596
+ hidden_states: torch.Tensor,
597
+ attention_mask: Optional[torch.Tensor] = None,
598
+ position_ids: Optional[torch.LongTensor] = None,
599
+ past_key_value: Optional[Cache] = None,
600
+ output_attentions: bool = False,
601
+ use_cache: bool = False,
602
+ cache_position: Optional[torch.LongTensor] = None,
603
+ ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]:
604
+ if output_attentions:
605
+ # TODO: Improve this warning with e.g. `model.config.attn_implementation = "manual"`
606
+ # once this is implemented.
607
+ logger.warning_once(
608
+ "InternLM2Model uses InternLM2SdpaAttention, but `torch.nn.functional.scaled_dot_product_attention` "
609
+ "does not support `output_attentions=True`. Falling back to the manual attention implementation, "
610
+ "but specifying the manual implementation will be required from Transformers version v5.0.0 onwards. "
611
+ 'This warning can be removed using the argument `attn_implementation="eager"` when loading the model.'
612
+ )
613
+ return super().forward(
614
+ hidden_states=hidden_states,
615
+ attention_mask=attention_mask,
616
+ position_ids=position_ids,
617
+ past_key_value=past_key_value,
618
+ output_attentions=output_attentions,
619
+ use_cache=use_cache,
620
+ cache_position=cache_position,
621
+ )
622
+
623
+ bsz, q_len, _ = hidden_states.size()
624
+
625
+ qkv_states = self.wqkv(hidden_states)
626
+
627
+ qkv_states = rearrange(
628
+ qkv_states,
629
+ "b q (h gs d) -> b q h gs d",
630
+ gs=2 + self.num_key_value_groups,
631
+ d=self.head_dim,
632
+ )
633
+
634
+ query_states = qkv_states[..., : self.num_key_value_groups, :]
635
+ query_states = rearrange(query_states, "b q h gs d -> b q (h gs) d")
636
+ key_states = qkv_states[..., -2, :]
637
+ value_states = qkv_states[..., -1, :]
638
+
639
+ query_states = query_states.transpose(1, 2)
640
+ key_states = key_states.transpose(1, 2)
641
+ value_states = value_states.transpose(1, 2)
642
+
643
+ cos, sin = self.rotary_emb(value_states, position_ids)
644
+ query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin)
645
+
646
+ if past_key_value is not None:
647
+ # sin and cos are specific to RoPE models; cache_position needed for the static cache
648
+ cache_kwargs = {"sin": sin, "cos": cos, "cache_position": cache_position}
649
+ key_states, value_states = past_key_value.update(key_states, value_states, self.layer_idx, cache_kwargs)
650
+
651
+ key_states = repeat_kv(key_states, self.num_key_value_groups)
652
+ value_states = repeat_kv(value_states, self.num_key_value_groups)
653
+
654
+ causal_mask = attention_mask
655
+ if attention_mask is not None:
656
+ causal_mask = causal_mask[:, :, :, : key_states.shape[-2]]
657
+
658
+ # SDPA with memory-efficient backend is currently (torch==2.1.2) bugged with non-contiguous inputs with
659
+ # custom attn_mask, Reference: https://github.com/pytorch/pytorch/issues/112577.
660
+ if query_states.device.type == "cuda" and causal_mask is not None:
661
+ query_states = query_states.contiguous()
662
+ key_states = key_states.contiguous()
663
+ value_states = value_states.contiguous()
664
+
665
+ # We dispatch to SDPA's Flash Attention or Efficient kernels via this `is_causal` if statement instead of
666
+ # an inline conditional assignment in SDPA to support both torch.compile's dynamic shapes and full graph
667
+ # options. An inline conditional prevents dynamic shapes from compiling.
668
+ is_causal = bool(causal_mask is None and q_len > 1)
669
+
670
+ attn_output = torch.nn.functional.scaled_dot_product_attention( # pylint: disable=E1102
671
+ query_states,
672
+ key_states,
673
+ value_states,
674
+ attn_mask=causal_mask,
675
+ dropout_p=0.0,
676
+ is_causal=is_causal,
677
+ )
678
+
679
+ attn_output = attn_output.transpose(1, 2).contiguous()
680
+ attn_output = attn_output.view(bsz, q_len, self.hidden_size)
681
+
682
+ attn_output = self.wo(attn_output)
683
+
684
+ return attn_output, None, past_key_value
685
+
686
+
687
  INTERNLM2_ATTENTION_CLASSES = {
688
  "eager": InternLM2Attention,
689
  "flash_attention_2": InternLM2FlashAttention2,
690
+ "sdpa": InternLM2SdpaAttention,
691
  }
692
 
693
+
694
+ # Modified from transformers.models.llama.modeling_llama.LlamaDecoderLayer with Llama->InternLM2
695
  class InternLM2DecoderLayer(nn.Module):
696
+ """InternLM2 Decoder Layer. This module is a single layer of the InternLM2 model."""
697
+
698
+ def __init__(self, config: InternLM2Config, layer_idx: int):
699
  super().__init__()
700
  self.hidden_size = config.hidden_size
701
+ self.layer_idx = layer_idx
702
 
703
+ self.attention = INTERNLM2_ATTENTION_CLASSES[config.attn_implementation](config=config, layer_idx=layer_idx)
704
 
705
  self.feed_forward = InternLM2MLP(config)
706
  self.attention_norm = InternLM2RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
 
711
  hidden_states: torch.Tensor,
712
  attention_mask: Optional[torch.Tensor] = None,
713
  position_ids: Optional[torch.LongTensor] = None,
714
+ past_key_value: Optional[Cache] = None,
715
  output_attentions: Optional[bool] = False,
716
  use_cache: Optional[bool] = False,
717
+ cache_position: Optional[torch.LongTensor] = None,
718
  ) -> Tuple[torch.FloatTensor, Optional[Tuple[torch.FloatTensor, torch.FloatTensor]]]:
719
  """
720
  Args:
 
730
  (see `past_key_values`).
731
  past_key_value (`Tuple(torch.FloatTensor)`, *optional*): cached past key and value projection states
732
  """
 
 
 
 
 
 
733
  residual = hidden_states
734
 
735
  hidden_states = self.attention_norm(hidden_states)
 
742
  past_key_value=past_key_value,
743
  output_attentions=output_attentions,
744
  use_cache=use_cache,
745
+ cache_position=cache_position,
746
  )
747
  hidden_states = residual + hidden_states
748
 
 
786
  InternLM2_START_DOCSTRING,
787
  )
788
  class InternLM2PreTrainedModel(PreTrainedModel):
789
+ """
790
+ InternLM2 pretraiend model's base class.
791
+ """
792
+
793
  config_class = InternLM2Config
794
  base_model_prefix = "model"
795
  supports_gradient_checkpointing = True
796
  _no_split_modules = ["InternLM2DecoderLayer"]
797
+ _skip_keys_device_placement = ["past_key_values"]
798
+ _supports_flash_attn_2 = True
799
+ _supports_sdpa = True
800
+ _supports_cache_class = True
801
+ _supports_quantized_cache = True
802
+ _supports_static_cache = True
803
 
804
  def _init_weights(self, module):
805
  std = self.config.initializer_range
 
848
  config.n_positions - 1]`.
849
 
850
  [What are position IDs?](../glossary#position-ids)
851
+ past_key_values (`Cache` or `tuple(tuple(torch.FloatTensor))`, *optional*):
852
+ Pre-computed hidden-states (key and values in the self-attention blocks and in the cross-attention
853
+ blocks) that can be used to speed up sequential decoding. This typically consists in the `past_key_values`
854
+ returned by the model at a previous stage of decoding, when `use_cache=True` or `config.use_cache=True`.
855
+
856
+ Two formats are allowed:
857
+ - a [`~cache_utils.Cache`] instance;
858
+ - Tuple of `tuple(torch.FloatTensor)` of length `config.n_layers`, with each tuple having 2 tensors of
859
+ shape `(batch_size, num_heads, sequence_length, embed_size_per_head)`). This is also known as the legacy
860
+ cache format.
861
 
862
+ The model will output the same cache format that is fed as input. If no `past_key_values` are passed, the
863
+ legacy cache format will be returned.
864
 
865
  If `past_key_values` are used, the user can optionally input only the last `input_ids` (those that don't
866
  have their past key value states given to this model) of shape `(batch_size, 1)` instead of all `input_ids`
 
880
  more detail.
881
  return_dict (`bool`, *optional*):
882
  Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.
883
+ cache_position (`torch.LongTensor` of shape `(sequence_length)`, *optional*):
884
+ Indices depicting the position of the input sequence tokens in the sequence. Contrarily to `position_ids`,
885
+ this tensor is not affected by padding. It is used to update the cache in the correct position and to infer
886
+ the complete sequence length.
887
  """
888
 
889
 
890
+ # Modified from transformers.models.llama.modeling_llama.LlamaModel with Llama->InternLM2
891
  @add_start_docstrings(
892
  "The bare InternLM2 Model outputting raw hidden-states without any specific head on top.",
893
  InternLM2_START_DOCSTRING,
 
910
 
911
  self.tok_embeddings = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx)
912
 
913
+ self.layers = nn.ModuleList(
914
+ [InternLM2DecoderLayer(config, layer_idx) for layer_idx in range(config.num_hidden_layers)]
915
+ )
916
  self.norm = InternLM2RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
917
 
918
  self.gradient_checkpointing = False
 
925
  def set_input_embeddings(self, value):
926
  self.tok_embeddings = value
927
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
928
  @add_start_docstrings_to_model_forward(InternLM2_INPUTS_DOCSTRING)
929
  def forward(
930
  self,
931
  input_ids: torch.LongTensor = None,
932
  attention_mask: Optional[torch.Tensor] = None,
933
  position_ids: Optional[torch.LongTensor] = None,
934
+ past_key_values: Optional[Union[Cache, List[torch.FloatTensor]]] = None,
935
  inputs_embeds: Optional[torch.FloatTensor] = None,
936
  use_cache: Optional[bool] = None,
937
  output_attentions: Optional[bool] = None,
938
  output_hidden_states: Optional[bool] = None,
939
  return_dict: Optional[bool] = None,
940
+ cache_position: Optional[torch.LongTensor] = None,
941
  ) -> Union[Tuple, BaseModelOutputWithPast]:
942
  output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
943
  output_hidden_states = (
944
  output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
945
  )
946
  use_cache = use_cache if use_cache is not None else self.config.use_cache
 
947
  return_dict = return_dict if return_dict is not None else self.config.use_return_dict
948
 
949
+ if (input_ids is None) ^ (inputs_embeds is not None):
950
+ raise ValueError(
951
+ "You cannot specify both input_ids and inputs_embeds at the same time, and must specify either one"
952
+ )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
953
 
954
+ if self.gradient_checkpointing and self.training and use_cache:
955
+ logger.warning_once(
956
+ "`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`."
 
957
  )
958
+ use_cache = False
959
 
960
  if inputs_embeds is None:
961
  inputs_embeds = self.tok_embeddings(input_ids)
962
 
963
+ return_legacy_cache = False
964
+ if use_cache and not isinstance(past_key_values, Cache): # kept for BC (non `Cache` `past_key_values` inputs)
965
+ return_legacy_cache = True
966
+ past_key_values = DynamicCache.from_legacy_cache(past_key_values)
967
+
968
+ if cache_position is None:
969
+ past_seen_tokens = past_key_values.get_seq_length() if past_key_values is not None else 0
970
+ cache_position = torch.arange(
971
+ past_seen_tokens, past_seen_tokens + inputs_embeds.shape[1], device=inputs_embeds.device
 
972
  )
973
+ if position_ids is None:
974
+ position_ids = cache_position.unsqueeze(0)
975
+
976
+ causal_mask = self._update_causal_mask(
977
+ attention_mask, inputs_embeds, cache_position, past_key_values, output_attentions
978
+ )
979
 
980
  # embed positions
981
  hidden_states = inputs_embeds
982
 
 
 
 
 
 
 
 
983
  # decoder layers
984
  all_hidden_states = () if output_hidden_states else None
985
  all_self_attns = () if output_attentions else None
986
+ next_decoder_cache = None
987
 
988
+ for decoder_layer in self.layers:
989
  if output_hidden_states:
990
  all_hidden_states += (hidden_states,)
991
 
 
 
992
  if self.gradient_checkpointing and self.training:
993
+ layer_outputs = self._gradient_checkpointing_func(
994
+ decoder_layer.__call__,
 
 
 
 
 
 
 
 
995
  hidden_states,
996
+ causal_mask,
997
  position_ids,
998
+ past_key_values,
999
+ output_attentions,
1000
+ use_cache,
1001
+ cache_position,
1002
  )
1003
  else:
1004
  layer_outputs = decoder_layer(
1005
  hidden_states,
1006
+ attention_mask=causal_mask,
1007
  position_ids=position_ids,
1008
+ past_key_value=past_key_values,
1009
  output_attentions=output_attentions,
1010
  use_cache=use_cache,
1011
+ cache_position=cache_position,
1012
  )
1013
 
1014
  hidden_states = layer_outputs[0]
1015
 
1016
  if use_cache:
1017
+ next_decoder_cache = layer_outputs[2 if output_attentions else 1]
1018
 
1019
  if output_attentions:
1020
  all_self_attns += (layer_outputs[1],)
 
1026
  all_hidden_states += (hidden_states,)
1027
 
1028
  next_cache = next_decoder_cache if use_cache else None
1029
+ if return_legacy_cache:
1030
+ next_cache = next_cache.to_legacy_cache()
1031
+
1032
  if not return_dict:
1033
  return tuple(v for v in [hidden_states, next_cache, all_hidden_states, all_self_attns] if v is not None)
1034
  return BaseModelOutputWithPast(
 
1038
  attentions=all_self_attns,
1039
  )
1040
 
1041
+ def _update_causal_mask(
1042
+ self,
1043
+ attention_mask: torch.Tensor,
1044
+ input_tensor: torch.Tensor,
1045
+ cache_position: torch.Tensor,
1046
+ past_key_values: Cache,
1047
+ output_attentions: bool,
1048
+ ):
1049
+ # TODO: As of torch==2.2.0, the `attention_mask` passed to the model in `generate` is 2D and of dynamic length
1050
+ # even when the static KV cache is used. This is an issue for torch.compile which then recaptures cudagraphs at
1051
+ # each decode steps due to the dynamic shapes. (`recording cudagraph tree for symint key 13`, etc.), which is
1052
+ # VERY slow. A workaround is `@torch.compiler.disable`, but this prevents using `fullgraph=True`.
1053
+ # See more context in https://github.com/huggingface/transformers/pull/29114
1054
 
1055
+ if self.config.attn_implementation == "flash_attention_2":
1056
+ if attention_mask is not None and 0.0 in attention_mask:
1057
+ return attention_mask
1058
+ return None
1059
+
1060
+ # For SDPA, when possible, we will rely on its `is_causal` argument instead of its `attn_mask` argument, in
1061
+ # order to dispatch on Flash Attention 2. This feature is not compatible with static cache, as SDPA will fail
1062
+ # to infer the attention mask.
1063
+ past_seen_tokens = past_key_values.get_seq_length() if past_key_values is not None else 0
1064
+ using_static_cache = isinstance(past_key_values, StaticCache)
1065
+
1066
+ # When output attentions is True, sdpa implementation's forward method calls the eager implementation's forward
1067
+ if self.config.attn_implementation == "sdpa" and not using_static_cache and not output_attentions:
1068
+ if AttentionMaskConverter._ignore_causal_mask_sdpa(
1069
+ attention_mask,
1070
+ inputs_embeds=input_tensor,
1071
+ past_key_values_length=past_seen_tokens,
1072
+ is_training=self.training,
1073
+ ):
1074
+ return None
1075
+
1076
+ dtype, device = input_tensor.dtype, input_tensor.device
1077
+ min_dtype = torch.finfo(dtype).min
1078
+ sequence_length = input_tensor.shape[1]
1079
+ if using_static_cache:
1080
+ target_length = past_key_values.get_max_length()
1081
+ else:
1082
+ target_length = (
1083
+ attention_mask.shape[-1]
1084
+ if isinstance(attention_mask, torch.Tensor)
1085
+ else past_seen_tokens + sequence_length + 1
1086
+ )
1087
+
1088
+ if attention_mask is not None and attention_mask.dim() == 4:
1089
+ # in this case we assume that the mask comes already in inverted form and requires no inversion or slicing
1090
+ if attention_mask.max() != 0:
1091
+ raise ValueError("Custom 4D attention mask should be passed in inverted form with max==0`")
1092
+ causal_mask = attention_mask
1093
+ else:
1094
+ causal_mask = torch.full((sequence_length, target_length), fill_value=min_dtype, dtype=dtype, device=device)
1095
+ if sequence_length != 1:
1096
+ causal_mask = torch.triu(causal_mask, diagonal=1)
1097
+ causal_mask *= torch.arange(target_length, device=device) > cache_position.reshape(-1, 1)
1098
+ causal_mask = causal_mask[None, None, :, :].expand(input_tensor.shape[0], 1, -1, -1)
1099
+ if attention_mask is not None:
1100
+ causal_mask = causal_mask.clone() # copy to contiguous memory for in-place edit
1101
+ mask_length = attention_mask.shape[-1]
1102
+ padding_mask = causal_mask[:, :, :, :mask_length] + attention_mask[:, None, None, :]
1103
+ padding_mask = padding_mask == 0
1104
+ causal_mask[:, :, :, :mask_length] = causal_mask[:, :, :, :mask_length].masked_fill(
1105
+ padding_mask, min_dtype
1106
+ )
1107
+ if (
1108
+ self.config.attn_implementation == "sdpa"
1109
+ and attention_mask is not None
1110
+ and attention_mask.device.type == "cuda"
1111
+ and not output_attentions
1112
+ ):
1113
+ # Attend to all tokens in fully masked rows in the causal_mask, for example the relevant first rows when
1114
+ # using left padding. This is required by F.scaled_dot_product_attention memory-efficient attention path.
1115
+ # Details: https://github.com/pytorch/pytorch/issues/110213
1116
+ causal_mask = AttentionMaskConverter._unmask_unattended(causal_mask, min_dtype) # pylint: disable=E1120
1117
+
1118
+ return causal_mask
1119
+
1120
+
1121
+ # Modified from transformers.models.llama.modeling_llama.LlamaForCausalLM
1122
  class InternLM2ForCausalLM(InternLM2PreTrainedModel):
1123
+ """Causal language model (CLM) for InternLM2."""
1124
 
1125
+ _auto_class = "AutoModelForCausalLM"
1126
  _tied_weights_keys = ["output.weight"]
1127
 
1128
  def __init__(self, config):
 
1159
  input_ids: torch.LongTensor = None,
1160
  attention_mask: Optional[torch.Tensor] = None,
1161
  position_ids: Optional[torch.LongTensor] = None,
1162
+ past_key_values: Optional[Union[Cache, List[torch.FloatTensor]]] = None,
1163
  inputs_embeds: Optional[torch.FloatTensor] = None,
1164
  labels: Optional[torch.LongTensor] = None,
1165
  use_cache: Optional[bool] = None,
1166
  output_attentions: Optional[bool] = None,
1167
  output_hidden_states: Optional[bool] = None,
1168
  return_dict: Optional[bool] = None,
1169
+ cache_position: Optional[torch.LongTensor] = None,
1170
  ) -> Union[Tuple, CausalLMOutputWithPast]:
1171
  r"""
1172
  Args:
 
1182
  ```python
1183
  >>> from transformers import AutoTokenizer, InternLM2ForCausalLM
1184
 
1185
+ >>> model = InternLM2ForCausalLM.from_pretrained("meta-InternLM2/InternLM2-2-7b-hf")
1186
+ >>> tokenizer = AutoTokenizer.from_pretrained("meta-InternLM2/InternLM2-2-7b-hf")
1187
 
1188
  >>> prompt = "Hey, are you conscious? Can you talk to me?"
1189
  >>> inputs = tokenizer(prompt, return_tensors="pt")
 
1211
  output_attentions=output_attentions,
1212
  output_hidden_states=output_hidden_states,
1213
  return_dict=return_dict,
1214
+ cache_position=cache_position,
1215
  )
1216
 
1217
  hidden_states = outputs[0]
1218
+ if self.config.pretraining_tp > 1:
1219
+ output_slices = self.output.weight.split(self.vocab_size // self.config.pretraining_tp, dim=0)
1220
+ logits = [
1221
+ F.linear(hidden_states, output_slices[i]) # pylint: disable=not-callable
1222
+ for i in range(self.config.pretraining_tp)
1223
+ ]
1224
+ logits = torch.cat(logits, dim=-1)
1225
+ else:
1226
+ logits = self.output(hidden_states)
1227
  logits = logits.float()
1228
 
1229
  loss = None
 
1252
  )
1253
 
1254
  def prepare_inputs_for_generation(
1255
+ self,
1256
+ input_ids,
1257
+ past_key_values=None,
1258
+ attention_mask=None,
1259
+ inputs_embeds=None,
1260
+ cache_position=None,
1261
+ use_cache=True,
1262
+ **kwargs,
1263
  ):
1264
+ past_length = 0
1265
  if past_key_values is not None:
1266
+ if isinstance(past_key_values, Cache):
1267
+ past_length = cache_position[0] if cache_position is not None else past_key_values.get_seq_length()
1268
+ max_cache_length = (
1269
+ torch.tensor(past_key_values.get_max_length(), device=input_ids.device)
1270
+ if past_key_values.get_max_length() is not None
1271
+ else None
1272
+ )
1273
+ cache_length = past_length if max_cache_length is None else torch.min(max_cache_length, past_length)
1274
+ # TODO joao: remove this `else` after `generate` prioritizes `Cache` objects
1275
  else:
1276
+ cache_length = past_length = past_key_values[0][0].shape[2]
1277
+ max_cache_length = None
1278
+
1279
+ # Keep only the unprocessed tokens:
1280
+ # 1 - If the length of the attention_mask exceeds the length of input_ids, then we are in a setting where
1281
+ # some of the inputs are exclusively passed as part of the cache (e.g. when passing input_embeds as input)
1282
+ if attention_mask is not None and attention_mask.shape[1] > input_ids.shape[1]:
1283
+ input_ids = input_ids[:, -(attention_mask.shape[1] - past_length) :]
1284
+ # 2 - If the past_length is smaller than input_ids', then input_ids holds all input tokens. We can discard
1285
+ # input_ids based on the past_length.
1286
+ elif past_length < input_ids.shape[1]:
1287
+ input_ids = input_ids[:, past_length:]
1288
+ # 3 - Otherwise (past_length >= input_ids.shape[1]), let's assume input_ids only has unprocessed tokens.
1289
+
1290
+ # If we are about to go beyond the maximum cache length, we need to crop the input attention mask.
1291
+ if (
1292
+ max_cache_length is not None
1293
+ and attention_mask is not None
1294
+ and cache_length + input_ids.shape[1] > max_cache_length
1295
+ ):
1296
+ attention_mask = attention_mask[:, -max_cache_length:] # pylint: disable=E1130
1297
 
1298
  position_ids = kwargs.get("position_ids", None)
1299
  if attention_mask is not None and position_ids is None:
 
1307
  if inputs_embeds is not None and past_key_values is None:
1308
  model_inputs = {"inputs_embeds": inputs_embeds}
1309
  else:
1310
+ # The `contiguous()` here is necessary to have a static stride during decoding. torchdynamo otherwise
1311
+ # recompiles graphs as the stride of the inputs is a guard.
1312
+ # Ref: https://github.com/huggingface/transformers/pull/29114
1313
+ # TODO: use `next_tokens` directly instead.
1314
+ model_inputs = {"input_ids": input_ids.contiguous()}
1315
+
1316
+ input_length = position_ids.shape[-1] if position_ids is not None else input_ids.shape[-1]
1317
+ if cache_position is None:
1318
+ cache_position = torch.arange(past_length, past_length + input_length, device=input_ids.device)
1319
+ elif use_cache:
1320
+ cache_position = cache_position[-input_length:]
1321
 
1322
  model_inputs.update(
1323
  {
1324
  "position_ids": position_ids,
1325
+ "cache_position": cache_position,
1326
  "past_key_values": past_key_values,
1327
+ "use_cache": use_cache,
1328
  "attention_mask": attention_mask,
1329
  }
1330
  )
 
1339
  )
1340
  return reordered_past
1341
 
1342
+ def build_inputs(self, tokenizer, query: str, history: List[Tuple[str, str]] = None, meta_instruction=""):
1343
+ if history is None:
1344
+ history = []
1345
  if tokenizer.add_bos_token:
1346
  prompt = ""
1347
  else:
 
1358
  self,
1359
  tokenizer,
1360
  query: str,
1361
+ history: Optional[List[Tuple[str, str]]] = None,
1362
  streamer: Optional[BaseStreamer] = None,
1363
  max_new_tokens: int = 1024,
1364
  do_sample: bool = True,
1365
  temperature: float = 0.8,
1366
  top_p: float = 0.8,
1367
  meta_instruction: str = "You are an AI assistant whose name is InternLM (书生·浦语).\n"
1368
+ "- InternLM (书生·浦语) is a conversational language model that is developed by Shanghai AI Laboratory "
1369
+ "(上海人工智能实验室). It is designed to be helpful, honest, and harmless.\n"
1370
+ "- InternLM (书生·浦语) can understand and communicate fluently in the language chosen by the user such "
1371
+ "as English and 中文.",
1372
  **kwargs,
1373
  ):
1374
+ if history is None:
1375
+ history = []
1376
  inputs = self.build_inputs(tokenizer, query, history, meta_instruction)
1377
  inputs = {k: v.to(self.device) for k, v in inputs.items() if torch.is_tensor(v)}
1378
  # also add end-of-assistant token in eos token id to avoid unnecessary generation
 
1398
  self,
1399
  tokenizer,
1400
  query: str,
1401
+ history: List[Tuple[str, str]] = None,
1402
  max_new_tokens: int = 1024,
1403
  do_sample: bool = True,
1404
  temperature: float = 0.8,
1405
  top_p: float = 0.8,
1406
  **kwargs,
1407
  ):
1408
+ if history is None:
1409
+ history = []
1410
  """
1411
  Return a generator in format: (response, history)
1412
  Eg.
 
1422
  response_queue = queue.Queue(maxsize=20)
1423
 
1424
  class ChatStreamer(BaseStreamer):
1425
+ """
1426
+ Streamer used in generate to print words one by one.
1427
+ """
1428
+
1429
  def __init__(self, tokenizer) -> None:
1430
  super().__init__()
1431
  self.tokenizer = tokenizer
 
1486
  return consumer()
1487
 
1488
 
1489
+ # Copied from transformers.models.llama.modeling_llama.LlamaForSequenceClassification with Llama->InternLM2
1490
  @add_start_docstrings(
1491
  """
1492
  The InternLM2 Model transformer with a sequence classification head on top (linear layer).
1493
 
1494
+ [`InternLM2ForSequenceClassification`] uses the last token in order to do the classification, as other causal models
1495
+ (e.g. GPT-2) do.
1496
 
1497
  Since it does classification on the last token, it requires to know the position of the last token. If a
1498
  `pad_token_id` is defined in the configuration, it finds the last token that is not a padding token in each row. If
 
1503
  InternLM2_START_DOCSTRING,
1504
  )
1505
  class InternLM2ForSequenceClassification(InternLM2PreTrainedModel):
1506
+ """Sequence Classification Head for InternLM2 Model."""
1507
+
1508
  def __init__(self, config):
1509
  super().__init__(config)
1510
  self.num_labels = config.num_labels
 
1526
  input_ids: torch.LongTensor = None,
1527
  attention_mask: Optional[torch.Tensor] = None,
1528
  position_ids: Optional[torch.LongTensor] = None,
1529
+ past_key_values: Optional[Union[Cache, List[torch.FloatTensor]]] = None,
1530
  inputs_embeds: Optional[torch.FloatTensor] = None,
1531
  labels: Optional[torch.LongTensor] = None,
1532
  use_cache: Optional[bool] = None,
 
1567
  sequence_lengths = -1
1568
  else:
1569
  if input_ids is not None:
1570
+ # if no pad token found, use modulo instead of reverse indexing for ONNX compatibility
1571
+ sequence_lengths = torch.eq(input_ids, self.config.pad_token_id).int().argmax(-1) - 1
1572
+ sequence_lengths = sequence_lengths % input_ids.shape[-1]
1573
+ sequence_lengths = sequence_lengths.to(logits.device)
1574
  else:
1575
  sequence_lengths = -1
1576
 
 
1582
  if self.config.problem_type is None:
1583
  if self.num_labels == 1:
1584
  self.config.problem_type = "regression"
1585
+ elif self.num_labels > 1 and (labels.dtype in (torch.long, torch.int)):
1586
  self.config.problem_type = "single_label_classification"
1587
  else:
1588
  self.config.problem_type = "multi_label_classification"
 
1610
  hidden_states=transformer_outputs.hidden_states,
1611
  attentions=transformer_outputs.attentions,
1612
  )
1613
+
1614
+
1615
+ # Copied from transformers.models.llama.modeling_llama.LlamaForQuestionAnswering with Llama->InternLM2
1616
+ @add_start_docstrings(
1617
+ """
1618
+ The InternLM2 Model transformer with a span classification head on top for extractive question-answering tasks like
1619
+ SQuAD (a linear layer on top of the hidden-states output to compute `span start logits` and `span end logits`).
1620
+ """,
1621
+ InternLM2_START_DOCSTRING,
1622
+ )
1623
+ class InternLM2ForQuestionAnswering(InternLM2PreTrainedModel):
1624
+ """Question Answering model for InternLM2."""
1625
+
1626
+ base_model_prefix = "transformer"
1627
+
1628
+ def __init__(self, config):
1629
+ super().__init__(config)
1630
+ self.transformer = InternLM2Model(config)
1631
+ self.qa_outputs = nn.Linear(config.hidden_size, 2)
1632
+
1633
+ # Initialize weights and apply final processing
1634
+ self.post_init()
1635
+
1636
+ def get_input_embeddings(self):
1637
+ return self.transformer.tok_embeddings
1638
+
1639
+ def set_input_embeddings(self, value):
1640
+ self.transformer.tok_embeddings = value
1641
+
1642
+ @add_start_docstrings_to_model_forward(InternLM2_INPUTS_DOCSTRING)
1643
+ def forward(
1644
+ self,
1645
+ input_ids: Optional[torch.LongTensor] = None,
1646
+ attention_mask: Optional[torch.FloatTensor] = None,
1647
+ position_ids: Optional[torch.LongTensor] = None,
1648
+ past_key_values: Optional[Union[Cache, List[torch.FloatTensor]]] = None,
1649
+ inputs_embeds: Optional[torch.FloatTensor] = None,
1650
+ start_positions: Optional[torch.LongTensor] = None,
1651
+ end_positions: Optional[torch.LongTensor] = None,
1652
+ output_attentions: Optional[bool] = None,
1653
+ output_hidden_states: Optional[bool] = None,
1654
+ return_dict: Optional[bool] = None,
1655
+ ) -> Union[Tuple, QuestionAnsweringModelOutput]:
1656
+ r"""
1657
+ start_positions (`torch.LongTensor` of shape `(batch_size,)`, *optional*):
1658
+ Labels for position (index) of the start of the labelled span for computing the token classification loss.
1659
+ Positions are clamped to the length of the sequence (`sequence_length`). Position outside of the sequence
1660
+ are not taken into account for computing the loss.
1661
+ end_positions (`torch.LongTensor` of shape `(batch_size,)`, *optional*):
1662
+ Labels for position (index) of the end of the labelled span for computing the token classification loss.
1663
+ Positions are clamped to the length of the sequence (`sequence_length`). Position outside of the sequence
1664
+ are not taken into account for computing the loss.
1665
+ """
1666
+ return_dict = return_dict if return_dict is not None else self.config.use_return_dict
1667
+
1668
+ outputs = self.transformer(
1669
+ input_ids,
1670
+ attention_mask=attention_mask,
1671
+ position_ids=position_ids,
1672
+ past_key_values=past_key_values,
1673
+ inputs_embeds=inputs_embeds,
1674
+ output_attentions=output_attentions,
1675
+ output_hidden_states=output_hidden_states,
1676
+ return_dict=return_dict,
1677
+ )
1678
+
1679
+ sequence_output = outputs[0]
1680
+
1681
+ logits = self.qa_outputs(sequence_output)
1682
+ start_logits, end_logits = logits.split(1, dim=-1)
1683
+ start_logits = start_logits.squeeze(-1).contiguous()
1684
+ end_logits = end_logits.squeeze(-1).contiguous()
1685
+
1686
+ total_loss = None
1687
+ if start_positions is not None and end_positions is not None:
1688
+ # If we are on multi-GPU, split add a dimension
1689
+ if len(start_positions.size()) > 1:
1690
+ start_positions = start_positions.squeeze(-1).to(start_logits.device)
1691
+ if len(end_positions.size()) > 1:
1692
+ end_positions = end_positions.squeeze(-1).to(end_logits.device)
1693
+ # sometimes the start/end positions are outside our model inputs, we ignore these terms
1694
+ ignored_index = start_logits.size(1)
1695
+ start_positions = start_positions.clamp(0, ignored_index)
1696
+ end_positions = end_positions.clamp(0, ignored_index)
1697
+
1698
+ loss_fct = CrossEntropyLoss(ignore_index=ignored_index)
1699
+ start_loss = loss_fct(start_logits, start_positions)
1700
+ end_loss = loss_fct(end_logits, end_positions)
1701
+ total_loss = (start_loss + end_loss) / 2
1702
+
1703
+ if not return_dict:
1704
+ output = (start_logits, end_logits) + outputs[2:]
1705
+ return ((total_loss,) + output) if total_loss is not None else output
1706
+
1707
+ return QuestionAnsweringModelOutput(
1708
+ loss=total_loss,
1709
+ start_logits=start_logits,
1710
+ end_logits=end_logits,
1711
+ hidden_states=outputs.hidden_states,
1712
+ attentions=outputs.attentions,
1713
+ )
1714
+
1715
+
1716
+ # Copied from transformers.models.llama.modeling_llama.LlamaForTokenClassification with Llama->InternLM2
1717
+ @add_start_docstrings(
1718
+ """
1719
+ The InternLM2 Model transformer with a token classification head on top (a linear layer on top of the hidden-states
1720
+ output) e.g. for Named-Entity-Recognition (NER) tasks.
1721
+ """,
1722
+ InternLM2_START_DOCSTRING,
1723
+ )
1724
+ class InternLM2ForTokenClassification(InternLM2PreTrainedModel):
1725
+ """Token classification model for InternLM2."""
1726
+
1727
+ def __init__(self, config):
1728
+ super().__init__(config)
1729
+ self.num_labels = config.num_labels
1730
+ self.model = InternLM2Model(config)
1731
+ if getattr(config, "classifier_dropout", None) is not None:
1732
+ classifier_dropout = config.classifier_dropout
1733
+ elif getattr(config, "hidden_dropout", None) is not None:
1734
+ classifier_dropout = config.hidden_dropout
1735
+ else:
1736
+ classifier_dropout = 0.1
1737
+ self.dropout = nn.Dropout(classifier_dropout)
1738
+ self.score = nn.Linear(config.hidden_size, config.num_labels)
1739
+
1740
+ # Initialize weights and apply final processing
1741
+ self.post_init()
1742
+
1743
+ def get_input_embeddings(self):
1744
+ return self.model.tok_embeddings
1745
+
1746
+ def set_input_embeddings(self, value):
1747
+ self.model.tok_embeddings = value
1748
+
1749
+ @add_start_docstrings_to_model_forward(InternLM2_INPUTS_DOCSTRING)
1750
+ def forward(
1751
+ self,
1752
+ input_ids: torch.LongTensor = None,
1753
+ attention_mask: Optional[torch.Tensor] = None,
1754
+ position_ids: Optional[torch.LongTensor] = None,
1755
+ past_key_values: Optional[List[torch.FloatTensor]] = None,
1756
+ inputs_embeds: Optional[torch.FloatTensor] = None,
1757
+ labels: Optional[torch.LongTensor] = None,
1758
+ use_cache: Optional[bool] = None,
1759
+ output_attentions: Optional[bool] = None,
1760
+ output_hidden_states: Optional[bool] = None,
1761
+ return_dict: Optional[bool] = None,
1762
+ ) -> Union[Tuple, SequenceClassifierOutputWithPast]:
1763
+ r"""
1764
+ labels (`torch.LongTensor` of shape `(batch_size,)`, *optional*):
1765
+ Labels for computing the sequence classification/regression loss. Indices should be in `[0, ...,
1766
+ config.num_labels - 1]`. If `config.num_labels == 1` a regression loss is computed (Mean-Square loss), If
1767
+ `config.num_labels > 1` a classification loss is computed (Cross-Entropy).
1768
+ """
1769
+ return_dict = return_dict if return_dict is not None else self.config.use_return_dict
1770
+
1771
+ outputs = self.model(
1772
+ input_ids,
1773
+ attention_mask=attention_mask,
1774
+ position_ids=position_ids,
1775
+ past_key_values=past_key_values,
1776
+ inputs_embeds=inputs_embeds,
1777
+ use_cache=use_cache,
1778
+ output_attentions=output_attentions,
1779
+ output_hidden_states=output_hidden_states,
1780
+ return_dict=return_dict,
1781
+ )
1782
+ sequence_output = outputs[0]
1783
+ sequence_output = self.dropout(sequence_output)
1784
+ logits = self.score(sequence_output)
1785
+
1786
+ loss = None
1787
+ if labels is not None:
1788
+ loss_fct = CrossEntropyLoss()
1789
+ loss = loss_fct(logits.view(-1, self.num_labels), labels.view(-1))
1790
+
1791
+ if not return_dict:
1792
+ output = (logits,) + outputs[2:]
1793
+ return ((loss,) + output) if loss is not None else output
1794
+
1795
+ return TokenClassifierOutput(
1796
+ loss=loss,
1797
+ logits=logits,
1798
+ hidden_states=outputs.hidden_states,
1799
+ attentions=outputs.attentions,
1800
+ )