Team Ai
Modelpublic

WisdomShell/Shell-7B-Chat

sourceHugging Faceupdated 2y agoView on Hugging Face
1likes18downloads
modeling_codeshell.py1088 linesDownload Raw Back to root
1# coding=utf-82# Copyright 2023 WisdomShell Inc. All Rights Reserved.3 4# Licensed under the Apache License, Version 2.0 (the "License");5# you may not use this file except in compliance with the License.6# You may obtain a copy of the License at7#8#     http://www.apache.org/licenses/LICENSE-2.09#10# Unless required by applicable law or agreed to in writing, software11# distributed under the License is distributed on an "AS IS" BASIS,12# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.13# See the License for the specific language governing permissions and14# limitations under the License.15 16# This code is based on Bigcode's GPTBigCode model. It has been modified from17# its original forms to accommodate minor architectural differences compared to 18# GPTBigCode model that trained the model.19 20# Copyright 2023 The Bigcode team and HuggingFace Inc. team.21# Licensed under the Apache License, Version 2.0 (the "License");22# you may not use this file except in compliance with the License.23# You may obtain a copy of the License at24#25#     http://www.apache.org/licenses/LICENSE-2.026#27# Unless required by applicable law or agreed to in writing, software28# distributed under the License is distributed on an "AS IS" BASIS,29# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.30# See the License for the specific language governing permissions and31# limitations under the License.32"""PyTorch CodeShell model."""33import os34import math35from typing import List, Optional, Tuple, Union, Callable36from threading import Thread37from queue import Queue38 39 40import torch41import torch.utils.checkpoint42from torch import nn43from torch.nn import BCEWithLogitsLoss, CrossEntropyLoss, MSELoss44 45from transformers import LogitsProcessorList, StoppingCriteriaList, StoppingCriteria, PreTrainedModel, PretrainedConfig46from transformers.generation.utils import GenerationConfig47 48from transformers.activations import ACT2FN49from transformers.modeling_outputs import (50    BaseModelOutputWithPastAndCrossAttentions,51    CausalLMOutputWithCrossAttentions,52)53from transformers.modeling_utils import PreTrainedModel54from transformers.utils import (55    add_start_docstrings,56    add_start_docstrings_to_model_forward,57)58from .configuration_codeshell import CodeShellConfig59 60# Fused kernels61# Use separate functions for each case because conditionals prevent kernel fusion.62# TODO: Could have better fused kernels depending on scaling, dropout and head mask.63#  Is it doable without writing 32 functions?64@torch.jit.script65def upcast_masked_softmax(66    x: torch.Tensor, mask: torch.Tensor, mask_value: torch.Tensor, scale: float, softmax_dtype: torch.dtype67):68    input_dtype = x.dtype69    x = x.to(softmax_dtype) * scale70    x = torch.where(mask, x, mask_value)71    x = torch.nn.functional.softmax(x, dim=-1).to(input_dtype)72    return x73 74 75@torch.jit.script76def upcast_softmax(x: torch.Tensor, scale: float, softmax_dtype: torch.dtype):77    78    input_dtype = x.dtype79    x = x.to(softmax_dtype) * scale80    x = torch.nn.functional.softmax(x, dim=-1).to(input_dtype)81    return x82 83 84@torch.jit.script85def masked_softmax(x: torch.Tensor, mask: torch.Tensor, mask_value: torch.Tensor):86    x = torch.where(mask, x, mask_value)87    x = torch.nn.functional.softmax(x, dim=-1)88    return x89 90 91class CodeShellRotaryEmbedding(torch.nn.Module):92    def __init__(self, dim, max_position_embeddings=2048, base=10000, device=None):93        super().__init__()94 95        self.dim = dim96        self.max_position_embeddings = max_position_embeddings97        self.base = base98        inv_freq = 1.0 / (self.base ** (torch.arange(0, self.dim, 2).float().to(device) / self.dim))99        self.register_buffer("inv_freq", inv_freq)100 101        # Build here to make `torch.jit.trace` work.102        self._set_cos_sin_cache(103            seq_len=max_position_embeddings, device=self.inv_freq.device, dtype=torch.get_default_dtype()104        )105 106    def _set_cos_sin_cache(self, seq_len, device, dtype):107        self.max_seq_len_cached = seq_len108        t = torch.arange(self.max_seq_len_cached, device=device, dtype=self.inv_freq.dtype)109 110        freqs = torch.einsum("i,j->ij", t, self.inv_freq)111        # Different from paper, but it uses a different permutation in order to obtain the same calculation112        emb = torch.cat((freqs, freqs), dim=-1)113        self.register_buffer("cos_cached", emb.cos()[None, None, :, :].to(dtype), persistent=False)114        self.register_buffer("sin_cached", emb.sin()[None, None, :, :].to(dtype), persistent=False)115 116    def forward(self, x, seq_len=None):117        # x: [bs, num_attention_heads, seq_len, head_size]118        if seq_len > self.max_seq_len_cached:119            self._set_cos_sin_cache(seq_len=seq_len, device=x.device, dtype=x.dtype)120 121        return (122            self.cos_cached[:, :, :seq_len, ...].to(dtype=x.dtype),123            self.sin_cached[:, :, :seq_len, ...].to(dtype=x.dtype),124        )125 126 127class CodeShellLinearScalingRotaryEmbedding(CodeShellRotaryEmbedding):128    """CodeShellRotaryEmbedding extended with linear scaling. Credits to the Reddit user /u/kaiokendev"""129 130    def __init__(self, dim, max_position_embeddings=2048, base=10000, device=None, scaling_factor=1.0):131        self.scaling_factor = scaling_factor132        super().__init__(dim, max_position_embeddings, base, device)133 134    def _set_cos_sin_cache(self, seq_len, device, dtype):135        self.max_seq_len_cached = seq_len136        t = torch.arange(self.max_seq_len_cached, device=device, dtype=self.inv_freq.dtype)137        t = t / self.scaling_factor138 139        freqs = torch.einsum("i,j->ij", t, self.inv_freq)140        # Different from paper, but it uses a different permutation in order to obtain the same calculation141        emb = torch.cat((freqs, freqs), dim=-1)142        self.register_buffer("cos_cached", emb.cos()[None, None, :, :].to(dtype), persistent=False)143        self.register_buffer("sin_cached", emb.sin()[None, None, :, :].to(dtype), persistent=False)144 145 146class CodeShellDynamicNTKScalingRotaryEmbedding(CodeShellRotaryEmbedding):147    """ShellRotaryEmbedding extended with Dynamic NTK scaling. Credits to the Reddit users /u/bloc97 and /u/emozilla"""148 149    def __init__(self, dim, max_position_embeddings=2048, base=10000, device=None, scaling_factor=1.0):150        self.scaling_factor = scaling_factor151        super().__init__(dim, max_position_embeddings, base, device)152 153    def _set_cos_sin_cache(self, seq_len, device, dtype):154        self.max_seq_len_cached = seq_len155 156        if seq_len > self.max_position_embeddings:157            base = self.base * (158                (self.scaling_factor * seq_len / self.max_position_embeddings) - (self.scaling_factor - 1)159            ) ** (self.dim / (self.dim - 2))160            inv_freq = 1.0 / (base ** (torch.arange(0, self.dim, 2).float().to(device) / self.dim))161            self.register_buffer("inv_freq", inv_freq)162 163        t = torch.arange(self.max_seq_len_cached, device=device, dtype=self.inv_freq.dtype)164 165        freqs = torch.einsum("i,j->ij", t, self.inv_freq)166        # Different from paper, but it uses a different permutation in order to obtain the same calculation167        emb = torch.cat((freqs, freqs), dim=-1)168        self.register_buffer("cos_cached", emb.cos()[None, None, :, :].to(dtype), persistent=False)169        self.register_buffer("sin_cached", emb.sin()[None, None, :, :].to(dtype), persistent=False)170 171def rotate_half(x):172    """Rotates half the hidden dims of the input."""173    x1 = x[..., : x.shape[-1] // 2]174    x2 = x[..., x.shape[-1] // 2 :]175    return torch.cat((-x2, x1), dim=-1)176 177 178def apply_rotary_pos_emb(q, k, cos, sin, position_ids):179    # The first two dimensions of cos and sin are always 1, so we can `squeeze` them.180    cos = cos.squeeze(1).squeeze(0)  # [seq_len, dim]181    sin = sin.squeeze(1).squeeze(0)  # [seq_len, dim]182    cos = cos[position_ids].unsqueeze(1)  # [bs, 1, seq_len, dim]183    sin = sin[position_ids].unsqueeze(1)  # [bs, 1, seq_len, dim]184    q_embed = (q * cos) + (rotate_half(q) * sin)185    k_embed = (k * cos) + (rotate_half(k) * sin)186    return q_embed, k_embed187 188def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor:189    """190    This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch,191    num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim)192    """193    batch, num_key_value_heads, slen, head_dim = hidden_states.shape194    if n_rep == 1:195        return hidden_states196    hidden_states = hidden_states[:, :, None, :, :].expand(batch, num_key_value_heads, n_rep, slen, head_dim)197    return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim)198 199class CodeShellAttention(nn.Module):200    def __init__(self, config, layer_idx=None):201        super().__init__()202        self.mask_value = None203        204        self.position_embedding_type = config.position_embedding_type205        self.rope_scaling = config.rope_scaling206        self.max_position_embeddings = config.max_position_embeddings207        208        self.group_query_attention = config.group_query_attention209        self.num_query_groups = config.num_query_groups210        self.num_key_value_groups = config.num_attention_heads // config.num_query_groups211        212        self.embed_dim = config.hidden_size213        self.num_heads = config.num_attention_heads214        self.head_dim = self.embed_dim // self.num_heads215        self.kv_heads = config.num_query_groups if self.group_query_attention else self.num_heads216        self.kv_dim = self.kv_heads * self.head_dim217        self.split_size = self.embed_dim218        if self.head_dim * self.num_heads != self.embed_dim:219            raise ValueError(220                f"`embed_dim` must be divisible by num_heads (got `embed_dim`: {self.embed_dim} and `num_heads`:"221                f" {self.num_heads})."222            )223 224        self.layer_idx = layer_idx225 226        self.c_attn = nn.Linear(self.embed_dim, self.embed_dim + 2 * self.kv_dim)227        self.c_proj = nn.Linear(self.embed_dim, self.embed_dim)228 229        self.attn_dropout = nn.Dropout(config.attn_pdrop)230        self.resid_dropout = nn.Dropout(config.resid_pdrop)231 232        if self.position_embedding_type == "rope":233            self._init_rope()234 235    def _init_rope(self):236        if self.rope_scaling is None:237            self.rotary_emb = CodeShellRotaryEmbedding(self.head_dim, max_position_embeddings=self.max_position_embeddings)238        else:239            scaling_type = self.rope_scaling["type"]240            scaling_factor = self.rope_scaling["factor"]241            if scaling_type == "linear":242                self.rotary_emb = CodeShellLinearScalingRotaryEmbedding(243                    self.head_dim, max_position_embeddings=self.max_position_embeddings, scaling_factor=scaling_factor244                )245            elif scaling_type == "dynamic":246                self.rotary_emb = CodeShellDynamicNTKScalingRotaryEmbedding(247                    self.head_dim, max_position_embeddings=self.max_position_embeddings, scaling_factor=scaling_factor248                )249            else:250                raise ValueError(f"Unknown RoPE scaling type {scaling_type}")251 252 253    def _get_mask_value(self, device, dtype):254        # torch.where expects a tensor. We use a cache to avoid recreating it every time.255        if self.mask_value is None or self.mask_value.dtype != dtype or self.mask_value.device != device:256            self.mask_value = torch.full([], torch.finfo(dtype).min, dtype=dtype, device=device)257        return self.mask_value258 259    def forward(260        self,261        hidden_states: torch.Tensor,262        layer_past: Optional[torch.Tensor] = None,263        attention_mask: Optional[torch.Tensor] = None,264        position_ids: Optional[torch.LongTensor] = None,265        head_mask: Optional[torch.Tensor] = None,266        use_cache: Optional[bool] = False,267        output_attentions: Optional[bool] = False,268    ) -> Union[269        Tuple[torch.Tensor, Optional[torch.Tensor]],270        Tuple[torch.Tensor, Optional[torch.Tensor], Tuple[torch.Tensor, ...]],271    ]:272        bsz, q_len, _ = hidden_states.size()273        query_states, key_states, value_states = self.c_attn(hidden_states).split((self.embed_dim, self.kv_dim, self.kv_dim), dim=2)274        275        query_states = query_states.view(bsz, q_len, self.num_heads, self.head_dim).transpose(1, 2)276        key_states = key_states.view(bsz, q_len, self.num_query_groups, self.head_dim).transpose(1, 2)277        value_states = value_states.view(bsz, q_len, self.num_query_groups, self.head_dim).transpose(1, 2)278        279        kv_seq_len = key_states.shape[-2]280        if layer_past is not None:281            kv_seq_len += layer_past[0].shape[-2]282        cos, sin = self.rotary_emb(value_states, seq_len=kv_seq_len)283        query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin, position_ids)284 285        if layer_past is not None:286            # reuse k, v, self_attention287            key_states = torch.cat([layer_past[0], key_states], dim=2)288            value_states = torch.cat([layer_past[1], value_states], dim=2)289 290        layer_past = (key_states, value_states) if use_cache else None291 292        # repeat k/v heads if n_kv_heads < n_heads293        key_states = repeat_kv(key_states, self.num_heads // self.kv_heads)294        value_states = repeat_kv(value_states, self.num_heads // self.kv_heads)295    296        attn_weights = torch.matmul(query_states, key_states.transpose(2, 3)) / math.sqrt(self.head_dim)297 298        if attn_weights.size() != (bsz, self.num_heads, q_len, kv_seq_len):299            raise ValueError(300                f"Attention weights should be of size {(bsz, self.num_heads, q_len, kv_seq_len)}, but is"301                f" {attn_weights.size()}"302            )303 304        if attention_mask is not None:305            if attention_mask.size() != (bsz, 1, q_len, kv_seq_len):306                raise ValueError(307                    f"Attention mask should be of size {(bsz, 1, q_len, kv_seq_len)}, but is {attention_mask.size()}"308                )309            mask_value = self._get_mask_value(attn_weights.device, attn_weights.dtype)310            # The fused kernel is very slow when the key length is not a multiple of 8, so we skip fusion.311            attn_weights = torch.where(attention_mask, attn_weights, mask_value)312 313        # upcast attention to fp32314        attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query_states.dtype)315        attn_weights = self.attn_dropout(attn_weights)316        attn_output = torch.matmul(attn_weights, value_states)317 318        if attn_output.size() != (bsz, self.num_heads, q_len, self.head_dim):319            raise ValueError(320                f"`attn_output` should be of size {(bsz, self.num_heads, q_len, self.head_dim)}, but is"321                f" {attn_output.size()}"322            )323 324        attn_output = attn_output.transpose(1, 2).contiguous()325        attn_output = attn_output.reshape(bsz, q_len, self.embed_dim)326 327        attn_output = self.c_proj(attn_output)328        attn_output = self.resid_dropout(attn_output)329        330        outputs = (attn_output, layer_past)331        if output_attentions:332            outputs += (attn_weights,)333 334        return outputs # a, present, (attentions)335 336 337class CodeShellMLP(nn.Module):338    def __init__(self, intermediate_size, config):339        super().__init__()340        embed_dim = config.hidden_size341        self.c_fc = nn.Linear(embed_dim, intermediate_size)342        self.c_proj = nn.Linear(intermediate_size, embed_dim)343        self.act = ACT2FN[config.activation_function]344        self.dropout = nn.Dropout(config.resid_pdrop)345 346    # Copied from transformers.models.gpt2.modeling_gpt2.GPT2MLP.forward347    def forward(self, hidden_states: Optional[Tuple[torch.Tensor]]) -> torch.Tensor:348        hidden_states = self.c_fc(hidden_states)349        hidden_states = self.act(hidden_states)350        hidden_states = self.c_proj(hidden_states)351        hidden_states = self.dropout(hidden_states)352        return hidden_states353 354 355class CodeShellBlock(nn.Module):356    def __init__(self, config, layer_idx=None):357        super().__init__()358        hidden_size = config.hidden_size359        self.inner_dim = config.n_inner if config.n_inner is not None else 4 * hidden_size360 361        self.ln_1 = nn.LayerNorm(hidden_size, eps=config.layer_norm_epsilon)362        self.attn = CodeShellAttention(config, layer_idx=layer_idx)363        self.ln_2 = nn.LayerNorm(hidden_size, eps=config.layer_norm_epsilon)364 365        self.mlp = CodeShellMLP(self.inner_dim, config)366 367    def forward(368        self,369        hidden_states: Optional[Tuple[torch.Tensor]],370        layer_past: Optional[torch.Tensor] = None,371        attention_mask: Optional[torch.Tensor] = None,372        position_ids: Optional[torch.LongTensor] = None,373        head_mask: Optional[torch.Tensor] = None,374        encoder_hidden_states: Optional[torch.Tensor] = None,375        encoder_attention_mask: Optional[torch.Tensor] = None,376        use_cache: Optional[bool] = False,377        output_attentions: Optional[bool] = False,378    ) -> Union[379        Tuple[torch.Tensor], Tuple[torch.Tensor, torch.Tensor], Tuple[torch.Tensor, torch.Tensor, torch.Tensor]380    ]:381        residual = hidden_states382        hidden_states = self.ln_1(hidden_states)383        attn_outputs = self.attn(384            hidden_states,385            layer_past=layer_past,386            attention_mask=attention_mask,387            position_ids=position_ids,388            head_mask=head_mask,389            use_cache=use_cache,390            output_attentions=output_attentions,391        )392        attn_output = attn_outputs[0]  # output_attn: a, present, (attentions)393        394        outputs = attn_outputs[1:]395        # residual connection396        hidden_states = attn_output + residual397 398        residual = hidden_states399        hidden_states = self.ln_2(hidden_states)400        feed_forward_hidden_states = self.mlp(hidden_states)401        # residual connection402        hidden_states = residual + feed_forward_hidden_states403 404        if use_cache:405            outputs = (hidden_states,) + outputs406        else:407            outputs = (hidden_states,) + outputs[1:]408 409        return outputs  # hidden_states, present, (attentions, cross_attentions)410 411 412class CodeShellPreTrainedModel(PreTrainedModel):413    """414    An abstract class to handle weights initialization and a simple interface for downloading and loading pretrained415    models.416    """417 418    config_class = CodeShellConfig419    base_model_prefix = "transformer"420    supports_gradient_checkpointing = True421    _no_split_modules = ["ShellBlock"]422    _skip_keys_device_placement = "past_key_values"423 424    def __init__(self, *inputs, **kwargs):425        super().__init__(*inputs, **kwargs)426 427    def _init_weights(self, module):428        """Initialize the weights."""429        if isinstance(module, (CodeShellMLP, CodeShellAttention)):430            # Reinitialize selected weights subject to the OpenAI GPT-2 Paper Scheme:431            #   > A modified initialization which accounts for the accumulation on the residual path with model depth. Scale432            #   > the weights of residual layers at initialization by a factor of 1/โˆšN where N is the # of residual layers.433            #   >   -- GPT-2 :: https://openai.com/blog/better-language-models/434            #435            # Reference (Megatron-LM): https://github.com/NVIDIA/Megatron-LM/blob/main/megatron/model/gpt_model.py436            module.c_proj.weight.data.normal_(437                mean=0.0, std=(self.config.initializer_range / math.sqrt(2 * self.config.n_layer))438            )439            module.c_proj._is_hf_initialized = True440        elif isinstance(module, nn.Linear):441            # Slightly different from the TF version which uses truncated_normal for initialization442            # cf https://github.com/pytorch/pytorch/pull/5617443            module.weight.data.normal_(mean=0.0, std=self.config.initializer_range)444            if module.bias is not None:445                module.bias.data.zero_()446        elif isinstance(module, nn.Embedding):447            module.weight.data.normal_(mean=0.0, std=self.config.initializer_range)448            if module.padding_idx is not None:449                module.weight.data[module.padding_idx].zero_()450        elif isinstance(module, nn.LayerNorm):451            module.bias.data.zero_()452            module.weight.data.fill_(1.0)453 454    # Copied from transformers.models.gpt2.modeling_gpt2.GPT2PreTrainedModel._set_gradient_checkpointing with GPT2->Shell455    def _set_gradient_checkpointing(self, module, value=False):456        if isinstance(module, CodeShellModel):457            module.gradient_checkpointing = value458 459 460GPT_BIGCODE_START_DOCSTRING = r"""461 462    This model inherits from [`PreTrainedModel`]. Check the superclass documentation for the generic methods the463    library implements for all its model (such as downloading or saving, resizing the input embeddings, pruning heads464    etc.)465 466    This model is also a PyTorch [torch.nn.Module](https://pytorch.org/docs/stable/nn.html#torch.nn.Module) subclass.467    Use it as a regular PyTorch Module and refer to the PyTorch documentation for all matter related to general usage468    and behavior.469 470    Parameters:471        config ([`CodeShellConfig`]): Model configuration class with all the parameters of the model.472            Initializing with a config file does not load the weights associated with the model, only the473            configuration. Check out the [`~PreTrainedModel.from_pretrained`] method to load the model weights.474"""475 476GPT_BIGCODE_INPUTS_DOCSTRING = r"""477    Args:478        input_ids (`torch.Tensor` of shape `(batch_size, input_ids_length)`):479            `input_ids_length` = `sequence_length` if `past_key_values` is `None` else480            `past_key_values[0][0].shape[-2]` (`sequence_length` of input past key value states). Indices of input481            sequence tokens in the vocabulary.482 483            If `past_key_values` is used, only `input_ids` that do not have their past calculated should be passed as484            `input_ids`.485 486            Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and487            [`PreTrainedTokenizer.__call__`] for details.488 489            [What are input IDs?](../glossary#input-ids)490        past_key_values (`Tuple[torch.Tensor]` of length `config.n_layers`):491            Contains precomputed hidden-states (key and values in the attention blocks) as computed by the model (see492            `past_key_values` output below). Can be used to speed up sequential decoding. The `input_ids` which have493            their past given to this model should not be passed as `input_ids` as they have already been computed.494        attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):495            Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:496 497            - 1 for tokens that are **not masked**,498            - 0 for tokens that are **masked**.499 500            If `past_key_values` is used, `attention_mask` needs to contain the masking strategy that was used for501            `past_key_values`. In other words, the `attention_mask` always has to have the length:502            `len(past_key_values) + len(input_ids)`503 504            [What are attention masks?](../glossary#attention-mask)505        token_type_ids (`torch.Tensor` of shape `(batch_size, input_ids_length)`, *optional*):506            Segment token indices to indicate first and second portions of the inputs. Indices are selected in `[0,507            1]`:508 509            - 0 corresponds to a *sentence A* token,510            - 1 corresponds to a *sentence B* token.511 512            [What are token type IDs?](../glossary#token-type-ids)513        position_ids (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):514            Indices of positions of each input sequence tokens in the position embeddings. Selected in the range `[0,515            config.max_position_embeddings - 1]`.516 517            [What are position IDs?](../glossary#position-ids)518        head_mask (`torch.Tensor` of shape `(num_heads,)` or `(num_layers, num_heads)`, *optional*):519            Mask to nullify selected heads of the self-attention modules. Mask values selected in `[0, 1]`:520 521            - 1 indicates the head is **not masked**,522            - 0 indicates the head is **masked**.523 524        inputs_embeds (`torch.Tensor` of shape `(batch_size, sequence_length, hidden_size)`, *optional*):525            Optionally, instead of passing `input_ids` you can choose to directly pass an embedded representation. This526            is useful if you want more control over how to convert `input_ids` indices into associated vectors than the527            model's internal embedding lookup matrix.528 529            If `past_key_values` is used, optionally only the last `inputs_embeds` have to be input (see530            `past_key_values`).531        use_cache (`bool`, *optional*):532            If set to `True`, `past_key_values` key value states are returned and can be used to speed up decoding (see533            `past_key_values`).534        output_attentions (`bool`, *optional*):535            Whether or not to return the attentions tensors of all attention layers. See `attentions` under returned536            tensors for more detail.537        output_hidden_states (`bool`, *optional*):538            Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors for539            more detail.540        return_dict (`bool`, *optional*):541            Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.542"""543 544 545@add_start_docstrings(546    "The bare GPT_BIGCODE Model transformer outputting raw hidden-states without any specific head on top.",547    GPT_BIGCODE_START_DOCSTRING,548)549class CodeShellModel(CodeShellPreTrainedModel):550    def __init__(self, config):551        super().__init__(config)552        self.group_query_attention = config.group_query_attention553        self.num_query_groups = config.num_query_groups554        self.position_embedding_type = config.position_embedding_type555        self.embed_dim = config.hidden_size556 557        self.wte = nn.Embedding(config.vocab_size, self.embed_dim)558        if self.position_embedding_type == "learned_absolute":559            self.wpe = nn.Embedding(config.max_position_embeddings, self.embed_dim)560        else:561            pass562 563        self.drop = nn.Dropout(config.embd_pdrop)564        self.h = nn.ModuleList([CodeShellBlock(config, layer_idx=i) for i in range(config.num_hidden_layers)])565        self.ln_f = nn.LayerNorm(self.embed_dim, eps=config.layer_norm_epsilon)566 567        max_positions = config.max_position_embeddings568        self.register_buffer(569            "bias", torch.tril(torch.ones((max_positions, max_positions), dtype=torch.bool)), persistent=False570        )571 572        self.gradient_checkpointing = False573 574        # Initialize weights and apply final processing575        self.post_init()576 577    def get_input_embeddings(self):578        return self.wte579 580    def set_input_embeddings(self, new_embeddings):581        self.wte = new_embeddings582 583    @add_start_docstrings_to_model_forward(GPT_BIGCODE_INPUTS_DOCSTRING)584    def forward(585        self,586        input_ids: Optional[torch.Tensor] = None,587        past_key_values: Optional[List[torch.Tensor]] = None,588        attention_mask: Optional[torch.Tensor] = None,589        token_type_ids: Optional[torch.Tensor] = None,590        position_ids: Optional[torch.Tensor] = None,591        head_mask: Optional[torch.Tensor] = None,592        inputs_embeds: Optional[torch.Tensor] = None,593        encoder_hidden_states: Optional[torch.Tensor] = None,594        encoder_attention_mask: Optional[torch.Tensor] = None,595        use_cache: Optional[bool] = None,596        output_attentions: Optional[bool] = None,597        output_hidden_states: Optional[bool] = None,598        return_dict: Optional[bool] = None,599    ) -> Union[Tuple, BaseModelOutputWithPastAndCrossAttentions]:600        output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions601        output_hidden_states = (602            output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states603        )604        use_cache = use_cache if use_cache is not None else self.config.use_cache605        return_dict = return_dict if return_dict is not None else self.config.use_return_dict606 607        if input_ids is not None and inputs_embeds is not None:608            raise ValueError("You cannot specify both input_ids and inputs_embeds at the same time")609        elif input_ids is not None:610            input_shape = input_ids.size()611            input_ids = input_ids.reshape(-1, input_shape[-1])612            batch_size = input_ids.shape[0]613        elif inputs_embeds is not None:614            input_shape = inputs_embeds.size()[:-1]615            batch_size = inputs_embeds.shape[0]616        else:617            raise ValueError("You have to specify either input_ids or inputs_embeds")618 619        if batch_size <= 0:620            raise ValueError("batch_size has to be defined and > 0")621 622        device = input_ids.device if input_ids is not None else inputs_embeds.device623 624        if token_type_ids is not None:625            token_type_ids = token_type_ids.reshape(-1, input_shape[-1])626        if position_ids is not None:627            position_ids = position_ids.reshape(-1, input_shape[-1])628 629        if past_key_values is None:630            past_length = 0631            past_key_values = tuple([None] * len(self.h))632        else:633            past_length = past_key_values[0][0].size(-2)634 635        if attention_mask is not None and len(attention_mask.shape) == 2 and position_ids is None:636            # create position_ids on the fly for batch generation637            position_ids = attention_mask.long().cumsum(-1) - 1638            position_ids.masked_fill_(attention_mask == 0, 1)639            if past_length > 0:640                position_ids = position_ids[:, past_length : input_shape[-1] + past_length :]641        elif position_ids is None:642            position_ids = torch.arange(past_length, input_shape[-1] + past_length, dtype=torch.long, device=device)643            position_ids = position_ids.unsqueeze(0).reshape(-1, input_shape[-1])644 645        # Self-attention mask.646        query_length = input_shape[-1]647        key_length = past_length + query_length648        self_attention_mask = self.bias[None, key_length - query_length : key_length, :key_length]649 650        if attention_mask is not None:651            self_attention_mask = self_attention_mask * attention_mask.reshape(batch_size, 1, -1).to(652                dtype=torch.bool, device=self_attention_mask.device653            )654 655        # MQA models: (batch_size, query_length, n_heads, key_length)656        # MHA models: (batch_size, n_heads, query_length, key_length)657        attention_mask = self_attention_mask.unsqueeze(1)658 659        encoder_attention_mask = None660 661        # Prepare head mask if needed662        # 1.0 in head_mask indicate we keep the head663        # attention_probs has shape bsz x n_heads x N x N664        # head_mask has shape n_layer x batch x n_heads x N x N665        head_mask = self.get_head_mask(head_mask, self.config.n_layer)666 667        if inputs_embeds is None:668            inputs_embeds = self.wte(input_ids)669        670        hidden_states = inputs_embeds671        if self.position_embedding_type == "learned_absolute":672            position_embeds = self.wpe(position_ids)673            hidden_states = hidden_states + position_embeds674 675        if token_type_ids is not None:676            token_type_embeds = self.wte(token_type_ids)677            hidden_states = hidden_states + token_type_embeds678 679        hidden_states = self.drop(hidden_states)680 681        output_shape = input_shape + (hidden_states.size(-1),)682 683        presents = [] if use_cache else None684        all_self_attentions = () if output_attentions else None685        all_hidden_states = () if output_hidden_states else None686        for i, (block, layer_past) in enumerate(zip(self.h, past_key_values)):687            if output_hidden_states:688                all_hidden_states = all_hidden_states + (hidden_states,)689 690            if self.gradient_checkpointing and self.training:691 692                def create_custom_forward(module):693                    def custom_forward(*inputs):694                        # None for past_key_value695                        return module(*inputs, use_cache, output_attentions)696 697                    return custom_forward698 699                outputs = torch.utils.checkpoint.checkpoint(700                    create_custom_forward(block),701                    hidden_states,702                    None,703                    attention_mask,704                    position_ids,705                    head_mask[i],706                    encoder_hidden_states,707                    encoder_attention_mask,708                )709            else:710                outputs = block(711                    hidden_states,712                    layer_past=layer_past,713                    attention_mask=attention_mask,714                    position_ids=position_ids,715                    head_mask=head_mask[i],716                    encoder_hidden_states=encoder_hidden_states,717                    encoder_attention_mask=encoder_attention_mask,718                    use_cache=use_cache,719                    output_attentions=output_attentions,720                )721 722            hidden_states = outputs[0]723            if use_cache:724                presents.append(outputs[1])725 726            if output_attentions:727                all_self_attentions = all_self_attentions + (outputs[2 if use_cache else 1],)728        729        hidden_states = self.ln_f(hidden_states)730        hidden_states = hidden_states.reshape(output_shape)731        # Add last hidden state732        if output_hidden_states:733            all_hidden_states = all_hidden_states + (hidden_states,)734        735        736        if not return_dict:737            return tuple(738                v739                for v in [hidden_states, presents, all_hidden_states, all_self_attentions]740                if v is not None741            )742 743        return BaseModelOutputWithPastAndCrossAttentions(744            last_hidden_state=hidden_states,745            past_key_values=presents,746            hidden_states=all_hidden_states,747            attentions=all_self_attentions,748        )749    750class EndOfFunctionCriteria(StoppingCriteria):751    """Custom `StoppingCriteria` which checks if all generated functions in the batch are completed."""752    def __init__(self, input_lengths, eof_strings, tokenizer):753        self.input_lengths = input_lengths754        self.eof_strings = eof_strings755        self.tokenizer = tokenizer756 757    def __call__(self, input_ids, scores, **kwargs):758        """Returns true if all generated sequences contain any of the end-of-function strings."""759        decoded_generations = []760        for _input_ids, input_length in zip(input_ids, self.input_lengths):761            decoded_generations.append(self.tokenizer.decode(_input_ids[input_length:]))762        done = []763        for decoded_generation in decoded_generations:764            done.append(765                any(766                    [767                        stop_string in decoded_generation768                        for stop_string in self.eof_strings769                    ]770                )771            )772        return all(done)773 774class TextIterStreamer:775    def __init__(self, tokenizer, skip_prompt=False, skip_special_tokens=False):776        self.tokenizer = tokenizer777        self.skip_prompt = skip_prompt778        self.skip_special_tokens = skip_special_tokens779        self.tokens = []780        self.text_queue = Queue()781        self.next_tokens_are_prompt = True782 783    def put(self, value):784        if self.skip_prompt and self.next_tokens_are_prompt:785            self.next_tokens_are_prompt = False786        else:787            if len(value.shape) > 1:788                value = value[0]789            self.tokens.extend(value.tolist())790            self.text_queue.put(791                self.tokenizer.decode(self.tokens, skip_special_tokens=self.skip_special_tokens))792 793    def end(self):794        self.text_queue.put(None)795 796    def __iter__(self):797        return self798 799    def __next__(self):800        value = self.text_queue.get()801        if value is None:802            raise StopIteration()803        else:804            return value805 806 807@add_start_docstrings(808    """809    The GPT_BIGCODE Model transformer with a language modeling head on top (linear layer with weights tied to the input810    embeddings).811    """,812    GPT_BIGCODE_START_DOCSTRING,813)814class CodeShellForCausalLM(CodeShellPreTrainedModel):815    _tied_weights_keys = ["lm_head.weight"]816 817    def __init__(self, config):818        super().__init__(config)819        self.transformer = CodeShellModel(config)820        self.lm_head = nn.Linear(config.n_embd, config.vocab_size, bias=False)821 822        # Initialize weights and apply final processing823        self.post_init()824 825    def quantize(self, bits: int):826        try:827            import bitsandbytes828            from .quantizer import quantize829        except ImportError:830            raise ImportError(f"Needs bitsandbytes to run quantize.")831        return quantize(self, bits)832 833    def get_output_embeddings(self):834        return self.lm_head835 836    def set_output_embeddings(self, new_embeddings):837        self.lm_head = new_embeddings838 839    def prepare_inputs_for_generation(self, input_ids, past_key_values=None, inputs_embeds=None, **kwargs):840        token_type_ids = kwargs.get("token_type_ids", None)841        # only last token for inputs_ids if past is defined in kwargs842        if past_key_values:843            input_ids = input_ids[:, -1].unsqueeze(-1)844            if token_type_ids is not None:845                token_type_ids = token_type_ids[:, -1].unsqueeze(-1)846 847        attention_mask = kwargs.get("attention_mask", None)848        position_ids = kwargs.get("position_ids", None)849 850        if attention_mask is not None and position_ids is None:851            # create position_ids on the fly for batch generation852            position_ids = attention_mask.long().cumsum(-1) - 1853            position_ids.masked_fill_(attention_mask == 0, 1)854            if past_key_values:855                position_ids = position_ids[:, -1].unsqueeze(-1)856        else:857            position_ids = None858 859        # if `inputs_embeds` are passed, we only want to use them in the 1st generation step860        if inputs_embeds is not None and past_key_values is None:861            model_inputs = {"inputs_embeds": inputs_embeds}862        else:863            model_inputs = {"input_ids": input_ids}864 865        model_inputs.update(866            {867                "past_key_values": past_key_values,868                "use_cache": kwargs.get("use_cache"),869                "position_ids": position_ids,870                "attention_mask": attention_mask,871                "token_type_ids": token_type_ids,872            }873        )874        return model_inputs875 876    @add_start_docstrings_to_model_forward(GPT_BIGCODE_INPUTS_DOCSTRING)877    def forward(878        self,879        input_ids: Optional[torch.Tensor] = None,880        past_key_values: Optional[Tuple[Tuple[torch.Tensor]]] = None,881        attention_mask: Optional[torch.Tensor] = None,882        token_type_ids: Optional[torch.Tensor] = None,883        position_ids: Optional[torch.Tensor] = None,884        head_mask: Optional[torch.Tensor] = None,885        inputs_embeds: Optional[torch.Tensor] = None,886        encoder_hidden_states: Optional[torch.Tensor] = None,887        encoder_attention_mask: Optional[torch.Tensor] = None,888        labels: Optional[torch.Tensor] = None,889        use_cache: Optional[bool] = None,890        output_attentions: Optional[bool] = None,891        output_hidden_states: Optional[bool] = None,892        return_dict: Optional[bool] = None,893    ) -> Union[Tuple, CausalLMOutputWithCrossAttentions]:894        r"""895        labels (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):896            Labels for language modeling. Note that the labels **are shifted** inside the model, i.e. you can set897            `labels = input_ids` Indices are selected in `[-100, 0, ..., config.vocab_size]` All labels set to `-100`898            are ignored (masked), the loss is only computed for labels in `[0, ..., config.vocab_size]`899        """900        return_dict = return_dict if return_dict is not None else self.config.use_return_dict901 902        transformer_outputs = self.transformer(903            input_ids,904            past_key_values=past_key_values,905            attention_mask=attention_mask,906            token_type_ids=token_type_ids,907            position_ids=position_ids,908            head_mask=head_mask,909            inputs_embeds=inputs_embeds,910            encoder_hidden_states=encoder_hidden_states,911            encoder_attention_mask=encoder_attention_mask,912            use_cache=use_cache,913            output_attentions=output_attentions,914            output_hidden_states=output_hidden_states,915            return_dict=return_dict,916        )917        hidden_states = transformer_outputs[0]918        lm_logits = self.lm_head(hidden_states)919        loss = None920        if labels is not None:921            # Shift so that tokens < n predict n922            shift_logits = lm_logits[..., :-1, :].contiguous()923            shift_labels = labels[..., 1:].contiguous().to(shift_logits.device)924            # Flatten the tokens925            loss_fct = CrossEntropyLoss()926            loss = loss_fct(shift_logits.reshape(-1, shift_logits.size(-1)), shift_labels.reshape(-1))927 928        if not return_dict:929            output = (lm_logits,) + transformer_outputs[1:]930            return ((loss,) + output) if loss is not None else output931 932        return CausalLMOutputWithCrossAttentions(933            loss=loss,934            logits=lm_logits,935            past_key_values=transformer_outputs.past_key_values,936            hidden_states=transformer_outputs.hidden_states,937            attentions=transformer_outputs.attentions,938        )939 940    @staticmethod941    def _reorder_cache(past_key_values, beam_idx):942        reordered_past = ()943        for layer_past in past_key_values:944            reordered_past += (945                tuple(past_state.index_select(0, beam_idx.to(past_state.device)) for past_state in layer_past),946            )947        return reordered_past948    949 950    def build_chat_input(self, query, history, tokenizer, max_new_tokens=None):951        user_name = "<human>:"952        ai_name = "<assistant>:"953        stop = "<|endoftext|>"954 955        prompt = ''956        for q, r in history:957            prompt += f"{user_name}{q}{stop}"958            prompt += f"{ai_name}{r}{stop}"959        prompt += f"{user_name}{query}{stop}"960        prompt += ai_name.rstrip()961 962        max_new_tokens = max_new_tokens or self.generation_config.max_new_tokens or 1024963        max_input_tokens = self.config.n_positions - max_new_tokens964 965        input_tokens = tokenizer.encode(prompt)966        input_tokens = input_tokens[-max_input_tokens:]  # truncate left967        return torch.LongTensor([input_tokens]).to(self.device)968 969    def chat(self, query, history, tokenizer, stream=False,970            generation_config: Optional[GenerationConfig]=None):971        generation_config = generation_config or self.generation_config972        input_ids = self.build_chat_input(query, history, tokenizer, generation_config.max_new_tokens)973        stopping_criteria = StoppingCriteriaList(974            [EndOfFunctionCriteria([len(input_ids[0])], ["<|endoftext|>", "<human>:"], tokenizer)]975        )976        977        if stream:978            streamer = TextIterStreamer(tokenizer, skip_prompt=True, skip_special_tokens=True)979            Thread(target=self.generate, kwargs=dict(980                inputs=input_ids, streamer=streamer,981                stopping_criteria = stopping_criteria,982                generation_config=generation_config,983            )).start()984            return streamer985        else:986            outputs = self.generate(input_ids, generation_config=generation_config, stopping_criteria = stopping_criteria)987            response = tokenizer.decode(outputs[0][len(input_ids[0]):], skip_special_tokens=True)988            return response989        990    def generate_stream(self, prompt, tokenizer, generation_config=None, **kwargs):991        generation_config = generation_config or self.generation_config992        max_input_tokens = self.config.n_positions - self.generation_config.max_new_tokens993 994        input_ids = tokenizer.encode(prompt)995        input_ids = input_ids[-max_input_tokens:]  # truncate left996 997        stopping_criteria = StoppingCriteriaList(998            [EndOfFunctionCriteria([len(input_ids[0])], ["<|endoftext|>", "<human>:"], tokenizer)]999        )1000 1001        streamer = TextIterStreamer(tokenizer, skip_prompt=True, skip_special_tokens=True)1002        Thread(target=self.generate, kwargs=dict(1003            inputs=input_ids, stopping_criteria=stopping_criteria, **kwargs1004        )).start()1005        return streamer1006 1007 1008class CodeShell4bitForCausalLM(CodeShellForCausalLM):1009    def __init__(self, config):1010        CodeShellPreTrainedModel.__init__(self, config)  1011        self.transformer = CodeShellModel(config)1012        self.lm_head = nn.Linear(config.n_embd, config.vocab_size, bias=False)1013 1014        try:1015            import bitsandbytes1016            from .quantizer import quantize_offline1017            quantize_offline(self)1018        except ImportError:1019            raise ImportError(f"Needs bitsandbytes to run quantize.")1020        1021        self.post_init()1022 1023    @classmethod1024    def from_pretrained(1025        cls,1026        pretrained_model_name_or_path: Optional[Union[str, os.PathLike]],1027        *model_args,1028        config: Optional[Union[PretrainedConfig, str, os.PathLike]] = None,1029        cache_dir: Optional[Union[str, os.PathLike]] = None,1030        ignore_mismatched_sizes: bool = False,1031        force_download: bool = False,1032        local_files_only: bool = False,1033        token: Optional[Union[str, bool]] = None,1034        revision: str = "main",1035        use_safetensors: bool = None,1036        **kwargs,1037    ):1038        if not isinstance(config, PretrainedConfig):1039            config_path = config if config is not None else pretrained_model_name_or_path1040            config, _ = cls.config_class.from_pretrained(1041                config_path,1042                cache_dir=cache_dir,1043                return_unused_kwargs=True,1044                force_download=force_download,1045                resume_download=False,1046                proxies=None,1047                local_files_only=local_files_only,1048                token=token,1049                revision=revision,1050                subfolder="",1051                _from_auto=False,1052                _from_pipeline=None,1053                **kwargs,1054            )1055            1056        # Load config if we don't provide a configuration1057        from .quantizer import load_state_dict_for_qunantied_model1058        model = cls(config)1059        state_dict = torch.load(os.path.join(pretrained_model_name_or_path, 'pytorch_model.bin'), map_location="cpu") 1060        model = load_state_dict_for_qunantied_model(model, state_dict)1061        model.eval()1062        1063        # If it is a model with generation capabilities, attempt to load the generation config1064        if model.can_generate():1065            try:1066                model.generation_config = GenerationConfig.from_pretrained(1067                    pretrained_model_name_or_path,1068                    cache_dir=cache_dir,1069                    force_download=force_download,1070                    resume_download=False,1071                    proxies=None,1072                    local_files_only=local_files_only,1073                    token=token,1074                    revision=revision,1075                    subfolder="",1076                    _from_auto=False,1077                    _from_pipeline=None,1078                    **kwargs,1079                )1080            except (OSError, TypeError):1081                pass1082 1083        device_map = kwargs.pop("device_map", None)1084        if device_map is not None:1085            model = model.to(torch.device(device_map))1086        1087        return model1088