WisdomShell/Shell-7B-Chat
118
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 