Team Ai
Modelpublic

codefuse-ai/CodeFuse-DevOps-Model-7B-Base

sourceHugging Faceotherupdated 3y agoView on Hugging Face
1likes18downloads
modeling_qwen.py1220 linesDownload Raw Back to root
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()

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