MathLLMs/MathCoder-VL-8B
677
1# Copyright (c) The InternLM team and The HuggingFace Inc. team. All rights reserved.2#3# This code is based on transformers/src/transformers/models/llama/modeling_llama.py4#5# Licensed under the Apache License, Version 2.0 (the "License");6# you may not use this file except in compliance with the License.7# You may obtain a copy of the License at8#9# http://www.apache.org/licenses/LICENSE-2.010#11# Unless required by applicable law or agreed to in writing, software12# distributed under the License is distributed on an "AS IS" BASIS,13# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.14# See the License for the specific language governing permissions and15# limitations under the License.16""" PyTorch InternLM2 model."""17import math18import queue19import threading20import warnings21from typing import List, Optional, Tuple, Union22 23import torch24import torch.nn.functional as F25import torch.utils.checkpoint26from einops import rearrange27from torch import nn28from torch.nn import BCEWithLogitsLoss, CrossEntropyLoss, MSELoss29from transformers.activations import ACT2FN30from transformers.modeling_outputs import (BaseModelOutputWithPast,31 CausalLMOutputWithPast,32 SequenceClassifierOutputWithPast)33from transformers.modeling_utils import PreTrainedModel34from transformers.utils import (add_start_docstrings,35 add_start_docstrings_to_model_forward, logging,36 replace_return_docstrings)37 38try:39 from transformers.generation.streamers import BaseStreamer40except: # noqa # pylint: disable=bare-except41 BaseStreamer = None42 43from .configuration_internlm2 import InternLM2Config44 45logger = logging.get_logger(__name__)46 47_CONFIG_FOR_DOC = 'InternLM2Config'48 49flash_attn_func, flash_attn_varlen_func = None, None50pad_input, index_first_axis, unpad_input = None, None, None51try:52 from flash_attn import flash_attn_func as _flash_attn_func53 from flash_attn import flash_attn_varlen_func as _flash_attn_varlen_func54 from flash_attn.bert_padding import index_first_axis as _index_first_axis55 from flash_attn.bert_padding import pad_input as _pad_input56 from flash_attn.bert_padding import unpad_input as _unpad_input57 58 flash_attn_func, flash_attn_varlen_func = _flash_attn_func, _flash_attn_varlen_func59 pad_input, index_first_axis, unpad_input = _pad_input, _index_first_axis, _unpad_input60 has_flash_attn = True61except:62 has_flash_attn = False63 64 65def _import_flash_attn():66 global flash_attn_func, flash_attn_varlen_func67 global pad_input, index_first_axis, unpad_input68 try:69 from flash_attn import flash_attn_func as _flash_attn_func70 from flash_attn import \71 flash_attn_varlen_func as _flash_attn_varlen_func72 from flash_attn.bert_padding import \73 index_first_axis as _index_first_axis74 from flash_attn.bert_padding import pad_input as _pad_input75 from flash_attn.bert_padding import unpad_input as _unpad_input76 flash_attn_func, flash_attn_varlen_func = _flash_attn_func, _flash_attn_varlen_func77 pad_input, index_first_axis, unpad_input = _pad_input, _index_first_axis, _unpad_input78 except ImportError:79 raise ImportError('flash_attn is not installed.')80 81 82# Copied from transformers.models.llama.modeling_llama._get_unpad_data83def _get_unpad_data(attention_mask):84 seqlens_in_batch = attention_mask.sum(dim=-1, dtype=torch.int32)85 indices = torch.nonzero(attention_mask.flatten(), as_tuple=False).flatten()86 max_seqlen_in_batch = seqlens_in_batch.max().item()87 cu_seqlens = F.pad(torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.torch.int32), (1, 0))88 return (89 indices,90 cu_seqlens,91 max_seqlen_in_batch,92 )93 94 95# Copied from transformers.models.bart.modeling_bart._make_causal_mask96def _make_causal_mask(97 input_ids_shape: torch.Size, dtype: torch.dtype, device: torch.device, past_key_values_length: int = 098):99 """100 Make causal mask used for bi-directional self-attention.101 """102 bsz, tgt_len = input_ids_shape103 mask = torch.full((tgt_len, tgt_len), torch.tensor(torch.finfo(dtype).min, device=device), device=device)104 mask_cond = torch.arange(mask.size(-1), device=device)105 mask.masked_fill_(mask_cond < (mask_cond + 1).view(mask.size(-1), 1), 0)106 mask = mask.to(dtype)107 108 if past_key_values_length > 0:109 mask = torch.cat([torch.zeros(tgt_len, past_key_values_length, dtype=dtype, device=device), mask], dim=-1)110 return mask[None, None, :, :].expand(bsz, 1, tgt_len, tgt_len + past_key_values_length)111 112 113# Copied from transformers.models.bart.modeling_bart._expand_mask114def _expand_mask(mask: torch.Tensor, dtype: torch.dtype, tgt_len: Optional[int] = None):115 """116 Expands attention_mask from `[bsz, seq_len]` to `[bsz, 1, tgt_seq_len, src_seq_len]`.117 """118 bsz, src_len = mask.size()119 tgt_len = tgt_len if tgt_len is not None else src_len120 121 expanded_mask = mask[:, None, None, :].expand(bsz, 1, tgt_len, src_len).to(dtype)122 123 inverted_mask = 1.0 - expanded_mask124 125 return inverted_mask.masked_fill(inverted_mask.to(torch.bool), torch.finfo(dtype).min)126 127 128# Copied from transformers.models.llama.modeling_llama.LlamaRMSNorm with Llama->InternLM2129class InternLM2RMSNorm(nn.Module):130 def __init__(self, hidden_size, eps=1e-6):131 """132 InternLM2RMSNorm is equivalent to T5LayerNorm133 """134 super().__init__()135 self.weight = nn.Parameter(torch.ones(hidden_size))136 self.variance_epsilon = eps137 138 def forward(self, hidden_states):139 input_dtype = hidden_states.dtype140 hidden_states = hidden_states.to(torch.float32)141 variance = hidden_states.pow(2).mean(-1, keepdim=True)142 hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)143 return self.weight * hidden_states.to(input_dtype)144 145 146# Copied from transformers.model.llama.modeling_llama.LlamaRotaryEmbedding with Llama->InternLM2147class InternLM2RotaryEmbedding(nn.Module):148 def __init__(self, dim, max_position_embeddings=2048, base=10000, device=None):149 super().__init__()150 151 self.dim = dim152 self.max_position_embeddings = max_position_embeddings153 self.base = base154 inv_freq = 1.0 / (self.base ** (torch.arange(0, self.dim, 2).float().to(device) / self.dim))155 self.register_buffer('inv_freq', inv_freq, persistent=False)156 157 # Build here to make `torch.jit.trace` work.158 self._set_cos_sin_cache(159 seq_len=max_position_embeddings, device=self.inv_freq.device, dtype=torch.get_default_dtype()160 )161 162 def _set_cos_sin_cache(self, seq_len, device, dtype):163 self.max_seq_len_cached = seq_len164 t = torch.arange(self.max_seq_len_cached, device=device).to(dtype=self.inv_freq.dtype)165 166 freqs = torch.einsum('i,j->ij', t, self.inv_freq)167 # Different from paper, but it uses a different permutation in order to obtain the same calculation168 emb = torch.cat((freqs, freqs), dim=-1)169 self.register_buffer('cos_cached', emb.cos().to(dtype), persistent=False)170 self.register_buffer('sin_cached', emb.sin().to(dtype), persistent=False)171 172 def forward(self, x, seq_len=None):173 # x: [bs, num_attention_heads, seq_len, head_size]174 if seq_len > self.max_seq_len_cached:175 self._set_cos_sin_cache(seq_len=seq_len, device=x.device, dtype=torch.float32)176 177 return (178 self.cos_cached[:seq_len].to(dtype=x.dtype),179 self.sin_cached[:seq_len].to(dtype=x.dtype),180 )181 182 183# Copied from transformers.model.llama.modeling_llama.LlamaLinearScalingRotaryEmbedding with Llama->InternLM2184class InternLM2LinearScalingRotaryEmbedding(InternLM2RotaryEmbedding):185 """InternLM2RotaryEmbedding extended with linear scaling. Credits to the Reddit user /u/kaiokendev"""186 187 def __init__(self, dim, max_position_embeddings=2048, base=10000, device=None, scaling_factor=1.0):188 self.scaling_factor = scaling_factor189 super().__init__(dim, max_position_embeddings, base, device)190 191 def _set_cos_sin_cache(self, seq_len, device, dtype):192 self.max_seq_len_cached = seq_len193 t = torch.arange(self.max_seq_len_cached, device=device).to(dtype=self.inv_freq.dtype)194 t = t / self.scaling_factor195 196 freqs = torch.einsum('i,j->ij', t, self.inv_freq)197 # Different from paper, but it uses a different permutation in order to obtain the same calculation198 emb = torch.cat((freqs, freqs), dim=-1)199 self.register_buffer('cos_cached', emb.cos().to(dtype), persistent=False)200 self.register_buffer('sin_cached', emb.sin().to(dtype), persistent=False)201 202 203# Copied from transformers.model.llama.modeling_llama.LlamaDynamicNTKScalingRotaryEmbedding with Llama->InternLM2204class InternLM2DynamicNTKScalingRotaryEmbedding(InternLM2RotaryEmbedding):205 """InternLM2RotaryEmbedding extended with Dynamic NTK scaling.206 Credits to the Reddit users /u/bloc97 and /u/emozilla.207 """208 209 def __init__(self, dim, max_position_embeddings=2048, base=10000, device=None, scaling_factor=1.0):210 self.scaling_factor = scaling_factor211 super().__init__(dim, max_position_embeddings, base, device)212 213 def _set_cos_sin_cache(self, seq_len, device, dtype):214 self.max_seq_len_cached = seq_len215 216 if seq_len > self.max_position_embeddings:217 base = self.base * (218 (self.scaling_factor * seq_len / self.max_position_embeddings) - (self.scaling_factor - 1)219 ) ** (self.dim / (self.dim - 2))220 inv_freq = 1.0 / (base ** (torch.arange(0, self.dim, 2).float().to(device) / self.dim))221 self.register_buffer('inv_freq', inv_freq, persistent=False)222 223 t = torch.arange(self.max_seq_len_cached, device=device).to(dtype=self.inv_freq.dtype)224 225 freqs = torch.einsum('i,j->ij', t, self.inv_freq)226 # Different from paper, but it uses a different permutation in order to obtain the same calculation227 emb = torch.cat((freqs, freqs), dim=-1)228 self.register_buffer('cos_cached', emb.cos().to(dtype), persistent=False)229 self.register_buffer('sin_cached', emb.sin().to(dtype), persistent=False)230 231 232# Copied from transformers.model.llama.modeling_llama.rotate_half233def rotate_half(x):234 """Rotates half the hidden dims of the input."""235 x1 = x[..., : x.shape[-1] // 2]236 x2 = x[..., x.shape[-1] // 2 :]237 return torch.cat((-x2, x1), dim=-1)238 239 240# Copied from transformers.model.llama.modeling_llama.apply_rotary_pos_emb241def apply_rotary_pos_emb(q, k, cos, sin, position_ids, unsqueeze_dim=1):242 """Applies Rotary Position Embedding to the query and key tensors."""243 cos = cos[position_ids].unsqueeze(unsqueeze_dim)244 sin = sin[position_ids].unsqueeze(unsqueeze_dim)245 q_embed = (q * cos) + (rotate_half(q) * sin)246 k_embed = (k * cos) + (rotate_half(k) * sin)247 return q_embed, k_embed248 249 250class InternLM2MLP(nn.Module):251 def __init__(self, config):252 super().__init__()253 self.config = config254 self.hidden_size = config.hidden_size255 self.intermediate_size = config.intermediate_size256 self.w1 = nn.Linear(self.hidden_size, self.intermediate_size, bias=False)257 self.w3 = nn.Linear(self.hidden_size, self.intermediate_size, bias=False)258 self.w2 = nn.Linear(self.intermediate_size, self.hidden_size, bias=False)259 self.act_fn = ACT2FN[config.hidden_act]260 261 def forward(self, x):262 down_proj = self.w2(self.act_fn(self.w1(x)) * self.w3(x))263 264 return down_proj265 266 267# Copied from transformers.model.llama.modeling_llama.repeat_kv268def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor:269 """270 This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch,271 num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim)272 """273 batch, num_key_value_heads, slen, head_dim = hidden_states.shape274 if n_rep == 1:275 return hidden_states276 hidden_states = hidden_states[:, :, None, :, :].expand(batch, num_key_value_heads, n_rep, slen, head_dim)277 return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim)278 279 280# Modified from transformers.model.llama.modeling_llama.LlamaAttention281class InternLM2Attention(nn.Module):282 """Multi-headed attention from 'Attention Is All You Need' paper"""283 284 def __init__(self, config: InternLM2Config):285 super().__init__()286 self.config = config287 self.hidden_size = config.hidden_size288 self.num_heads = config.num_attention_heads289 self.head_dim = self.hidden_size // self.num_heads290 self.num_key_value_heads = config.num_key_value_heads291 self.num_key_value_groups = self.num_heads // self.num_key_value_heads292 self.max_position_embeddings = config.max_position_embeddings293 self.is_causal = True294 295 if (self.head_dim * self.num_heads) != self.hidden_size:296 raise ValueError(297 f'hidden_size must be divisible by num_heads (got `hidden_size`: {self.hidden_size}'298 f' and `num_heads`: {self.num_heads}).'299 )300 301 self.wqkv = nn.Linear(302 self.hidden_size,303 (self.num_heads + 2 * self.num_key_value_heads) * self.head_dim,304 bias=config.bias,305 )306 307 self.wo = nn.Linear(self.num_heads * self.head_dim, self.hidden_size, bias=config.bias)308 self._init_rope()309 310 def _init_rope(self):311 if self.config.rope_scaling is None:312 self.rotary_emb = InternLM2RotaryEmbedding(313 self.head_dim,314 max_position_embeddings=self.max_position_embeddings,315 base=self.config.rope_theta,316 )317 else:318 scaling_type = self.config.rope_scaling['type']319 scaling_factor = self.config.rope_scaling['factor']320 if scaling_type == 'dynamic':321 self.rotary_emb = InternLM2DynamicNTKScalingRotaryEmbedding(322 self.head_dim,323 max_position_embeddings=self.max_position_embeddings,324 base=self.config.rope_theta,325 scaling_factor=scaling_factor,326 )327 elif scaling_type == 'linear':328 self.rotary_emb = InternLM2LinearScalingRotaryEmbedding(329 self.head_dim,330 max_position_embeddings=self.max_position_embeddings,331 base=self.config.rope_theta,332 scaling_factor=scaling_factor,333 )334 else:335 raise ValueError("Currently we only support rotary embedding's type being 'dynamic' or 'linear'.")336 return self.rotary_emb337 338 def _shape(self, tensor: torch.Tensor, seq_len: int, bsz: int):339 return tensor.view(bsz, seq_len, self.num_heads, self.head_dim).transpose(1, 2).contiguous()340 341 def forward(342 self,343 hidden_states: torch.Tensor,344 attention_mask: Optional[torch.Tensor] = None,345 position_ids: Optional[torch.LongTensor] = None,346 past_key_value: Optional[Tuple[torch.Tensor]] = None,347 output_attentions: bool = False,348 use_cache: bool = False,349 **kwargs,350 ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]:351 if 'padding_mask' in kwargs:352 warnings.warn(353 'Passing `padding_mask` is deprecated and will be removed in v4.37. '354 'Please make sure use `attention_mask` instead.`'355 )356 357 bsz, q_len, _ = hidden_states.size()358 359 qkv_states = self.wqkv(hidden_states)360 361 qkv_states = rearrange(362 qkv_states,363 'b q (h gs d) -> b q h gs d',364 gs=2 + self.num_key_value_groups,365 d=self.head_dim,366 )367 368 query_states = qkv_states[..., : self.num_key_value_groups, :]369 query_states = rearrange(query_states, 'b q h gs d -> b q (h gs) d')370 key_states = qkv_states[..., -2, :]371 value_states = qkv_states[..., -1, :]372 373 query_states = query_states.transpose(1, 2)374 key_states = key_states.transpose(1, 2)375 value_states = value_states.transpose(1, 2)376 377 kv_seq_len = key_states.shape[-2]378 if past_key_value is not None:379 kv_seq_len += past_key_value[0].shape[-2]380 cos, sin = self.rotary_emb(value_states, seq_len=kv_seq_len)381 query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin, position_ids)382 383 if past_key_value is not None:384 # reuse k, v, self_attention385 key_states = torch.cat([past_key_value[0], key_states], dim=2)386 value_states = torch.cat([past_key_value[1], value_states], dim=2)387 388 past_key_value = (key_states, value_states) if use_cache else None389 390 key_states = repeat_kv(key_states, self.num_key_value_groups)391 value_states = repeat_kv(value_states, self.num_key_value_groups)392 393 attn_weights = torch.matmul(query_states, key_states.transpose(2, 3)) / math.sqrt(self.head_dim)394 395 if attn_weights.size() != (bsz, self.num_heads, q_len, kv_seq_len):396 raise ValueError(397 f'Attention weights should be of size {(bsz, self.num_heads, q_len, kv_seq_len)}, but is'398 f' {attn_weights.size()}'399 )400 401 if attention_mask is not None:402 if attention_mask.size() != (bsz, 1, q_len, kv_seq_len):403 raise ValueError(404 f'Attention mask should be of size {(bsz, 1, q_len, kv_seq_len)}, but is {attention_mask.size()}'405 )406 attn_weights = attn_weights + attention_mask407 408 # upcast attention to fp32409 attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query_states.dtype)410 attn_output = torch.matmul(attn_weights, value_states)411 412 if attn_output.size() != (bsz, self.num_heads, q_len, self.head_dim):413 raise ValueError(414 f'`attn_output` should be of size {(bsz, self.num_heads, q_len, self.head_dim)}, but is'415 f' {attn_output.size()}'416 )417 418 attn_output = attn_output.transpose(1, 2).contiguous()419 attn_output = attn_output.reshape(bsz, q_len, self.hidden_size)420 421 attn_output = self.wo(attn_output)422 423 if not output_attentions:424 attn_weights = None425 426 return attn_output, attn_weights, past_key_value427 428 429# Modified from transformers.model.llama.modeling_llama.InternLM2FlashAttention2430class InternLM2FlashAttention2(InternLM2Attention):431 """432 InternLM2 flash attention module. This module inherits from `InternLM2Attention` as the weights of the module stays433 untouched. The only required change would be on the forward pass where it needs to correctly call the public API of434 flash attention and deal with padding tokens in case the input contains any of them.435 """436 437 def forward(438 self,439 hidden_states: torch.Tensor,440 attention_mask: Optional[torch.LongTensor] = None,441 position_ids: Optional[torch.LongTensor] = None,442 past_key_value: Optional[Tuple[torch.Tensor]] = None,443 output_attentions: bool = False,444 use_cache: bool = False,445 **kwargs,446 ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]:447 # InternLM2FlashAttention2 attention does not support output_attentions448 if 'padding_mask' in kwargs:449 warnings.warn(450 'Passing `padding_mask` is deprecated and will be removed in v4.37. '451 'Please make sure use `attention_mask` instead.`'452 )453 454 # overwrite attention_mask with padding_mask455 attention_mask = kwargs.pop('padding_mask')456 457 output_attentions = False458 459 bsz, q_len, _ = hidden_states.size()460 461 qkv_states = self.wqkv(hidden_states)462 463 qkv_states = rearrange(464 qkv_states,465 'b q (h gs d) -> b q h gs d',466 gs=2 + self.num_key_value_groups,467 d=self.head_dim,468 )469 470 query_states = qkv_states[..., : self.num_key_value_groups, :]471 query_states = rearrange(query_states, 'b q h gs d -> b q (h gs) d')472 key_states = qkv_states[..., -2, :]473 value_states = qkv_states[..., -1, :]474 475 query_states = query_states.transpose(1, 2)476 key_states = key_states.transpose(1, 2)477 value_states = value_states.transpose(1, 2)478 479 kv_seq_len = key_states.shape[-2]480 if past_key_value is not None:481 kv_seq_len += past_key_value[0].shape[-2]482 483 cos, sin = self.rotary_emb(value_states, seq_len=kv_seq_len)484 485 query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin, position_ids)486 487 if past_key_value is not None:488 # reuse k, v, self_attention489 key_states = torch.cat([past_key_value[0], key_states], dim=2)490 value_states = torch.cat([past_key_value[1], value_states], dim=2)491 492 past_key_value = (key_states, value_states) if use_cache else None493 494 query_states = query_states.transpose(1, 2)495 key_states = key_states.transpose(1, 2)496 value_states = value_states.transpose(1, 2)497 498 attn_output = self._flash_attention_forward(499 query_states, key_states, value_states, attention_mask, q_len500 )501 attn_output = attn_output.reshape(bsz, q_len, self.hidden_size).contiguous()502 attn_output = self.wo(attn_output)503 504 if not output_attentions:505 attn_weights = None506 507 return attn_output, attn_weights, past_key_value508 509 def _flash_attention_forward(510 self, query_states, key_states, value_states, attention_mask, query_length, dropout=0.0, softmax_scale=None511 ):512 """513 Calls the forward method of Flash Attention - if the input hidden states contain at least one padding token514 first unpad the input, then computes the attention scores and pad the final attention scores.515 516 Args:517 query_states (`torch.Tensor`):518 Input query states to be passed to Flash Attention API519 key_states (`torch.Tensor`):520 Input key states to be passed to Flash Attention API521 value_states (`torch.Tensor`):522 Input value states to be passed to Flash Attention API523 attention_mask (`torch.Tensor`):524 The padding mask - corresponds to a tensor of size `(batch_size, seq_len)` where 0 stands for the525 position of padding tokens and 1 for the position of non-padding tokens.526 dropout (`int`, *optional*):527 Attention dropout528 softmax_scale (`float`, *optional*):529 The scaling of QK^T before applying softmax. Default to 1 / sqrt(head_dim)530 """531 # Contains at least one padding token in the sequence532 causal = self.is_causal and query_length != 1533 if attention_mask is not None:534 batch_size = query_states.shape[0]535 query_states, key_states, value_states, indices_q, cu_seq_lens, max_seq_lens = self._unpad_input(536 query_states, key_states, value_states, attention_mask, query_length537 )538 539 cu_seqlens_q, cu_seqlens_k = cu_seq_lens540 max_seqlen_in_batch_q, max_seqlen_in_batch_k = max_seq_lens541 542 attn_output_unpad = flash_attn_varlen_func(543 query_states,544 key_states,545 value_states,546 cu_seqlens_q=cu_seqlens_q,547 cu_seqlens_k=cu_seqlens_k,548 max_seqlen_q=max_seqlen_in_batch_q,549 max_seqlen_k=max_seqlen_in_batch_k,550 dropout_p=dropout,551 softmax_scale=softmax_scale,552 causal=causal,553 )554 555 attn_output = pad_input(attn_output_unpad, indices_q, batch_size, query_length)556 else:557 attn_output = flash_attn_func(558 query_states, key_states, value_states, dropout, softmax_scale=softmax_scale, causal=causal559 )560 561 return attn_output562 563 def _unpad_input(self, query_layer, key_layer, value_layer, attention_mask, query_length):564 indices_k, cu_seqlens_k, max_seqlen_in_batch_k = _get_unpad_data(attention_mask)565 batch_size, kv_seq_len, num_key_value_heads, head_dim = key_layer.shape566 567 key_layer = index_first_axis(568 key_layer.reshape(batch_size * kv_seq_len, num_key_value_heads, head_dim), indices_k569 )570 value_layer = index_first_axis(571 value_layer.reshape(batch_size * kv_seq_len, num_key_value_heads, head_dim), indices_k572 )573 574 if query_length == kv_seq_len:575 query_layer = index_first_axis(576 query_layer.reshape(batch_size * kv_seq_len, self.num_heads, head_dim), indices_k577 )578 cu_seqlens_q = cu_seqlens_k579 max_seqlen_in_batch_q = max_seqlen_in_batch_k580 indices_q = indices_k581 elif query_length == 1:582 max_seqlen_in_batch_q = 1583 cu_seqlens_q = torch.arange(584 batch_size + 1, dtype=torch.int32, device=query_layer.device585 ) # There is a memcpy here, that is very bad.586 indices_q = cu_seqlens_q[:-1]587 query_layer = query_layer.squeeze(1)588 else:589 # The -q_len: slice assumes left padding.590 attention_mask = attention_mask[:, -query_length:]591 query_layer, indices_q, cu_seqlens_q, max_seqlen_in_batch_q = unpad_input(query_layer, attention_mask)592 593 return (594 query_layer,595 key_layer,596 value_layer,597 indices_q.to(torch.int64),598 (cu_seqlens_q, cu_seqlens_k),599 (max_seqlen_in_batch_q, max_seqlen_in_batch_k),600 )601 602 603INTERNLM2_ATTENTION_CLASSES = {604 'eager': InternLM2Attention,605 'flash_attention_2': InternLM2FlashAttention2,606}607 608 609# Modified from transformers.model.llama.modeling_llama.LlamaDecoderLayer610class InternLM2DecoderLayer(nn.Module):611 def __init__(self, config: InternLM2Config):612 super().__init__()613 self.hidden_size = config.hidden_size614 615 self.attention = INTERNLM2_ATTENTION_CLASSES[config.attn_implementation](config=config)616 617 self.feed_forward = InternLM2MLP(config)618 self.attention_norm = InternLM2RMSNorm(config.hidden_size, eps=config.rms_norm_eps)619 self.ffn_norm = InternLM2RMSNorm(config.hidden_size, eps=config.rms_norm_eps)620 621 def forward(622 self,623 hidden_states: torch.Tensor,624 attention_mask: Optional[torch.Tensor] = None,625 position_ids: Optional[torch.LongTensor] = None,626 past_key_value: Optional[Tuple[torch.Tensor]] = None,627 output_attentions: Optional[bool] = False,628 use_cache: Optional[bool] = False,629 **kwargs,630 ) -> Tuple[torch.FloatTensor, Optional[Tuple[torch.FloatTensor, torch.FloatTensor]]]:631 """632 Args:633 hidden_states (`torch.FloatTensor`): input to the layer of shape `(batch, seq_len, embed_dim)`634 attention_mask (`torch.FloatTensor`, *optional*):635 attention mask of size `(batch_size, sequence_length)` if flash attention is used or `(batch_size, 1,636 query_sequence_length, key_sequence_length)` if default attention is used.637 output_attentions (`bool`, *optional*):638 Whether or not to return the attentions tensors of all attention layers. See `attentions` under639 returned tensors for more detail.640 use_cache (`bool`, *optional*):641 If set to `True`, `past_key_values` key value states are returned and can be used to speed up decoding642 (see `past_key_values`).643 past_key_value (`Tuple(torch.FloatTensor)`, *optional*): cached past key and value projection states644 """645 if 'padding_mask' in kwargs:646 warnings.warn(647 'Passing `padding_mask` is deprecated and will be removed in v4.37. '648 'Please make sure use `attention_mask` instead.`'649 )650 651 residual = hidden_states652 653 hidden_states = self.attention_norm(hidden_states)654 655 # Self Attention656 hidden_states, self_attn_weights, present_key_value = self.attention(657 hidden_states=hidden_states,658 attention_mask=attention_mask,659 position_ids=position_ids,660 past_key_value=past_key_value,661 output_attentions=output_attentions,662 use_cache=use_cache,663 **kwargs,664 )665 hidden_states = residual + hidden_states666 667 # Fully Connected668 residual = hidden_states669 hidden_states = self.ffn_norm(hidden_states)670 hidden_states = self.feed_forward(hidden_states)671 hidden_states = residual + hidden_states672 673 outputs = (hidden_states,)674 675 if output_attentions:676 outputs += (self_attn_weights,)677 678 if use_cache:679 outputs += (present_key_value,)680 681 return outputs682 683 684InternLM2_START_DOCSTRING = r"""685 This model inherits from [`PreTrainedModel`]. Check the superclass documentation for the generic methods the686 library implements for all its model (such as downloading or saving, resizing the input embeddings, pruning heads687 etc.)688 689 This model is also a PyTorch [torch.nn.Module](https://pytorch.org/docs/stable/nn.html#torch.nn.Module) subclass.690 Use it as a regular PyTorch Module and refer to the PyTorch documentation for all matter related to general usage691 and behavior.692 693 Parameters:694 config ([`InternLM2Config`]):695 Model configuration class with all the parameters of the model. Initializing with a config file does not696 load the weights associated with the model, only the configuration. Check out the697 [`~PreTrainedModel.from_pretrained`] method to load the model weights.698"""699 700 701# Copied from transformers.models.llama.modeling_llama.LlamaPreTrainedModel with Llama->InternLM2702@add_start_docstrings(703 'The bare InternLM2 Model outputting raw hidden-states without any specific head on top.',704 InternLM2_START_DOCSTRING,705)706class InternLM2PreTrainedModel(PreTrainedModel):707 config_class = InternLM2Config708 base_model_prefix = 'model'709 supports_gradient_checkpointing = True710 _no_split_modules = ['InternLM2DecoderLayer']711 _skip_keys_device_placement = 'past_key_values'712 _supports_flash_attn_2 = True713 714 def _init_weights(self, module):715 std = self.config.initializer_range716 if isinstance(module, nn.Linear):717 module.weight.data.normal_(mean=0.0, std=std)718 if module.bias is not None:719 module.bias.data.zero_()720 elif isinstance(module, nn.Embedding):721 module.weight.data.normal_(mean=0.0, std=std)722 if module.padding_idx is not None:723 module.weight.data[module.padding_idx].zero_()724 725 726InternLM2_INPUTS_DOCSTRING = r"""727 Args:728 input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`):729 Indices of input sequence tokens in the vocabulary. Padding will be ignored by default should you provide730 it.731 732 Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and733 [`PreTrainedTokenizer.__call__`] for details.734 735 [What are input IDs?](../glossary#input-ids)736 attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):737 Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:738 739 - 1 for tokens that are **not masked**,740 - 0 for tokens that are **masked**.741 742 [What are attention masks?](../glossary#attention-mask)743 744 Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and745 [`PreTrainedTokenizer.__call__`] for details.746 747 If `past_key_values` is used, optionally only the last `input_ids` have to be input (see748 `past_key_values`).749 750 If you want to change padding behavior, you should read [`modeling_opt._prepare_decoder_attention_mask`]751 and modify to your needs. See diagram 1 in [the paper](https://arxiv.org/abs/1910.13461) for more752 information on the default strategy.753 754 - 1 indicates the head is **not masked**,755 - 0 indicates the head is **masked**.756 position_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):757 Indices of positions of each input sequence tokens in the position embeddings. Selected in the range `[0,758 config.n_positions - 1]`.759 760 [What are position IDs?](../glossary#position-ids)761 past_key_values (`tuple(tuple(torch.FloatTensor))`, *optional*, returned when `use_cache=True` is passed or762 when `config.use_cache=True`):763 Tuple of `tuple(torch.FloatTensor)` of length `config.n_layers`, with each tuple having 2 tensors of shape764 `(batch_size, num_heads, sequence_length, embed_size_per_head)`) and 2 additional tensors of shape765 `(batch_size, num_heads, decoder_sequence_length, embed_size_per_head)`.766 767 Contains pre-computed hidden-states (key and values in the self-attention blocks and in the cross-attention768 blocks) that can be used (see `past_key_values` input) to speed up sequential decoding.769 770 If `past_key_values` are used, the user can optionally input only the last `input_ids` (those that don't771 have their past key value states given to this model) of shape `(batch_size, 1)` instead of all `input_ids`772 of shape `(batch_size, sequence_length)`.773 inputs_embeds (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`, *optional*):774 Optionally, instead of passing `input_ids` you can choose to directly pass an embedded representation. This775 is useful if you want more control over how to convert `input_ids` indices into associated vectors than the776 model's internal embedding lookup matrix.777 use_cache (`bool`, *optional*):778 If set to `True`, `past_key_values` key value states are returned and can be used to speed up decoding (see779 `past_key_values`).780 output_attentions (`bool`, *optional*):781 Whether or not to return the attentions tensors of all attention layers. See `attentions` under returned782 tensors for more detail.783 output_hidden_states (`bool`, *optional*):784 Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors for785 more detail.786 return_dict (`bool`, *optional*):787 Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.788"""789 790 791# Modified from transformers.model.llama.modeling_llama.LlamaModel792@add_start_docstrings(793 'The bare InternLM2 Model outputting raw hidden-states without any specific head on top.',794 InternLM2_START_DOCSTRING,795)796class InternLM2Model(InternLM2PreTrainedModel):797 """798 Transformer decoder consisting of *config.num_hidden_layers* layers. Each layer is a [`InternLM2DecoderLayer`]799 800 Args:801 config: InternLM2Config802 """803 804 _auto_class = 'AutoModel'805 806 def __init__(self, config: InternLM2Config):807 super().__init__(config)808 self.padding_idx = config.pad_token_id809 self.vocab_size = config.vocab_size810 self.config = config811 if not has_flash_attn:812 self.config.attn_implementation = 'eager'813 print('Warning: Flash attention is not available, using eager attention instead.')814 815 self.tok_embeddings = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx)816 817 self.layers = nn.ModuleList([InternLM2DecoderLayer(config) for _ in range(config.num_hidden_layers)])818 self.norm = InternLM2RMSNorm(config.hidden_size, eps=config.rms_norm_eps)819 820 self.gradient_checkpointing = False821 # Initialize weights and apply final processing822 self.post_init()823 824 def get_input_embeddings(self):825 return self.tok_embeddings826 827 def set_input_embeddings(self, value):828 self.tok_embeddings = value829 830 def _prepare_decoder_attention_mask(self, attention_mask, input_shape, inputs_embeds, past_key_values_length):831 # create causal mask832 # [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]833 combined_attention_mask = None834 if input_shape[-1] > 1:835 combined_attention_mask = _make_causal_mask(836 input_shape,837 inputs_embeds.dtype,838 device=inputs_embeds.device,839 past_key_values_length=past_key_values_length,840 )841 842 if attention_mask is not None:843 # [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]844 expanded_attn_mask = _expand_mask(attention_mask, inputs_embeds.dtype, tgt_len=input_shape[-1]).to(845 inputs_embeds.device846 )847 combined_attention_mask = (848 expanded_attn_mask if combined_attention_mask is None else expanded_attn_mask + combined_attention_mask849 )850 851 return combined_attention_mask852 853 @add_start_docstrings_to_model_forward(InternLM2_INPUTS_DOCSTRING)854 def forward(855 self,856 input_ids: torch.LongTensor = None,857 attention_mask: Optional[torch.Tensor] = None,858 position_ids: Optional[torch.LongTensor] = None,859 past_key_values: Optional[List[torch.FloatTensor]] = None,860 inputs_embeds: Optional[torch.FloatTensor] = None,861 use_cache: Optional[bool] = None,862 output_attentions: Optional[bool] = None,863 output_hidden_states: Optional[bool] = None,864 return_dict: Optional[bool] = None,865 ) -> Union[Tuple, BaseModelOutputWithPast]:866 output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions867 output_hidden_states = (868 output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states869 )870 use_cache = use_cache if use_cache is not None else self.config.use_cache871 872 return_dict = return_dict if return_dict is not None else self.config.use_return_dict873 874 if self.config.attn_implementation == 'flash_attention_2':875 _import_flash_attn()876 877 # retrieve input_ids and inputs_embeds878 if input_ids is not None and inputs_embeds is not None:879 raise ValueError('You cannot specify both input_ids and inputs_embeds at the same time')880 elif input_ids is not None:881 batch_size, seq_length = input_ids.shape[:2]882 elif inputs_embeds is not None:883 batch_size, seq_length = inputs_embeds.shape[:2]884 else:885 raise ValueError('You have to specify either input_ids or inputs_embeds')886 887 seq_length_with_past = seq_length888 past_key_values_length = 0889 if past_key_values is not None:890 past_key_values_length = past_key_values[0][0].shape[2]891 seq_length_with_past = seq_length_with_past + past_key_values_length892 893 if position_ids is None:894 device = input_ids.device if input_ids is not None else inputs_embeds.device895 position_ids = torch.arange(896 past_key_values_length, seq_length + past_key_values_length, dtype=torch.long, device=device897 )898 position_ids = position_ids.unsqueeze(0)899 900 if inputs_embeds is None:901 inputs_embeds = self.tok_embeddings(input_ids)902 903 if self.config.attn_implementation == 'flash_attention_2':904 # 2d mask is passed through the layers905 attention_mask = attention_mask if (attention_mask is not None and 0 in attention_mask) else None906 else:907 if attention_mask is None:908 attention_mask = torch.ones(909 (batch_size, seq_length_with_past), dtype=torch.bool, device=inputs_embeds.device910 )911 attention_mask = self._prepare_decoder_attention_mask(912 attention_mask, (batch_size, seq_length), inputs_embeds, past_key_values_length913 )914 915 # embed positions916 hidden_states = inputs_embeds917 918 if self.gradient_checkpointing and self.training:919 if use_cache:920 logger.warning_once(921 '`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`...'922 )923 use_cache = False924 925 # decoder layers926 all_hidden_states = () if output_hidden_states else None927 all_self_attns = () if output_attentions else None928 next_decoder_cache = () if use_cache else None929 930 for idx, decoder_layer in enumerate(self.layers):931 if output_hidden_states:932 all_hidden_states += (hidden_states,)933 934 past_key_value = past_key_values[idx] if past_key_values is not None else None935 936 if self.gradient_checkpointing and self.training:937 938 def create_custom_forward(module):939 def custom_forward(*inputs):940 # None for past_key_value941 return module(*inputs, output_attentions, None)942 943 return custom_forward944 945 layer_outputs = torch.utils.checkpoint.checkpoint(946 create_custom_forward(decoder_layer),947 hidden_states,948 attention_mask,949 position_ids,950 None,951 )952 else:953 layer_outputs = decoder_layer(954 hidden_states,955 attention_mask=attention_mask,956 position_ids=position_ids,957 past_key_value=past_key_value,958 output_attentions=output_attentions,959 use_cache=use_cache,960 )961 962 hidden_states = layer_outputs[0]963 964 if use_cache:965 next_decoder_cache += (layer_outputs[2 if output_attentions else 1],)966 967 if output_attentions:968 all_self_attns += (layer_outputs[1],)969 970 hidden_states = self.norm(hidden_states)971 972 # add hidden states from the last decoder layer973 if output_hidden_states:974 all_hidden_states += (hidden_states,)975 976 next_cache = next_decoder_cache if use_cache else None977 if not return_dict:978 return tuple(v for v in [hidden_states, next_cache, all_hidden_states, all_self_attns] if v is not None)979 return BaseModelOutputWithPast(980 last_hidden_state=hidden_states,981 past_key_values=next_cache,982 hidden_states=all_hidden_states,983 attentions=all_self_attns,984 )985 986 987# Modified from transformers.model.llama.modeling_llama.LlamaForCausalLM988class InternLM2ForCausalLM(InternLM2PreTrainedModel):989 _auto_class = 'AutoModelForCausalLM'990 991 _tied_weights_keys = ['output.weight']992 993 def __init__(self, config):994 super().__init__(config)995 self.model = InternLM2Model(config)996 self.vocab_size = config.vocab_size997 self.output = nn.Linear(config.hidden_size, config.vocab_size, bias=False)998 999 # Initialize weights and apply final processing1000 self.post_init()1001 1002 def get_input_embeddings(self):1003 return self.model.tok_embeddings1004 1005 def set_input_embeddings(self, value):1006 self.model.tok_embeddings = value1007 1008 def get_output_embeddings(self):1009 return self.output1010 1011 def set_output_embeddings(self, new_embeddings):1012 self.output = new_embeddings1013 1014 def set_decoder(self, decoder):1015 self.model = decoder1016 1017 def get_decoder(self):1018 return self.model1019 1020 @add_start_docstrings_to_model_forward(InternLM2_INPUTS_DOCSTRING)1021 @replace_return_docstrings(output_type=CausalLMOutputWithPast, config_class=_CONFIG_FOR_DOC)1022 def forward(1023 self,1024 input_ids: torch.LongTensor = None,1025 attention_mask: Optional[torch.Tensor] = None,1026 position_ids: Optional[torch.LongTensor] = None,1027 past_key_values: Optional[List[torch.FloatTensor]] = None,1028 inputs_embeds: Optional[torch.FloatTensor] = None,1029 labels: Optional[torch.LongTensor] = None,1030 use_cache: Optional[bool] = None,1031 output_attentions: Optional[bool] = None,1032 output_hidden_states: Optional[bool] = None,1033 return_dict: Optional[bool] = None,1034 ) -> Union[Tuple, CausalLMOutputWithPast]:1035 r"""1036 Args:1037 labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):1038 Labels for computing the masked language modeling loss. Indices should either be in `[0, ...,1039 config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored1040 (masked), the loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`.1041 1042 Returns:1043 1044 Example:1045 1046 ```python1047 >>> from transformers import AutoTokenizer, InternLM2ForCausalLM1048 1049 >>> model = InternLM2ForCausalLM.from_pretrained(PATH_TO_CONVERTED_WEIGHTS)1050 >>> tokenizer = AutoTokenizer.from_pretrained(PATH_TO_CONVERTED_TOKENIZER)1051 1052 >>> prompt = "Hey, are you conscious? Can you talk to me?"1053 >>> inputs = tokenizer(prompt, return_tensors="pt")1054 1055 >>> # Generate1056 >>> generate_ids = model.generate(inputs.input_ids, max_length=30)1057 >>> tokenizer.batch_decode(generate_ids, skip_special_tokens=True, clean_up_tokenization_spaces=False)[0]1058 "Hey, are you conscious? Can you talk to me?\nI'm not conscious, but I can talk to you."1059 ```"""1060 1061 output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions1062 output_hidden_states = (1063 output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states1064 )1065 return_dict = return_dict if return_dict is not None else self.config.use_return_dict1066 1067 # decoder outputs consists of (dec_features, layer_state, dec_hidden, dec_attn)1068 outputs = self.model(1069 input_ids=input_ids,1070 attention_mask=attention_mask,1071 position_ids=position_ids,1072 past_key_values=past_key_values,1073 inputs_embeds=inputs_embeds,1074 use_cache=use_cache,1075 output_attentions=output_attentions,1076 output_hidden_states=output_hidden_states,1077 return_dict=return_dict,1078 )1079 1080 hidden_states = outputs[0]1081 logits = self.output(hidden_states)1082 logits = logits.float()1083 1084 loss = None1085 if labels is not None:1086 # Shift so that tokens < n predict n1087 shift_logits = logits[..., :-1, :].contiguous()1088 shift_labels = labels[..., 1:].contiguous()1089 # Flatten the tokens1090 loss_fct = CrossEntropyLoss()1091 shift_logits = shift_logits.view(-1, self.config.vocab_size)1092 shift_labels = shift_labels.view(-1)1093 # Enable model parallelism1094 shift_labels = shift_labels.to(shift_logits.device)1095 loss = loss_fct(shift_logits, shift_labels)1096 1097 if not return_dict:1098 output = (logits,) + outputs[1:]1099 return (loss,) + output if loss is not None else output1100 1101 device = input_ids.device if input_ids is not None else inputs_embeds.device1102 output = CausalLMOutputWithPast(1103 loss=loss,1104 logits=logits,1105 past_key_values=outputs.past_key_values,1106 hidden_states=outputs.hidden_states,1107 attentions=outputs.attentions,1108 )1109 output['logits'] = output['logits'].to(device)1110 return output1111 1112 def prepare_inputs_for_generation(1113 self, input_ids, past_key_values=None, attention_mask=None, inputs_embeds=None, **kwargs1114 ):1115 if past_key_values is not None:1116 past_length = past_key_values[0][0].shape[2]1117 1118 # Some generation methods already pass only the last input ID1119 if input_ids.shape[1] > past_length:1120 remove_prefix_length = past_length1121 else:1122 # Default to old behavior: keep only final ID1123 remove_prefix_length = input_ids.shape[1] - 11124 1125 input_ids = input_ids[:, remove_prefix_length:]1126 1127 position_ids = kwargs.get('position_ids', None)1128 if attention_mask is not None and position_ids is None:1129 # create position_ids on the fly for batch generation1130 position_ids = attention_mask.long().cumsum(-1) - 11131 position_ids.masked_fill_(attention_mask == 0, 1)1132 if past_key_values:1133 position_ids = position_ids[:, -input_ids.shape[1] :]1134 1135 # if `inputs_embeds` are passed, we only want to use them in the 1st generation step1136 if inputs_embeds is not None and past_key_values is None:1137 model_inputs = {'inputs_embeds': inputs_embeds}1138 else:1139 model_inputs = {'input_ids': input_ids}1140 1141 model_inputs.update(1142 {1143 'position_ids': position_ids,1144 'past_key_values': past_key_values,1145 'use_cache': kwargs.get('use_cache'),1146 'attention_mask': attention_mask,1147 }1148 )1149 return model_inputs1150 1151 @staticmethod1152 def _reorder_cache(past_key_values, beam_idx):1153 reordered_past = ()1154 for layer_past in past_key_values:1155 reordered_past += (1156 tuple(past_state.index_select(0, beam_idx.to(past_state.device)) for past_state in layer_past),1157 )1158 return reordered_past1159 1160 def build_inputs(self, tokenizer, query: str, history: List[Tuple[str, str]] = [], meta_instruction=''):1161 if tokenizer.add_bos_token:1162 prompt = ''1163 else:1164 prompt = tokenizer.bos_token1165 if meta_instruction:1166 prompt += f"""<|im_start|>system\n{meta_instruction}<|im_end|>\n"""1167 for record in history:1168 prompt += f"""<|im_start|>user\n{record[0]}<|im_end|>\n<|im_start|>assistant\n{record[1]}<|im_end|>\n"""1169 prompt += f"""<|im_start|>user\n{query}<|im_end|>\n<|im_start|>assistant\n"""1170 return tokenizer([prompt], return_tensors='pt')1171 1172 @torch.no_grad()1173 def chat(1174 self,1175 tokenizer,1176 query: str,1177 history: List[Tuple[str, str]] = [],1178 streamer: Optional[BaseStreamer] = None,1179 max_new_tokens: int = 1024,1180 do_sample: bool = True,1181 temperature: float = 0.8,1182 top_p: float = 0.8,1183 meta_instruction: str = 'You are an AI assistant whose name is InternLM (书生·浦语).\n'1184 '- InternLM (书生·浦语) is a conversational language model that is developed by Shanghai AI Laboratory (上海人工智能实验室). It is designed to be helpful, honest, and harmless.\n'1185 '- InternLM (书生·浦语) can understand and communicate fluently in the language chosen by the user such as English and 中文.',1186 **kwargs,1187 ):1188 inputs = self.build_inputs(tokenizer, query, history, meta_instruction)1189 inputs = {k: v.to(self.device) for k, v in inputs.items() if torch.is_tensor(v)}1190 # also add end-of-assistant token in eos token id to avoid unnecessary generation1191 eos_token_id = [tokenizer.eos_token_id, tokenizer.convert_tokens_to_ids(['<|im_end|>'])[0]]1192 outputs = self.generate(1193 **inputs,1194 streamer=streamer,1195 max_new_tokens=max_new_tokens,1196 do_sample=do_sample,1197 temperature=temperature,1198 top_p=top_p,1199 eos_token_id=eos_token_id,1200 **kwargs,