MathLLMs/MathCoder-VL-2B
730
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 713 def _init_weights(self, module):714 std = self.config.initializer_range715 if isinstance(module, nn.Linear):716 module.weight.data.normal_(mean=0.0, std=std)717 if module.bias is not None:718 module.bias.data.zero_()719 elif isinstance(module, nn.Embedding):720 module.weight.data.normal_(mean=0.0, std=std)721 if module.padding_idx is not None:722 module.weight.data[module.padding_idx].zero_()723 724 725InternLM2_INPUTS_DOCSTRING = r"""726 Args:727 input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`):728 Indices of input sequence tokens in the vocabulary. Padding will be ignored by default should you provide729 it.730 731 Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and732 [`PreTrainedTokenizer.__call__`] for details.733 734 [What are input IDs?](../glossary#input-ids)735 attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):736 Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:737 738 - 1 for tokens that are **not masked**,739 - 0 for tokens that are **masked**.740 741 [What are attention masks?](../glossary#attention-mask)742 743 Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and744 [`PreTrainedTokenizer.__call__`] for details.745 746 If `past_key_values` is used, optionally only the last `input_ids` have to be input (see747 `past_key_values`).748 749 If you want to change padding behavior, you should read [`modeling_opt._prepare_decoder_attention_mask`]750 and modify to your needs. See diagram 1 in [the paper](https://arxiv.org/abs/1910.13461) for more751 information on the default strategy.752 753 - 1 indicates the head is **not masked**,754 - 0 indicates the head is **masked**.755 position_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):756 Indices of positions of each input sequence tokens in the position embeddings. Selected in the range `[0,757 config.n_positions - 1]`.758 759 [What are position IDs?](../glossary#position-ids)760 past_key_values (`tuple(tuple(torch.FloatTensor))`, *optional*, returned when `use_cache=True` is passed or761 when `config.use_cache=True`):762 Tuple of `tuple(torch.FloatTensor)` of length `config.n_layers`, with each tuple having 2 tensors of shape763 `(batch_size, num_heads, sequence_length, embed_size_per_head)`) and 2 additional tensors of shape764 `(batch_size, num_heads, decoder_sequence_length, embed_size_per_head)`.765 766 Contains pre-computed hidden-states (key and values in the self-attention blocks and in the cross-attention767 blocks) that can be used (see `past_key_values` input) to speed up sequential decoding.768 769 If `past_key_values` are used, the user can optionally input only the last `input_ids` (those that don't770 have their past key value states given to this model) of shape `(batch_size, 1)` instead of all `input_ids`771 of shape `(batch_size, sequence_length)`.772 inputs_embeds (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`, *optional*):773 Optionally, instead of passing `input_ids` you can choose to directly pass an embedded representation. This774 is useful if you want more control over how to convert `input_ids` indices into associated vectors than the775 model's internal embedding lookup matrix.776 use_cache (`bool`, *optional*):777 If set to `True`, `past_key_values` key value states are returned and can be used to speed up decoding (see778 `past_key_values`).779 output_attentions (`bool`, *optional*):780 Whether or not to return the attentions tensors of all attention layers. See `attentions` under returned781 tensors for more detail.782 output_hidden_states (`bool`, *optional*):783 Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors for784 more detail.785 return_dict (`bool`, *optional*):786 Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.787"""788 789 790# Modified from transformers.model.llama.modeling_llama.LlamaModel791@add_start_docstrings(792 'The bare InternLM2 Model outputting raw hidden-states without any specific head on top.',793 InternLM2_START_DOCSTRING,794)795class InternLM2Model(InternLM2PreTrainedModel):796 """797 Transformer decoder consisting of *config.num_hidden_layers* layers. Each layer is a [`InternLM2DecoderLayer`]798 799 Args:800 config: InternLM2Config801 """802 803 _auto_class = 'AutoModel'804 805 def __init__(self, config: InternLM2Config):806 super().__init__(config)807 self.padding_idx = config.pad_token_id808 self.vocab_size = config.vocab_size809 self.config = config810 if not has_flash_attn:811 self.config.attn_implementation = 'eager'812 print('Warning: Flash attention is not available, using eager attention instead.')813 814 self.tok_embeddings = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx)815 816 self.layers = nn.ModuleList([InternLM2DecoderLayer(config) for _ in range(config.num_hidden_layers)])817 self.norm = InternLM2RMSNorm(config.hidden_size, eps=config.rms_norm_eps)818 819 self.gradient_checkpointing = False820 # Initialize weights and apply final processing821 self.post_init()822 823 def get_input_embeddings(self):824 return self.tok_embeddings825 826 def set_input_embeddings(self, value):827 self.tok_embeddings = value828 829 def _prepare_decoder_attention_mask(self, attention_mask, input_shape, inputs_embeds, past_key_values_length):830 # create causal mask831 # [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]832 combined_attention_mask = None833 if input_shape[-1] > 1:834 combined_attention_mask = _make_causal_mask(835 input_shape,836 inputs_embeds.dtype,837 device=inputs_embeds.device,838 past_key_values_length=past_key_values_length,839 )840 841 if attention_mask is not None:842 # [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]843 expanded_attn_mask = _expand_mask(attention_mask, inputs_embeds.dtype, tgt_len=input_shape[-1]).to(844 inputs_embeds.device845 )846 combined_attention_mask = (847 expanded_attn_mask if combined_attention_mask is None else expanded_attn_mask + combined_attention_mask848 )849 850 return combined_attention_mask851 852 @add_start_docstrings_to_model_forward(InternLM2_INPUTS_DOCSTRING)853 def forward(854 self,855 input_ids: torch.LongTensor = None,856 attention_mask: Optional[torch.Tensor] = None,857 position_ids: Optional[torch.LongTensor] = None,858 past_key_values: Optional[List[torch.FloatTensor]] = None,859 inputs_embeds: Optional[torch.FloatTensor] = None,860 use_cache: Optional[bool] = None,861 output_attentions: Optional[bool] = None,862 output_hidden_states: Optional[bool] = None,863 return_dict: Optional[bool] = None,864 ) -> Union[Tuple, BaseModelOutputWithPast]:865 output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions866 output_hidden_states = (867 output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states868 )869 use_cache = use_cache if use_cache is not None else self.config.use_cache870 871 return_dict = return_dict if return_dict is not None else self.config.use_return_dict872 873 if self.config.attn_implementation == 'flash_attention_2':874 _import_flash_attn()875 876 # retrieve input_ids and inputs_embeds877 if input_ids is not None and inputs_embeds is not None:878 raise ValueError('You cannot specify both input_ids and inputs_embeds at the same time')879 elif input_ids is not None:880 batch_size, seq_length = input_ids.shape[:2]881 elif inputs_embeds is not None:882 batch_size, seq_length = inputs_embeds.shape[:2]883 else:884 raise ValueError('You have to specify either input_ids or inputs_embeds')885 886 seq_length_with_past = seq_length887 past_key_values_length = 0888 if past_key_values is not None:889 past_key_values_length = past_key_values[0][0].shape[2]890 seq_length_with_past = seq_length_with_past + past_key_values_length891 892 if position_ids is None:893 device = input_ids.device if input_ids is not None else inputs_embeds.device894 position_ids = torch.arange(895 past_key_values_length, seq_length + past_key_values_length, dtype=torch.long, device=device896 )897 position_ids = position_ids.unsqueeze(0)898 899 if inputs_embeds is None:900 inputs_embeds = self.tok_embeddings(input_ids)901 902 if self.config.attn_implementation == 'flash_attention_2':903 # 2d mask is passed through the layers904 attention_mask = attention_mask if (attention_mask is not None and 0 in attention_mask) else None905 else:906 if attention_mask is None:907 attention_mask = torch.ones(908 (batch_size, seq_length_with_past), dtype=torch.bool, device=inputs_embeds.device909 )910 attention_mask = self._prepare_decoder_attention_mask(911 attention_mask, (batch_size, seq_length), inputs_embeds, past_key_values_length912 )913 914 # embed positions915 hidden_states = inputs_embeds916 917 if self.gradient_checkpointing and self.training:918 if use_cache:919 logger.warning_once(920 '`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`...'921 )922 use_cache = False923 924 # decoder layers925 all_hidden_states = () if output_hidden_states else None926 all_self_attns = () if output_attentions else None927 next_decoder_cache = () if use_cache else None928 929 for idx, decoder_layer in enumerate(self.layers):930 if output_hidden_states:931 all_hidden_states += (hidden_states,)932 933 past_key_value = past_key_values[idx] if past_key_values is not None else None934 935 if self.gradient_checkpointing and self.training:936 937 def create_custom_forward(module):938 def custom_forward(*inputs):939 # None for past_key_value940 return module(*inputs, output_attentions, None)941 942 return custom_forward943 944 layer_outputs = torch.utils.checkpoint.checkpoint(945 create_custom_forward(decoder_layer),946 hidden_states,947 attention_mask,948 position_ids,949 None,950 )951 else:952 layer_outputs = decoder_layer(953 hidden_states,954 attention_mask=attention_mask,955 position_ids=position_ids,956 past_key_value=past_key_value,957 output_attentions=output_attentions,958 use_cache=use_cache,959 )960 961 hidden_states = layer_outputs[0]962 963 if use_cache:964 next_decoder_cache += (layer_outputs[2 if output_attentions else 1],)965 966 if output_attentions:967 all_self_attns += (layer_outputs[1],)968 969 hidden_states = self.norm(hidden_states)970 971 # add hidden states from the last decoder layer972 if output_hidden_states:973 all_hidden_states += (hidden_states,)974 975 next_cache = next_decoder_cache if use_cache else None976 if not return_dict:977 return tuple(v for v in [hidden_states, next_cache, all_hidden_states, all_self_attns] if v is not None)978 return BaseModelOutputWithPast(979 last_hidden_state=hidden_states,980 past_key_values=next_cache,981 hidden_states=all_hidden_states,982 attentions=all_self_attns,983 )984 985 986# Modified from transformers.model.llama.modeling_llama.LlamaForCausalLM987class InternLM2ForCausalLM(InternLM2PreTrainedModel):988 _auto_class = 'AutoModelForCausalLM'989 990 _tied_weights_keys = ['output.weight']991 992 def __init__(self, config):993 super().__init__(config)994 self.model = InternLM2Model(config)995 self.vocab_size = config.vocab_size996 self.output = nn.Linear(config.hidden_size, config.vocab_size, bias=False)997 998 # Initialize weights and apply final processing999 self.post_init()1000 1001 def get_input_embeddings(self):1002 return self.model.tok_embeddings1003 1004 def set_input_embeddings(self, value):1005 self.model.tok_embeddings = value1006 1007 def get_output_embeddings(self):1008 return self.output1009 1010 def set_output_embeddings(self, new_embeddings):1011 self.output = new_embeddings1012 1013 def set_decoder(self, decoder):1014 self.model = decoder1015 1016 def get_decoder(self):1017 return self.model1018 1019 @add_start_docstrings_to_model_forward(InternLM2_INPUTS_DOCSTRING)1020 @replace_return_docstrings(output_type=CausalLMOutputWithPast, config_class=_CONFIG_FOR_DOC)1021 def forward(1022 self,1023 input_ids: torch.LongTensor = None,1024 attention_mask: Optional[torch.Tensor] = None,1025 position_ids: Optional[torch.LongTensor] = None,1026 past_key_values: Optional[List[torch.FloatTensor]] = None,1027 inputs_embeds: Optional[torch.FloatTensor] = None,1028 labels: Optional[torch.LongTensor] = None,1029 use_cache: Optional[bool] = None,1030 output_attentions: Optional[bool] = None,1031 output_hidden_states: Optional[bool] = None,1032 return_dict: Optional[bool] = None,1033 ) -> Union[Tuple, CausalLMOutputWithPast]:1034 r"""1035 Args:1036 labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):1037 Labels for computing the masked language modeling loss. Indices should either be in `[0, ...,1038 config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored1039 (masked), the loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`.1040 1041 Returns:1042 1043 Example:1044 1045 ```python1046 >>> from transformers import AutoTokenizer, InternLM2ForCausalLM1047 1048 >>> model = InternLM2ForCausalLM.from_pretrained(PATH_TO_CONVERTED_WEIGHTS)1049 >>> tokenizer = AutoTokenizer.from_pretrained(PATH_TO_CONVERTED_TOKENIZER)1050 1051 >>> prompt = "Hey, are you conscious? Can you talk to me?"1052 >>> inputs = tokenizer(prompt, return_tensors="pt")1053 1054 >>> # Generate1055 >>> generate_ids = model.generate(inputs.input_ids, max_length=30)1056 >>> tokenizer.batch_decode(generate_ids, skip_special_tokens=True, clean_up_tokenization_spaces=False)[0]1057 "Hey, are you conscious? Can you talk to me?\nI'm not conscious, but I can talk to you."1058 ```"""1059 1060 output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions1061 output_hidden_states = (1062 output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states1063 )1064 return_dict = return_dict if return_dict is not None else self.config.use_return_dict1065 1066 # decoder outputs consists of (dec_features, layer_state, dec_hidden, dec_attn)1067 outputs = self.model(1068 input_ids=input_ids,1069 attention_mask=attention_mask,1070 position_ids=position_ids,1071 past_key_values=past_key_values,1072 inputs_embeds=inputs_embeds,1073 use_cache=use_cache,1074 output_attentions=output_attentions,1075 output_hidden_states=output_hidden_states,1076 return_dict=return_dict,1077 )1078 1079 hidden_states = outputs[0]1080 logits = self.output(hidden_states)1081 logits = logits.float()1082 1083 loss = None1084 if labels is not None:1085 # Shift so that tokens < n predict n1086 shift_logits = logits[..., :-1, :].contiguous()1087 shift_labels = labels[..., 1:].contiguous()1088 # Flatten the tokens1089 loss_fct = CrossEntropyLoss()1090 shift_logits = shift_logits.view(-1, self.config.vocab_size)1091 shift_labels = shift_labels.view(-1)1092 # Enable model parallelism1093 shift_labels = shift_labels.to(shift_logits.device)1094 loss = loss_fct(shift_logits, shift_labels)1095 1096 if not return_dict:1097 output = (logits,) + outputs[1:]1098 return (loss,) + output if loss is not None else output1099 1100 device = input_ids.device if input_ids is not None else inputs_embeds.device1101 output = CausalLMOutputWithPast(1102 loss=loss,1103 logits=logits,1104 past_key_values=outputs.past_key_values,1105 hidden_states=outputs.hidden_states,1106 attentions=outputs.attentions,1107 )1108 output['logits'] = output['logits'].to(device)1109 return output1110 1111 def prepare_inputs_for_generation(1112 self, input_ids, past_key_values=None, attention_mask=None, inputs_embeds=None, **kwargs1113 ):1114 if past_key_values is not None:1115 past_length = past_key_values[0][0].shape[2]1116 1117 # Some generation methods already pass only the last input ID1118 if input_ids.shape[1] > past_length:1119 remove_prefix_length = past_length1120 else:1121 # Default to old behavior: keep only final ID1122 remove_prefix_length = input_ids.shape[1] - 11123 1124 input_ids = input_ids[:, remove_prefix_length:]1125 1126 position_ids = kwargs.get('position_ids', None)1127 if attention_mask is not None and position_ids is None:1128 # create position_ids on the fly for batch generation1129 position_ids = attention_mask.long().cumsum(-1) - 11130 position_ids.masked_fill_(attention_mask == 0, 1)1131 if past_key_values:1132 position_ids = position_ids[:, -input_ids.shape[1] :]1133 1134 # if `inputs_embeds` are passed, we only want to use them in the 1st generation step1135 if inputs_embeds is not None and past_key_values is None:1136 model_inputs = {'inputs_embeds': inputs_embeds}1137 else:1138 model_inputs = {'input_ids': input_ids}1139 1140 model_inputs.update(1141 {1142 'position_ids': position_ids,1143 'past_key_values': past_key_values,1144 'use_cache': kwargs.get('use_cache'),1145 'attention_mask': attention_mask,1146 }1147 )1148 return model_inputs1149 1150 @staticmethod1151 def _reorder_cache(past_key_values, beam_idx):1152 reordered_past = ()1153 for layer_past in past_key_values:1154 reordered_past += (1155 tuple(past_state.index_select(0, beam_idx.to(past_state.device)) for past_state in layer_past),1156 )1157 return reordered_past1158 1159 def build_inputs(self, tokenizer, query: str, history: List[Tuple[str, str]] = [], meta_instruction=''):1160 if tokenizer.add_bos_token:1161 prompt = ''1162 else:1163 prompt = tokenizer.bos_token1164 if meta_instruction:1165 prompt += f"""<|im_start|>system\n{meta_instruction}<|im_end|>\n"""1166 for record in history:1167 prompt += f"""<|im_start|>user\n{record[0]}<|im_end|>\n<|im_start|>assistant\n{record[1]}<|im_end|>\n"""1168 prompt += f"""<|im_start|>user\n{query}<|im_end|>\n<|im_start|>assistant\n"""1169 return tokenizer([prompt], return_tensors='pt')1170 1171 @torch.no_grad()1172 def chat(1173 self,1174 tokenizer,1175 query: str,1176 history: List[Tuple[str, str]] = [],1177 streamer: Optional[BaseStreamer] = None,1178 max_new_tokens: int = 1024,1179 do_sample: bool = True,1180 temperature: float = 0.8,1181 top_p: float = 0.8,1182 meta_instruction: str = 'You are an AI assistant whose name is InternLM (书生·浦语).\n'1183 '- InternLM (书生·浦语) is a conversational language model that is developed by Shanghai AI Laboratory (上海人工智能实验室). It is designed to be helpful, honest, and harmless.\n'1184 '- InternLM (书生·浦语) can understand and communicate fluently in the language chosen by the user such as English and 中文.',1185 **kwargs,1186 ):1187 inputs = self.build_inputs(tokenizer, query, history, meta_instruction)1188 inputs = {k: v.to(self.device) for k, v in inputs.items() if torch.is_tensor(v)}1189 # also add end-of-assistant token in eos token id to avoid unnecessary generation1190 eos_token_id = [tokenizer.eos_token_id, tokenizer.convert_tokens_to_ids(['<|im_end|>'])[0]]1191 outputs = self.generate(1192 **inputs,1193 streamer=streamer,1194 max_new_tokens=max_new_tokens,1195 do_sample=do_sample,1196 temperature=temperature,1197 top_p=top_p,1198 eos_token_id=eos_token_id,1199 **kwargs,1200 )