CoolFace
Modelpublic

OpenGVLab/VisualPRM-8B

sourceHugging Facemitupdated 1y agoView on Hugging Face
17likes109downloads
modeling_internlm2.py1416 linesDownload Raw Back to root
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,

Showing the first 1,200 of 1416 lines. Download the file for the rest.