codefuse-ai/CodeFuse-DevOps-Model-7B-Chat
1021
1# Copyright (c) Alibaba Cloud.2#3# This source code is licensed under the license found in the4# LICENSE file in the root directory of this source tree.5 6import importlib7import math8from typing import TYPE_CHECKING, Optional, Tuple, Union, Callable, List, Any, Generator9 10import torch11import torch.nn.functional as F12import torch.utils.checkpoint13from torch.cuda.amp import autocast14 15from torch.nn import CrossEntropyLoss16from transformers import PreTrainedTokenizer, GenerationConfig, StoppingCriteriaList17from transformers.generation.logits_process import LogitsProcessorList18 19if TYPE_CHECKING:20 from transformers.generation.streamers import BaseStreamer21from transformers.generation.utils import GenerateOutput22from transformers.modeling_outputs import (23 BaseModelOutputWithPast,24 CausalLMOutputWithPast,25)26from transformers.modeling_utils import PreTrainedModel27from transformers.utils import logging28 29try:30 from einops import rearrange31except ImportError:32 rearrange = None33from torch import nn34 35SUPPORT_CUDA = torch.cuda.is_available()36SUPPORT_BF16 = SUPPORT_CUDA and torch.cuda.is_bf16_supported()37SUPPORT_FP16 = SUPPORT_CUDA and torch.cuda.get_device_capability(0)[0] >= 738 39from .configuration_qwen import QWenConfig40from .qwen_generation_utils import (41 HistoryType,42 make_context,43 decode_tokens,44 get_stop_words_ids,45 StopWordsLogitsProcessor,46)47 48# from loguru import logger49logger = logging.get_logger(__name__)50 51_CHECKPOINT_FOR_DOC = "qwen"52_CONFIG_FOR_DOC = "QWenConfig"53 54QWen_PRETRAINED_MODEL_ARCHIVE_LIST = ["qwen-7b"]55 56_ERROR_BAD_CHAT_FORMAT = """\57We detect you are probably using the pretrained model (rather than chat model) for chatting, since the chat_format in generation_config is not "chatml".58If you are directly using the model downloaded from Huggingface, please make sure you are using our "Qwen/Qwen-7B-Chat" Huggingface model (rather than "Qwen/Qwen-7B") when you call model.chat().59我们检测到您可能在使用预训练模型(而非chat模型)进行多轮chat,因为您当前在generation_config指定的chat_format,并未设置为我们在对话中所支持的"chatml"格式。60如果您在直接使用我们从Huggingface提供的模型,请确保您在调用model.chat()时,使用的是"Qwen/Qwen-7B-Chat"模型(而非"Qwen/Qwen-7B"预训练模型)。61"""62 63_SENTINEL = object()64_ERROR_STREAM_IN_CHAT = """\65Pass argument `stream` to model.chat() is buggy, deprecated, and marked for removal. Please use model.chat_stream(...) instead of model.chat(..., stream=True).66向model.chat()传入参数stream的用法可能存在Bug,该用法已被废弃,将在未来被移除。请使用model.chat_stream(...)代替model.chat(..., stream=True)。67"""68 69apply_rotary_emb_func = None70rms_norm = None71flash_attn_unpadded_func = None72 73 74def _import_flash_attn():75 global apply_rotary_emb_func, rms_norm, flash_attn_unpadded_func76 try:77 from flash_attn.layers.rotary import apply_rotary_emb_func as __apply_rotary_emb_func78 apply_rotary_emb_func = __apply_rotary_emb_func79 print('Using flash_attn rope')80 except ImportError:81 logger.warn(82 "Warning: import flash_attn rotary fail, please install FlashAttention rotary to get higher efficiency "83 "https://github.com/Dao-AILab/flash-attention/tree/main/csrc/rotary"84 )85 86 try:87 from flash_attn.ops.rms_norm import rms_norm as __rms_norm88 rms_norm = __rms_norm89 print('Using flash_attn rms_norm')90 except ImportError:91 logger.warn(92 "Warning: import flash_attn rms_norm fail, please install FlashAttention layer_norm to get higher efficiency "93 "https://github.com/Dao-AILab/flash-attention/tree/main/csrc/layer_norm"94 )95 96 try:97 import flash_attn98 if not hasattr(flash_attn, '__version__'):99 from flash_attn.flash_attn_interface import flash_attn_unpadded_func as __flash_attn_unpadded_func100 else:101 if int(flash_attn.__version__.split(".")[0]) >= 2:102 from flash_attn.flash_attn_interface import flash_attn_varlen_func as __flash_attn_unpadded_func103 else:104 from flash_attn.flash_attn_interface import flash_attn_unpadded_func as __flash_attn_unpadded_func105 flash_attn_unpadded_func = __flash_attn_unpadded_func106 107 print('Using flash_attn attention func')108 except ImportError:109 logger.warn(110 "Warning: import flash_attn fail, please install FlashAttention to get higher efficiency "111 "https://github.com/Dao-AILab/flash-attention"112 )113 114 115class FlashSelfAttention(torch.nn.Module):116 def __init__(117 self,118 causal=False,119 softmax_scale=None,120 attention_dropout=0.0,121 ):122 super().__init__()123 assert flash_attn_unpadded_func is not None, (124 "Please install FlashAttention first, " "e.g., with pip install flash-attn"125 )126 assert (127 rearrange is not None128 ), "Please install einops first, e.g., with pip install einops"129 self.causal = causal130 self.softmax_scale = softmax_scale131 self.dropout_p = attention_dropout132 133 def forward(self, q, k, v):134 assert all((i.dtype in [torch.float16, torch.bfloat16] for i in (q, k, v)))135 assert all((i.is_cuda for i in (q, k, v)))136 batch_size, seqlen_q = q.shape[0], q.shape[1]137 seqlen_k = k.shape[1]138 q, k, v = [rearrange(x, "b s ... -> (b s) ...") for x in [q, k, v]]139 cu_seqlens_q = torch.arange(140 0,141 (batch_size + 1) * seqlen_q,142 step=seqlen_q,143 dtype=torch.int32,144 device=q.device,145 )146 147 if self.training:148 assert seqlen_k == seqlen_q149 150 is_causal = self.causal151 cu_seqlens_k = cu_seqlens_q152 else:153 is_causal = seqlen_q == seqlen_k154 cu_seqlens_k = torch.arange(155 0,156 (batch_size + 1) * seqlen_k,157 step=seqlen_k,158 dtype=torch.int32,159 device=q.device,160 )161 self.dropout_p = 0162 output = flash_attn_unpadded_func(163 q,164 k,165 v,166 cu_seqlens_q,167 cu_seqlens_k,168 seqlen_q,169 seqlen_k,170 self.dropout_p,171 softmax_scale=self.softmax_scale,172 causal=is_causal,173 )174 175 output = rearrange(output, "(b s) ... -> b s ...", b=batch_size)176 return output177 178 179class QWenAttention(nn.Module):180 def __init__(self, config, layer_number=None):181 super().__init__()182 183 max_positions = config.max_position_embeddings184 self.register_buffer(185 "bias",186 torch.tril(187 torch.ones((max_positions, max_positions), dtype=torch.bool)188 ).view(1, 1, max_positions, max_positions),189 persistent=False,190 )191 self.register_buffer("masked_bias", torch.tensor(-1e4), persistent=False)192 self.layer_number = max(1, layer_number)193 self.params_dtype = config.params_dtype194 self.seq_length = config.seq_length195 196 self.hidden_size = config.hidden_size197 self.split_size = config.hidden_size198 self.num_heads = config.num_attention_heads199 self.head_dim = self.hidden_size // self.num_heads200 201 self.use_flash_attn = config.use_flash_attn202 self.scale_attn_weights = True203 204 self.layer_idx = None205 206 self.projection_size = config.kv_channels * config.num_attention_heads207 208 assert self.projection_size % config.num_attention_heads == 0209 self.hidden_size_per_attention_head = (210 self.projection_size // config.num_attention_heads211 )212 213 self.c_attn = nn.Linear(config.hidden_size, 3 * self.projection_size)214 215 self.c_proj = nn.Linear(216 config.hidden_size, self.projection_size, bias=not config.no_bias217 )218 219 self.is_fp32 = not (config.bf16 or config.fp16)220 if (221 self.use_flash_attn222 and flash_attn_unpadded_func is not None223 and not self.is_fp32224 ):225 self.core_attention_flash = FlashSelfAttention(226 causal=True, attention_dropout=config.attn_pdrop227 )228 229 self.bf16 = config.bf16230 231 if config.rotary_pct == 1.0:232 self.rotary_ndims = None233 else:234 assert config.rotary_pct < 1235 self.rotary_ndims = int(236 self.hidden_size_per_attention_head * config.rotary_pct237 )238 dim = (239 self.rotary_ndims240 if self.rotary_ndims is not None241 else self.hidden_size_per_attention_head242 )243 self.rotary_emb = RotaryEmbedding(dim, base=config.rotary_emb_base)244 245 self.use_dynamic_ntk = config.use_dynamic_ntk246 self.use_logn_attn = config.use_logn_attn247 248 logn_list = [249 math.log(i, self.seq_length) if i > self.seq_length else 1250 for i in range(1, 32768)251 ]252 self.logn_tensor = torch.Tensor(logn_list)[None, :, None, None]253 self._ntk_cached = 1.0254 255 self.attn_dropout = nn.Dropout(config.attn_pdrop)256 257 def _attn(self, query, key, value, attention_mask=None, head_mask=None):258 attn_weights = torch.matmul(query, key.transpose(-1, -2))259 260 if self.scale_attn_weights:261 attn_weights = attn_weights / torch.full(262 [],263 value.size(-1) ** 0.5,264 dtype=attn_weights.dtype,265 device=attn_weights.device,266 )267 268 query_length, key_length = query.size(-2), key.size(-2)269 causal_mask = self.bias[270 :, :, key_length - query_length : key_length, :key_length271 ]272 mask_value = torch.finfo(attn_weights.dtype).min273 mask_value = torch.full([], mask_value, dtype=attn_weights.dtype).to(274 attn_weights.device275 )276 attn_weights = torch.where(277 causal_mask, attn_weights.to(attn_weights.dtype), mask_value278 )279 280 if attention_mask is not None:281 # Apply the attention mask282 attn_weights = attn_weights + attention_mask283 284 attn_weights = nn.functional.softmax(attn_weights, dim=-1)285 286 attn_weights = attn_weights.type(value.dtype)287 attn_weights = self.attn_dropout(attn_weights)288 289 if head_mask is not None:290 attn_weights = attn_weights * head_mask291 292 attn_output = torch.matmul(attn_weights, value)293 attn_output = attn_output.transpose(1, 2)294 295 return attn_output, attn_weights296 297 def _upcast_and_reordered_attn(298 self, query, key, value, attention_mask=None, head_mask=None299 ):300 bsz, num_heads, q_seq_len, dk = query.size()301 _, _, k_seq_len, _ = key.size()302 303 attn_weights = torch.empty(304 bsz * num_heads,305 q_seq_len,306 k_seq_len,307 dtype=torch.float32,308 device=query.device,309 )310 311 scale_factor = 1.0312 if self.scale_attn_weights:313 scale_factor /= float(value.size(-1)) ** 0.5314 315 with autocast(enabled=False):316 q, k = query.reshape(-1, q_seq_len, dk), key.transpose(-1, -2).reshape(317 -1, dk, k_seq_len318 )319 attn_weights = torch.baddbmm(320 attn_weights, q.float(), k.float(), beta=0, alpha=scale_factor321 )322 attn_weights = attn_weights.reshape(bsz, num_heads, q_seq_len, k_seq_len)323 324 query_length, key_length = query.size(-2), key.size(-2)325 causal_mask = self.bias[326 :, :, key_length - query_length : key_length, :key_length327 ]328 mask_value = torch.finfo(attn_weights.dtype).min329 mask_value = torch.tensor(mask_value, dtype=attn_weights.dtype).to(330 attn_weights.device331 )332 attn_weights = torch.where(causal_mask, attn_weights, mask_value)333 334 if attention_mask is not None:335 attn_weights = attn_weights + attention_mask336 337 attn_weights = nn.functional.softmax(attn_weights, dim=-1)338 339 if attn_weights.dtype != torch.float32:340 raise RuntimeError(341 "Error with upcasting, attn_weights does not have dtype torch.float32"342 )343 attn_weights = attn_weights.type(value.dtype)344 attn_weights = self.attn_dropout(attn_weights)345 346 if head_mask is not None:347 attn_weights = attn_weights * head_mask348 349 attn_output = torch.matmul(attn_weights, value)350 351 return attn_output, attn_weights352 353 def _split_heads(self, tensor, num_heads, attn_head_size):354 new_shape = tensor.size()[:-1] + (num_heads, attn_head_size)355 tensor = tensor.view(new_shape)356 return tensor357 358 def _merge_heads(self, tensor, num_heads, attn_head_size):359 tensor = tensor.contiguous()360 new_shape = tensor.size()[:-2] + (num_heads * attn_head_size,)361 return tensor.view(new_shape)362 363 def forward(364 self,365 hidden_states: Optional[Tuple[torch.FloatTensor]],366 layer_past: Optional[Tuple[torch.Tensor]] = None,367 attention_mask: Optional[torch.FloatTensor] = None,368 head_mask: Optional[torch.FloatTensor] = None,369 encoder_hidden_states: Optional[torch.Tensor] = None,370 encoder_attention_mask: Optional[torch.FloatTensor] = None,371 output_attentions: Optional[bool] = False,372 use_cache: Optional[bool] = False,373 ):374 375 mixed_x_layer = self.c_attn(hidden_states)376 query, key, value = mixed_x_layer.split(self.split_size, dim=2)377 378 query = self._split_heads(query, self.num_heads, self.head_dim)379 key = self._split_heads(key, self.num_heads, self.head_dim)380 value = self._split_heads(value, self.num_heads, self.head_dim)381 382 kv_seq_len = hidden_states.size()[1]383 if layer_past:384 # layer past[0] shape: bs * seq_len * head_num * dim385 kv_seq_len += layer_past[0].shape[1]386 if (387 self.use_dynamic_ntk388 and kv_seq_len == hidden_states.size()[1]389 and not self.training390 ):391 context_value = math.log(kv_seq_len / self.seq_length, 2) + 1392 ntk_alpha = 2 ** math.ceil(context_value) - 1393 ntk_alpha = max(ntk_alpha, 1)394 self._ntk_cached = ntk_alpha395 else:396 ntk_alpha = self._ntk_cached397 rotary_pos_emb = self.rotary_emb(kv_seq_len, ntk_alpha=ntk_alpha).to(398 hidden_states.device399 )400 401 if rotary_pos_emb is not None:402 if isinstance(rotary_pos_emb, tuple):403 rotary_pos_emb = rotary_pos_emb404 else:405 rotary_pos_emb = (rotary_pos_emb,) * 2406 407 if rotary_pos_emb is not None:408 q_pos_emb, k_pos_emb = rotary_pos_emb409 # Slice the pos emb for current inference410 cur_len = query.shape[1]411 q_pos_emb = q_pos_emb[:, -cur_len:, :, :]412 k_pos_emb = k_pos_emb[:, -cur_len:, :, :]413 query = apply_rotary_pos_emb(query, q_pos_emb)414 key = apply_rotary_pos_emb(key, k_pos_emb)415 416 if layer_past is not None:417 past_key, past_value = layer_past[0], layer_past[1]418 key = torch.cat((past_key, key), dim=1)419 value = torch.cat((past_value, value), dim=1)420 421 if use_cache:422 present = (key, value)423 else:424 present = None425 426 if self.use_logn_attn and not self.training:427 if self.logn_tensor.device != query.device or self.logn_tensor.dtype != query.dtype:428 self.logn_tensor = self.logn_tensor.to(query.device).type_as(query)429 seq_start = key.size(1) - query.size(1)430 seq_end = key.size(1)431 logn_tensor = self.logn_tensor[:, seq_start:seq_end, :, :]432 query = query * logn_tensor.expand_as(query)433 434 if (435 self.use_flash_attn436 and flash_attn_unpadded_func is not None437 and not self.is_fp32438 and query.is_cuda439 ):440 q, k, v = query, key, value441 context_layer = self.core_attention_flash(q, k, v)442 443 context_layer = rearrange(444 context_layer, "b s h d -> b s (h d)"445 ).contiguous()446 else:447 query = query.permute(0, 2, 1, 3)448 key = key.permute(0, 2, 1, 3)449 value = value.permute(0, 2, 1, 3)450 attn_output, attn_weight = self._attn(451 query, key, value, attention_mask, head_mask452 )453 context_layer = self._merge_heads(454 attn_output, self.num_heads, self.head_dim455 )456 457 attn_output = self.c_proj(context_layer)458 outputs = (attn_output, present)459 if output_attentions:460 if (461 self.use_flash_attn462 and flash_attn_unpadded_func is not None463 and not self.is_fp32464 ):465 raise ValueError("Cannot output attentions while using flash-attn")466 else:467 outputs += (attn_weight,)468 469 return outputs470 471 472class QWenMLP(nn.Module):473 def __init__(self, config):474 super().__init__()475 self.w1 = nn.Linear(476 config.hidden_size, config.ffn_hidden_size // 2, bias=not config.no_bias477 )478 self.w2 = nn.Linear(479 config.hidden_size, config.ffn_hidden_size // 2, bias=not config.no_bias480 )481 ff_dim_in = config.ffn_hidden_size // 2482 self.c_proj = nn.Linear(ff_dim_in, config.hidden_size, bias=not config.no_bias)483 484 def forward(self, hidden_states):485 a1 = self.w1(hidden_states)486 a2 = self.w2(hidden_states)487 intermediate_parallel = a1 * F.silu(a2)488 output = self.c_proj(intermediate_parallel)489 return output490 491 492class QWenBlock(nn.Module):493 def __init__(self, config, layer_idx=None, num_expert=1):494 super().__init__()495 self.num_expert = num_expert496 self.layer_number = layer_idx497 self.apply_residual_connection_post_layernorm = (498 config.apply_residual_connection_post_layernorm499 )500 hidden_size = config.hidden_size501 self.apply_residual_connection_post_layernorm = (502 config.apply_residual_connection_post_layernorm503 )504 self.bf16 = config.bf16505 506 self.ln_1 = RMSNorm(507 hidden_size,508 eps=config.layer_norm_epsilon,509 )510 self.attn = QWenAttention(config, layer_number=layer_idx)511 self.ln_2 = RMSNorm(512 hidden_size,513 eps=config.layer_norm_epsilon,514 )515 516 self.mlp = QWenMLP(config)517 518 def forward(519 self,520 hidden_states: Optional[Tuple[torch.FloatTensor]],521 layer_past: Optional[Tuple[torch.Tensor]] = None,522 attention_mask: Optional[torch.FloatTensor] = None,523 head_mask: Optional[torch.FloatTensor] = None,524 encoder_hidden_states: Optional[torch.Tensor] = None,525 encoder_attention_mask: Optional[torch.FloatTensor] = None,526 use_cache: Optional[bool] = False,527 output_attentions: Optional[bool] = False,528 ):529 layernorm_output = self.ln_1(hidden_states)530 531 attn_outputs = self.attn(532 layernorm_output,533 layer_past=layer_past,534 attention_mask=attention_mask,535 head_mask=head_mask,536 use_cache=use_cache,537 output_attentions=output_attentions,538 )539 attn_output = attn_outputs[0]540 541 outputs = attn_outputs[1:]542 543 if self.apply_residual_connection_post_layernorm:544 residual = layernorm_output545 else:546 residual = hidden_states547 layernorm_input = attn_output + residual548 549 layernorm_output = self.ln_2(layernorm_input)550 551 if self.apply_residual_connection_post_layernorm:552 residual = layernorm_output553 else:554 residual = layernorm_input555 556 mlp_output = self.mlp(layernorm_output)557 hidden_states = residual + mlp_output558 559 if use_cache:560 outputs = (hidden_states,) + outputs561 else:562 outputs = (hidden_states,) + outputs[1:]563 564 return outputs565 566 567class QWenPreTrainedModel(PreTrainedModel):568 config_class = QWenConfig569 base_model_prefix = "transformer"570 is_parallelizable = False571 supports_gradient_checkpointing = True572 _no_split_modules = ["QWenBlock"]573 574 def __init__(self, *inputs, **kwargs):575 super().__init__(*inputs, **kwargs)576 577 def _init_weights(self, module):578 """Initialize the weights."""579 if isinstance(module, nn.Linear):580 module.weight.data.normal_(mean=0.0, std=self.config.initializer_range)581 if module.bias is not None:582 module.bias.data.zero_()583 elif isinstance(module, nn.Embedding):584 module.weight.data.normal_(mean=0.0, std=self.config.initializer_range)585 if module.padding_idx is not None:586 module.weight.data[module.padding_idx].zero_()587 elif isinstance(module, RMSNorm):588 module.weight.data.fill_(1.0)589 590 for name, p in module.named_parameters():591 if name == "c_proj.weight":592 p.data.normal_(593 mean=0.0,594 std=(595 self.config.initializer_range596 / math.sqrt(2 * self.config.n_layer)597 ),598 )599 600 def _set_gradient_checkpointing(self, module, value=False):601 if isinstance(module, QWenModel):602 module.gradient_checkpointing = value603 604 605class QWenModel(QWenPreTrainedModel):606 _keys_to_ignore_on_load_missing = ["attn.masked_bias"]607 608 def __init__(self, config):609 super().__init__(config)610 self.vocab_size = config.padded_vocab_size611 self.num_hidden_layers = config.num_hidden_layers612 self.embed_dim = config.hidden_size613 614 max_sequence_length = config.max_position_embeddings615 self.position_embedding_type = config.pos_emb616 self.gradient_checkpointing = False617 618 if self.position_embedding_type == "learned":619 self.wpe = nn.Embedding(max_sequence_length, self.embed_dim)620 self.init_method(self.position_embeddings.weight)621 self._position_embeddings_key = "position_embeddings"622 self.init_method(self.position_embeddings.weight)623 else:624 self.wpe = None625 self._position_embeddings_key = ""626 627 self.wte = nn.Embedding(self.vocab_size, self.embed_dim)628 629 self.drop = nn.Dropout(config.embd_pdrop)630 self.h = nn.ModuleList(631 [632 QWenBlock(633 config,634 layer_idx=i,635 )636 for i in range(config.num_hidden_layers)637 ]638 )639 self.ln_f = RMSNorm(640 self.embed_dim,641 eps=config.layer_norm_epsilon,642 )643 644 self.post_init()645 646 def get_input_embeddings(self):647 return self.wte648 649 def set_input_embeddings(self, new_embeddings):650 self.wte = new_embeddings651 652 def forward(653 self,654 input_ids: Optional[torch.LongTensor] = None,655 past_key_values: Optional[Tuple[Tuple[torch.Tensor]]] = None,656 attention_mask: Optional[torch.FloatTensor] = None,657 token_type_ids: Optional[torch.LongTensor] = None,658 position_ids: Optional[torch.LongTensor] = None,659 head_mask: Optional[torch.FloatTensor] = None,660 inputs_embeds: Optional[torch.FloatTensor] = None,661 encoder_hidden_states: Optional[torch.Tensor] = None,662 encoder_attention_mask: Optional[torch.FloatTensor] = None,663 use_cache: Optional[bool] = None,664 output_attentions: Optional[bool] = None,665 output_hidden_states: Optional[bool] = None,666 return_dict: Optional[bool] = None,667 ):668 output_attentions = (669 output_attentions670 if output_attentions is not None671 else self.config.output_attentions672 )673 output_hidden_states = (674 output_hidden_states675 if output_hidden_states is not None676 else self.config.output_hidden_states677 )678 use_cache = use_cache if use_cache is not None else self.config.use_cache679 return_dict = (680 return_dict if return_dict is not None else self.config.use_return_dict681 )682 683 if input_ids is not None and inputs_embeds is not None:684 raise ValueError(685 "You cannot specify both input_ids and inputs_embeds at the same time"686 )687 elif input_ids is not None:688 input_shape = input_ids.size()689 input_ids = input_ids.view(-1, input_shape[-1])690 batch_size = input_ids.shape[0]691 elif inputs_embeds is not None:692 input_shape = inputs_embeds.size()[:-1]693 batch_size = inputs_embeds.shape[0]694 else:695 raise ValueError("You have to specify either input_ids or inputs_embeds")696 697 device = input_ids.device if input_ids is not None else inputs_embeds.device698 699 if token_type_ids is not None:700 token_type_ids = token_type_ids.view(-1, input_shape[-1])701 if position_ids is not None:702 position_ids = position_ids.view(-1, input_shape[-1])703 704 if past_key_values is None:705 past_length = 0706 past_key_values = tuple([None] * len(self.h))707 else:708 past_length = past_key_values[0][0].size(-2)709 710 if position_ids is None:711 position_ids = torch.arange(712 past_length,713 input_shape[-1] + past_length,714 dtype=torch.long,715 device=device,716 )717 position_ids = position_ids.unsqueeze(0).view(-1, input_shape[-1])718 719 if attention_mask is not None:720 if batch_size <= 0:721 raise ValueError("batch_size has to be defined and > 0")722 attention_mask = attention_mask.view(batch_size, -1)723 attention_mask = attention_mask[:, None, None, :]724 attention_mask = attention_mask.to(dtype=self.dtype)725 attention_mask = (1.0 - attention_mask) * torch.finfo(self.dtype).min726 # attention_mask中mask掉的部分是-inf, 看到的部分是0727 728 encoder_attention_mask = None729 head_mask = self.get_head_mask(head_mask, self.config.n_layer)730 731 if inputs_embeds is None:732 inputs_embeds = self.wte(input_ids)733 hidden_states = inputs_embeds734 if self.wpe is not None:735 position_embeds = self.wpe(position_ids)736 hidden_states = hidden_states + position_embeds737 738 hidden_states = self.drop(hidden_states)739 output_shape = input_shape + (hidden_states.size(-1),)740 741 if self.gradient_checkpointing and self.training:742 if use_cache:743 logger.warning_once(744 "`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`..."745 )746 use_cache = False747 748 presents = () if use_cache else None749 all_self_attentions = () if output_attentions else None750 all_hidden_states = () if output_hidden_states else None751 for i, (block, layer_past) in enumerate(zip(self.h, past_key_values)):752 753 if output_hidden_states:754 all_hidden_states = all_hidden_states + (hidden_states,)755 756 if self.gradient_checkpointing and self.training:757 758 def create_custom_forward(module):759 def custom_forward(*inputs):760 # None for past_key_value761 return module(*inputs, use_cache, output_attentions)762 763 return custom_forward764 765 outputs = torch.utils.checkpoint.checkpoint(766 create_custom_forward(block),767 hidden_states,768 None,769 attention_mask,770 head_mask[i],771 encoder_hidden_states,772 encoder_attention_mask,773 )774 else:775 outputs = block(776 hidden_states,777 layer_past=layer_past,778 attention_mask=attention_mask,779 head_mask=head_mask[i],780 encoder_hidden_states=encoder_hidden_states,781 encoder_attention_mask=encoder_attention_mask,782 use_cache=use_cache,783 output_attentions=output_attentions,784 )785 786 hidden_states = outputs[0]787 if use_cache is True:788 presents = presents + (outputs[2 if output_attentions else 1],)789 790 if output_attentions:791 all_self_attentions = all_self_attentions + (outputs[1],)792 793 hidden_states = self.ln_f(hidden_states)794 hidden_states = hidden_states.view(output_shape)795 796 if not return_dict:797 return tuple(798 v for v in [hidden_states, presents, all_hidden_states] if v is not None799 )800 801 return BaseModelOutputWithPast(802 last_hidden_state=hidden_states,803 past_key_values=presents,804 hidden_states=all_hidden_states,805 attentions=all_self_attentions,806 )807 808 809class QWenLMHeadModel(QWenPreTrainedModel):810 _keys_to_ignore_on_load_missing = [r"h\.\d+\.attn\.rotary_emb\.inv_freq"]811 _keys_to_ignore_on_load_unexpected = [r"h\.\d+\.attn\.masked_bias"]812 813 def __init__(self, config):814 super().__init__(config)815 assert (816 config.bf16 + config.fp16 + config.fp32 <= 1817 ), "Only one of \"bf16\", \"fp16\", \"fp32\" can be true"818 819 autoset_precision = config.bf16 + config.fp16 + config.fp32 == 0820 821 if autoset_precision:822 if SUPPORT_BF16:823 logger.warn(824 "The model is automatically converting to bf16 for faster inference. "825 "If you want to disable the automatic precision, please manually add bf16/fp16/fp32=True to \"AutoModelForCausalLM.from_pretrained\"."826 )827 config.bf16 = True828 elif SUPPORT_FP16:829 logger.warn(830 "The model is automatically converting to fp16 for faster inference. "831 "If you want to disable the automatic precision, please manually add bf16/fp16/fp32=True to \"AutoModelForCausalLM.from_pretrained\"."832 )833 config.fp16 = True834 else:835 config.fp32 = True836 837 if config.bf16 and SUPPORT_CUDA and not SUPPORT_BF16:838 logger.warn("Your device does NOT seem to support bf16, you can switch to fp16 or fp32 by by passing fp16/fp32=True in \"AutoModelForCausalLM.from_pretrained\".")839 if config.fp16 and SUPPORT_CUDA and not SUPPORT_FP16:840 logger.warn("Your device does NOT support faster inference with fp16, please switch to fp32 which is likely to be faster")841 if config.fp32:842 if SUPPORT_BF16:843 logger.warn("Your device support faster inference by passing bf16=True in \"AutoModelForCausalLM.from_pretrained\".")844 elif SUPPORT_FP16:845 logger.warn("Your device support faster inference by passing fp16=True in \"AutoModelForCausalLM.from_pretrained\".")846 847 if config.use_flash_attn == "auto":848 if config.bf16 or config.fp16:849 logger.warn("Try importing flash-attention for faster inference...")850 config.use_flash_attn = True851 else:852 config.use_flash_attn = False853 if config.use_flash_attn and config.fp32:854 logger.warn("Flash attention will be disabled because it does NOT support fp32.")855 856 if config.use_flash_attn:857 _import_flash_attn()858 859 self.transformer = QWenModel(config)860 self.lm_head = nn.Linear(config.n_embd, config.vocab_size, bias=False)861 862 if config.bf16:863 self.transformer.bfloat16()864 self.lm_head.bfloat16()865 if config.fp16:866 self.transformer.half()867 self.lm_head.half()868 self.post_init()869 870 def get_output_embeddings(self):871 return self.lm_head872 873 def set_output_embeddings(self, new_embeddings):874 self.lm_head = new_embeddings875 876 def prepare_inputs_for_generation(877 self, input_ids, past_key_values=None, inputs_embeds=None, **kwargs878 ):879 token_type_ids = kwargs.get("token_type_ids", None)880 if past_key_values:881 input_ids = input_ids[:, -1].unsqueeze(-1)882 if token_type_ids is not None:883 token_type_ids = token_type_ids[:, -1].unsqueeze(-1)884 885 attention_mask = kwargs.get("attention_mask", None)886 position_ids = kwargs.get("position_ids", None)887 888 if attention_mask is not None and position_ids is None:889 position_ids = attention_mask.long().cumsum(-1) - 1890 position_ids.masked_fill_(attention_mask == 0, 1)891 if past_key_values:892 position_ids = position_ids[:, -1].unsqueeze(-1)893 else:894 position_ids = None895 896 if inputs_embeds is not None and past_key_values is None:897 model_inputs = {"inputs_embeds": inputs_embeds}898 else:899 model_inputs = {"input_ids": input_ids}900 901 model_inputs.update(902 {903 "past_key_values": past_key_values,904 "use_cache": kwargs.get("use_cache"),905 "position_ids": position_ids,906 "attention_mask": attention_mask,907 "token_type_ids": token_type_ids,908 }909 )910 return model_inputs911 912 def forward(913 self,914 input_ids: Optional[torch.LongTensor] = None,915 past_key_values: Optional[Tuple[Tuple[torch.Tensor]]] = None,916 attention_mask: Optional[torch.FloatTensor] = None,917 token_type_ids: Optional[torch.LongTensor] = None,918 position_ids: Optional[torch.LongTensor] = None,919 head_mask: Optional[torch.FloatTensor] = None,920 inputs_embeds: Optional[torch.FloatTensor] = None,921 encoder_hidden_states: Optional[torch.Tensor] = None,922 encoder_attention_mask: Optional[torch.FloatTensor] = None,923 labels: Optional[torch.LongTensor] = None,924 use_cache: Optional[bool] = None,925 output_attentions: Optional[bool] = None,926 output_hidden_states: Optional[bool] = None,927 return_dict: Optional[bool] = None,928 ) -> Union[Tuple, CausalLMOutputWithPast]:929 930 return_dict = (931 return_dict if return_dict is not None else self.config.use_return_dict932 )933 934 transformer_outputs = self.transformer(935 input_ids,936 past_key_values=past_key_values,937 attention_mask=attention_mask,938 token_type_ids=token_type_ids,939 position_ids=position_ids,940 head_mask=head_mask,941 inputs_embeds=inputs_embeds,942 encoder_hidden_states=encoder_hidden_states,943 encoder_attention_mask=encoder_attention_mask,944 use_cache=use_cache,945 output_attentions=output_attentions,946 output_hidden_states=output_hidden_states,947 return_dict=return_dict,948 )949 hidden_states = transformer_outputs[0]950 951 lm_logits = self.lm_head(hidden_states)952 953 loss = None954 if labels is not None:955 labels = labels.to(lm_logits.device)956 shift_logits = lm_logits[..., :-1, :].contiguous()957 shift_labels = labels[..., 1:].contiguous()958 loss_fct = CrossEntropyLoss()959 loss = loss_fct(960 shift_logits.view(-1, shift_logits.size(-1)), shift_labels.view(-1)961 )962 963 if not return_dict:964 output = (lm_logits,) + transformer_outputs[1:]965 return ((loss,) + output) if loss is not None else output966 967 return CausalLMOutputWithPast(968 loss=loss,969 logits=lm_logits,970 past_key_values=transformer_outputs.past_key_values,971 hidden_states=transformer_outputs.hidden_states,972 attentions=transformer_outputs.attentions,973 )974 975 @staticmethod976 def _reorder_cache(977 past_key_values: Tuple[Tuple[torch.Tensor]], beam_idx: torch.Tensor978 ) -> Tuple[Tuple[torch.Tensor]]:979 980 return tuple(981 tuple(982 past_state.index_select(0, beam_idx.to(past_state.device))983 for past_state in layer_past984 )985 for layer_past in past_key_values986 )987 988 def chat(989 self,990 tokenizer: PreTrainedTokenizer,991 query: str,992 history: Optional[HistoryType],993 system: str = "You are a helpful assistant.",994 append_history: bool = True,995 stream: Optional[bool] = _SENTINEL,996 stop_words_ids: Optional[List[List[int]]] = None,997 **kwargs,998 ) -> Tuple[str, HistoryType]:999 assert stream is _SENTINEL, _ERROR_STREAM_IN_CHAT1000 assert self.generation_config.chat_format == 'chatml', _ERROR_BAD_CHAT_FORMAT1001 if history is None:1002 history = []1003 if stop_words_ids is None:1004 stop_words_ids = []1005 1006 raw_text, context_tokens = make_context(1007 tokenizer,1008 query,1009 history=history,1010 system=system,1011 max_window_size=6144,1012 chat_format=self.generation_config.chat_format,1013 )1014 1015 stop_words_ids.extend(get_stop_words_ids(1016 self.generation_config.chat_format, tokenizer1017 ))1018 input_ids = torch.tensor([context_tokens]).to(self.device)1019 outputs = self.generate(1020 input_ids,1021 stop_words_ids = stop_words_ids,1022 return_dict_in_generate = False,1023 **kwargs,1024 )1025 1026 response = decode_tokens(1027 outputs[0],1028 tokenizer,1029 raw_text_len=len(raw_text),1030 context_length=len(context_tokens),1031 chat_format=self.generation_config.chat_format,1032 verbose=False,1033 errors='replace'1034 )1035 1036 if append_history:1037 history.append((query, response))1038 1039 return response, history1040 1041 def chat_stream(1042 self,1043 tokenizer: PreTrainedTokenizer,1044 query: str,1045 history: Optional[HistoryType],1046 system: str = "You are a helpful assistant.",1047 stop_words_ids: Optional[List[List[int]]] = None,1048 logits_processor: Optional[LogitsProcessorList] = None,1049 **kwargs,1050 ) -> Generator[str, Any, None]:1051 assert self.generation_config.chat_format == 'chatml', _ERROR_BAD_CHAT_FORMAT1052 if history is None:1053 history = []1054 if stop_words_ids is None:1055 stop_words_ids = []1056 1057 raw_text, context_tokens = make_context(1058 tokenizer,1059 query,1060 history=history,1061 system=system,1062 max_window_size=6144,1063 chat_format=self.generation_config.chat_format,1064 )1065 1066 stop_words_ids.extend(get_stop_words_ids(1067 self.generation_config.chat_format, tokenizer1068 ))1069 if stop_words_ids is not None:1070 stop_words_logits_processor = StopWordsLogitsProcessor(1071 stop_words_ids=stop_words_ids,1072 eos_token_id=self.generation_config.eos_token_id,1073 )1074 if logits_processor is None:1075 logits_processor = LogitsProcessorList([stop_words_logits_processor])1076 else:1077 logits_processor.append(stop_words_logits_processor)1078 input_ids = torch.tensor([context_tokens]).to(self.device)1079 1080 from transformers_stream_generator.main import NewGenerationMixin, StreamGenerationConfig1081 self.__class__.generate_stream = NewGenerationMixin.generate1082 self.__class__.sample_stream = NewGenerationMixin.sample_stream1083 stream_config = StreamGenerationConfig(**self.generation_config.to_dict(), do_stream=True)1084 def stream_generator():1085 outputs = []1086 for token in self.generate_stream(1087 input_ids,1088 return_dict_in_generate=False,1089 generation_config=stream_config,1090 logits_processor=logits_processor,1091 seed=-1,1092 **kwargs):1093 outputs.append(token.item())1094 yield tokenizer.decode(outputs, skip_special_tokens=True, errors='ignore')1095 1096 return stream_generator()1097 1098 def generate(1099 self,1100 inputs: Optional[torch.Tensor] = None,1101 generation_config: Optional[GenerationConfig] = None,1102 logits_processor: Optional[LogitsProcessorList] = None,1103 stopping_criteria: Optional[StoppingCriteriaList] = None,1104 prefix_allowed_tokens_fn: Optional[1105 Callable[[int, torch.Tensor], List[int]]1106 ] = None,1107 synced_gpus: Optional[bool] = None,1108 assistant_model: Optional["PreTrainedModel"] = None,1109 streamer: Optional["BaseStreamer"] = None,1110 **kwargs,1111 ) -> Union[GenerateOutput, torch.LongTensor]:1112 # Process stop_words_ids.1113 stop_words_ids = kwargs.pop("stop_words_ids", None)1114 if stop_words_ids is None and generation_config is not None:1115 stop_words_ids = getattr(generation_config, "stop_words_ids", None)1116 if stop_words_ids is None:1117 stop_words_ids = getattr(self.generation_config, "stop_words_ids", None)1118 1119 if stop_words_ids is not None:1120 stop_words_logits_processor = StopWordsLogitsProcessor(1121 stop_words_ids=stop_words_ids,1122 eos_token_id=self.generation_config.eos_token_id,1123 )1124 if logits_processor is None:1125 logits_processor = LogitsProcessorList([stop_words_logits_processor])1126 else:1127 logits_processor.append(stop_words_logits_processor)1128 1129 return super().generate(1130 inputs,1131 generation_config=generation_config,1132 logits_processor=logits_processor,1133 stopping_criteria=stopping_criteria,1134 prefix_allowed_tokens_fn=prefix_allowed_tokens_fn,1135 synced_gpus=synced_gpus,1136 assistant_model=assistant_model,1137 streamer=streamer,1138 **kwargs,1139 )1140 1141 1142class RotaryEmbedding(torch.nn.Module):1143 def __init__(self, dim, base=10000):1144 super().__init__()1145 self.dim = dim1146 self.base = base1147 self.inv_freq = 1.0 / (base ** (torch.arange(0, dim, 2).float() / dim))1148 if importlib.util.find_spec("einops") is None:1149 raise RuntimeError("einops is required for Rotary Embedding")1150 1151 self._rotary_pos_emb_cache = None1152 self._seq_len_cached = 01153 self._ntk_alpha_cached = 1.01154 1155 def update_rotary_pos_emb_cache(self, max_seq_len, offset=0, ntk_alpha=1.0):1156 seqlen = max_seq_len + offset1157 if seqlen > self._seq_len_cached or ntk_alpha != self._ntk_alpha_cached:1158 base = self.base * ntk_alpha ** (self.dim / (self.dim - 2))1159 self.inv_freq = 1.0 / (1160 base1161 ** (1162 torch.arange(0, self.dim, 2, device=self.inv_freq.device).float()1163 / self.dim1164 )1165 )1166 self._seq_len_cached = max(2 * seqlen, 16)1167 self._ntk_alpha_cached = ntk_alpha1168 seq = torch.arange(self._seq_len_cached, device=self.inv_freq.device)1169 freqs = torch.outer(seq.type_as(self.inv_freq), self.inv_freq)1170 emb = torch.cat((freqs, freqs), dim=-1)1171 from einops import rearrange1172 1173 self._rotary_pos_emb_cache = rearrange(emb, "n d -> 1 n 1 d")1174 1175 def forward(self, max_seq_len, offset=0, ntk_alpha=1.0):1176 self.update_rotary_pos_emb_cache(max_seq_len, offset, ntk_alpha)1177 return self._rotary_pos_emb_cache[:, offset : offset + max_seq_len]1178 1179 1180def _rotate_half(x):1181 from einops import rearrange1182 1183 x = rearrange(x, "... (j d) -> ... j d", j=2)1184 x1, x2 = x.unbind(dim=-2)1185 return torch.cat((-x2, x1), dim=-1)1186 1187 1188def apply_rotary_pos_emb(t, freqs):1189 if apply_rotary_emb_func is not None:1190 t_ = t.float()1191 freqs = freqs.squeeze(0).squeeze(1)1192 cos = freqs[:, : freqs.shape[-1] // 2].cos()1193 sin = freqs[:, : freqs.shape[-1] // 2].sin()1194 output = apply_rotary_emb_func(t_, cos, sin).type_as(t)1195 return output1196 else:1197 rot_dim = freqs.shape[-1]1198 t_, t_pass_ = t[..., :rot_dim], t[..., rot_dim:]1199 t_ = t_.float()1200 t_pass_ = t_pass_.float()