Team Ai
Apppublic

modelscope/DiffSynth-Painter

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
14likes
kolors_text_encoder.py1553 linesDownload Raw Back to models
1"""2This model is copied from https://github.com/Kwai-Kolors/Kolors/tree/master/kolors/models.3We didn't modify this model.4The tensor operation is performed in the prompter.5"""6 7 8""" PyTorch ChatGLM model. """9 10import math11import copy12import warnings13import re14import sys15 16import torch17import torch.utils.checkpoint18import torch.nn.functional as F19from torch import nn20from torch.nn import CrossEntropyLoss, LayerNorm21from torch.nn import CrossEntropyLoss, LayerNorm, MSELoss, BCEWithLogitsLoss22from torch.nn.utils import skip_init23from typing import Optional, Tuple, Union, List, Callable, Dict, Any24from copy import deepcopy25 26from transformers.modeling_outputs import (27    BaseModelOutputWithPast,28    CausalLMOutputWithPast,29    SequenceClassifierOutputWithPast,30)31from transformers.modeling_utils import PreTrainedModel32from transformers.utils import logging33from transformers.generation.logits_process import LogitsProcessor34from transformers.generation.utils import LogitsProcessorList, StoppingCriteriaList, GenerationConfig, ModelOutput35from transformers import PretrainedConfig36from torch.nn.parameter import Parameter37import bz238import torch39import base6440import ctypes41from transformers.utils import logging42from typing import List43 44 45 46logger = logging.get_logger(__name__)47 48try:49    from cpm_kernels.kernels.base import LazyKernelCModule, KernelFunction, round_up50 51 52    class Kernel:53        def __init__(self, code: bytes, function_names: List[str]):54            self.code = code55            self._function_names = function_names56            self._cmodule = LazyKernelCModule(self.code)57 58            for name in self._function_names:59                setattr(self, name, KernelFunction(self._cmodule, name))60 61 62    quantization_code = "$QlpoOTFBWSZTWU9yuJUAQHN//////////f/n/8/n///n//bt4dTidcVx8X3V9FV/92/v4B7/AD5FBQFAAAChSgKpFCFAFVSigUAAAEKhSgUUqgFBKigqVREQAABQBQIANDTTIGI00BkZBkNGE0A0BkBkGQGRkaNAaAGQNBoGgDIAAYIGTI0DQAQAaGmmQMRpoDIyDIaMJoBoDIDIMgMjI0aA0AMgaDQNAGQAAwQMmRoGgAgA0NNMgYjTQGRkGQ0YTQDQGQGQZAZGRo0BoAZA0GgaAMgABggZMjQNABABoaaZAxGmgMjIMhowmgGgMgMgyAyMjRoDQAyBoNA0AZAADBAyZGgaAAmqU1NEgJqnptU/Sn4jRR6J6epk2pqb1Q/SgAPUGgyNNGjQ2SBpoAZAAGg0NB6mgDIAAAAA2oaApSREBNAARhGiYEaEwU8pvImlP0k2aam1GaGqbFNM1MHpTwmkepmyU9R6nqPKekHqNNPUxNGhp6n6p6QaZ6o9TG1GMqcoV9ly6nRanHlq6zPNbnGZNi6HSug+2nPiZ13XcnFYZW+45W11CumhzYhchOJ2GLLV1OBjBjGf4TptOddTSOcVxhqYZMYwZXZZY00zI1paX5X9J+b+f4e+x43RXSxXPOdquiGpduatGyXneN696M9t4HU2eR5XX/kPhP261NTx3JO1Ow7LyuDmeo9a7d351T1ZxnvnrvYnrXv/hXxPCeuYx2XsNmO003eg9J3Z6U7b23meJ4ri01OdzTk9BNO96brz+qT5nuvvH3ds/G+m/JcG/F2XYuhXlvO+jP7U3XgrzPN/lr8Sf1n6j4j7jZs+s/T0tNaNNYzTs12rxjwztHlnire3Nzc3N1wuBwOBwXBvZfoHpD7rFmR99V5vj3aXza3xdBbXMalubTg/jIv5dfAi54Pdc75j4z412n3Npj3Ld/ENm7a3b/Cod6h/ret1/5vn/C+l+gdslMvgPSLJ8d8q+U66fevYn/tW1chleEtNTGlcHCbLRlq0tHzF5tsbbZZfHjjLgZu42XCuC3NrdjTasZGNzgxPIrGqp7r3p7L2p5XjnpPSmTd5XtzqnB6U87zzg1Ol0zd0zsLszxR6lkxp35u6/teL0L0W922cR7Lu1lpL9CsHirzuM2T+BgsyViT6LHcm0/Vr6U/7LGGyJeqTEjt0PHWhF5mCT7R9mtlDwriYv0Tyr/OxYt6qp5r0mPVT0608TqnqMZaarU2nFwrTzzlrs1ed7z1ux60wyr4ydCaTi3enW8x68x0zU7tXSlcmPSW1mGpWJMg4zmPC2lK96tp0OE80y4MfEvnZj8zGluR6b22ki1Ou9V2nCd9xovcPvcYMZYy0lvN60ScZ45vN6yeCeeXFb1lVjnnCar5fwXwE2bzJ4HI1XVPXfXZMm44GUsMpYsmLB65TuVdm0cl0b+i/wGNN66XjeV7zuPpHcnK/juhhjdfId5jMdE5nN0dGmmm2zZs2cexD5n9p/dY352XsvXHaZNWWsmmS1atjR452nYudzvqv2HMRyvNNnlMcDl3R2+yx2uVrBubTW9icHDVtbNXlZm7jma1rM4VurZZd2y6nUau7ZXZ7bVU+mnoOVxZGMrVmvX60605JwmzGZhhhjTWtaaaMaaGTGmNMZasY0iX8VMUl8eepaIrzGSpemWOQyZORk2bNpjUybMmxqYmknCGCFynutfksaZpjTNMaaatM0xsxcGR0sociNqxNSmhhR1ZJPbsn8qyF0t2qH6iYBclclalbtTTcHTDsPaX6rlnElph2Jyumumtynv2Kk8GI7rsvXbIcJgHJOSaSXnnGaI3m87RtVXJOZ/YtgdTE6Wpha6ZlE8ayXkef1fh602r2WwvfMXtMdLlkfnLFdYYwYso+bWqm7yJqHXZGw2nrS5ZanSYnWlxBxMF1V940K2wdrI7R6OYf7DGGamMmTSbRhlS45xmVOumF1EyPCmHrrN8wwZOOrdNtLeMtzFzDlWnfTBxMk2NaXIZHBYxYLD4w8yju0ao65Vz1OIXoS9dLanwCe1PWrYuWMqf1if1z2k2yYfKJ741PDgno1ZQ8DRqvUny3mNoWTzGO6m1DkrJI8JiR5cSd+vZdGOO8nrMoc5+NDUFsMSXaZJeNlMmGLtJsovOsUp7I9S5VojKxF6bTVEelXqlfJobQr3LozSh2Jk7VcrVMfhXqszGWMzNqGhqZY0OadxkyyMssKugZR0KNFXBHlqwmJgTE/BNVMk6ItJXZMR0H47GpXv/DMOvNkmVuaV1PRfEdxuqc7Hcd+ZV/zTLaRxWk0nl9CdCeM6mn5rstHIBcpiuwmUZXeq81DacHI2rmrZ5SuE5mOZd6LQrZg9mx32TprA8BMo5jKN6yLTCi3WzQaZSuhzTtM1fUTGVpG8Tw+KXI0tjEpiWxtLYynOlktSbVlaI5kxP8TDH8kx50xoxi5KcA4pcja8KWLRlO/Ks6q06ergnvm1ca3Tq8Uw7LTUsmWyctXPWmpitl/uvGcWTGXGuAXDfhqazGmjkxcJW5hMMMMpYsXl2TZYtVOddG3XCarUt6Ptq9CZXSNzyuRzqRZOjsxdBbFVz6OA5HI43r1jityVlVpVkxmOsyaYWE1NTGq1sOVh36mHMcxtSvcy70edG0ZGR3I1Go1GRlV7mWWo1G0ZGRqlvH40l7o4m5xMWLLLYyNjnqc8556mdPqLJ31n/1nWOncxzG1tizrHs/Z+d2vP/B/l8wdJ6rHUn2nbbDq4p6htFtYzMMMTaZis1K5GKzGNmxhmUx2DDlZ/qNnIx41xnaMfCZWYaZWtNLTNW8ND4Fw1MyZOCdM428suKG1ehW8TesOydg7J+YYcD4cYR+8dFK6M4E3HM9ZfRNNL+Sn6rsl4DsrDl2HpPCnfxjGXtbZtYys1ttlyJ4T+BvexjGWRjMszK4Jpc77D3GyuVD7q0+G8m9G+2+rGm7cOR2y7FdtY2XUYx/oNlfRYxhMYyYZkyyg55enna9Kt/FFi6GMMwYwdwxWgxGMLKYmUyGExTKMZkMFhkymKuh0NOBNnBu+23LdwDoZYYzGGMxtORaTU1pjTGWTTGGtMrNWUsyyTTLLG1qy2ZjbK2DBllWqxMtBMaYZQmcE7zvvRcTkclUwdkxTaSdyySt/7fpL+T1v516Ji97fwr5JbLu305zMn5+GMTTZ9F+y7ExwmGVfG44yxn3dLv6l5i+Wth1jCrDq21nW9LqvvDzz3Vf3LLH/O/32TJ/erx3bXftO4eF+G956D952K/An4NfvOpjFjExjevP/UmE0fIoZXx6/w6lX/no3D0bLt+ixjieBM6ksRd0yB4Lt2SwYNE+gd1detlZWUnpiZfGfFaK+4PyCa/v18V8X75pe9fLXzp7l3VjF76vWZmHwGz1IZNWT7b8yddJ4q5kyrVdfru6atWc7bVYztL9Jf4GXvT+Y8m9/YsXP6H018a8D4XVOqvfzqeR+6yZOD8dPv0+U7/q5Pl+2dNb0MjzGVH5p6MNQ7cOWvw62U9aHE8DprDek+McLyvDz+te+9Zhq5+YTruufMcWMabqysTmZVWjKPfnK0wyVcrsuhjZRdLkHNvD72b9abriOSGIxiLixMOoalNPXzy+wT/tf+U6HHONfsz+xe8ufHBdQWWGWLA9if0rsnmrxK5LvRZQeWsTCsrmOYy8VteVfuRfcVTtDLItLIsMYxZLdU/DbtSemxF6Z6Zo5WBXE4tFdCyVMMXMTEMZXVlS6Xec2T4e0tHsRcEuWshcJ2YsNF5rUx1E8ifCq6Z+ZP7qdCeu/aTwFd53l16/o0NOw6O3dLavP4Hbi4RdmuDk6DoYaninC0+o4uZjbJ7Rxeu0/FbuFg+q7DVS6fQe0rZ6NDGUNNU6DEqOaLTicKnYZMnBWruljQxoaS3dZhocDge0bSTyOvdAbG5hxe2xji7E/L55xX13wWNDi6HCekcFxfCPGxY0MXC+s7afWaMdDyjyr+o8Rudm/NabOZvdl274zH4f5XK9z6On1Pe/K5TdPAslg77BjuO6Y3eO7GqvOPG/stknp1leyvLL0Z7bl9I4noMvLkzytLhWYzrOZzLXCORe028rORzOg4N/L0HlMOQ3Pgmnbb6KczlabORpu980q37TBqRu0/p3PO6234Bl03Ynuz+9W7gnsEcmvYaYY3aMYY0wx3pYd+ujsXauWdaY5Xkbtl23fPzFHiDB/QMo0yFjBllYxTQYYyxkrwn7JufwJ/PfgJ+C83X69ni6zvXcnyXabv0ncbLwsceS+RNlyN2mnneJtX0ngYO0+e+0+UnA+Wch3ji8hj5an4h+i6XBySU4n+R0roVcbw5yvHrmr4Yw8Y7x6c+9POPYHI5HI5HI5HI5HGXGww4nE4nrVyOR8XeqPEO7PLOiukYa3Novk5hV4cdtYZLI93e+uxff2jRo0aNGjRo0aNG1bVtW1dy3m83m8+tQ5ZzHw3nObwOu8La9Rc1dtkdS8A3eTk823tnktXWlxN6Oixe06zrN70Isd9jiOgZFq9yfkPqP/SLhN2Myl8jDM43bl1nbcb4cO57jlh8Jow6pzXZdL4dyODTuuhu77FyO27DdwdRxmvO+O+3N2+BdqyTwLHVczDVY4UPE4O66/ZO2cx1LFzVdSXtF7G4HMbrauOHRw6c8FdZ5m9fHZHYZXfTlZquyynSyTTKke6vcffSD9pzPA/G7n7jxPmuhc1DHMynPMrGL6AdewYmwu5ko+UUyTwrMv27rPH1v1nGqd87+p6N6LU8k3NEng53xXyHS97+44OSg/sy/hn+Se6yfYNjW0/uTgP+PvWYzLMmjhcLB/gGpri6H83/84eUXWT6T9Hsv7785z/7z4icpW+zfXypuR7rx/gMdZb1/wC678pcs8/2a3mDitGHxl9mfPlll5MafWWqxk/eYuTDgcNMzDGWLWvsuglNxs53GtN6uWpktlW1tZZYcuinMMWmnNnJydze3b2Y1McBxrBkXw799izLMZZYyy0TkbsGM4p03S2uVu5s/XXUdSdec6smVxZYYGpVmT8A+8ajuEyV5FatkvVru2x6uxGXXbH4A+jvgP4GMYy3iPLXzq/6z65+E005ey+cwMZD3fZcqc6xpjTFjQ0P3U+e++cPYmTIwj0nrK5NPTfl3WvpfLtXDcb2HQMudYOxFXQBor4L4T6vrOauFctYXJQ++NUWmJe5bmx1jDiZS1dTqWxo4GR8jm3fttpmPHppk9PEyv4/y8/sO07XacOmcqc0x2Vi9BvNJvN5oW8x4mOsydpidRxMYJPx06m1bqPzq9KtK8sxXNXFodD/+MYYaJTLwOhc9brCsV18oOR1i4tXChyTkq4lf4y1Ke+9axjDHqs1mfBbMXuP4Hzi+X7t8vzv7bHerrUPgPCxhjre4fXdfLNtNM+Jd+Zdh8xd8wP87uNPoPgv4W7/5P2BuxfsMabNnMnza+54Pdi5U671GPZY8CehX8Voeoo7FHpkeEc6715FwHZrIrUrHaviPUbPZHND+IhczrP6FcYvhOZ0Di/ETt0OI+YwNWR9r7tpf6WDeZKZDB1+z2IthOl1mPyb5FluvEx9h9d0NnM0Y1XPFkWIsk1WotJ0PBMmkvjvQTd0e71tfeV+8r8lQ/tpzpsmxJ+InrI/dj2UajUajVTUajatRqNRtGo1Go1Go4wjeMpZFMVV9CHbofPraLsJ3JpWV2XOoanCuFky4y3PPNxucK2uKC1Lbdb1eo+m5XomN6HfeZsabHLHRX/K+offtNGGmHWctcVcG44MdSqsOLY9VzX+Zxfxn2HPdWTpzWvkrtJ8M5zorrKcquRytJ5N5DZmcaW02l76nWO+BqPXm1A2Ry/0q71dH/mqrqeFjkYxjEXtsX8qubTk67rGycyqsdm4tZx5D6D5hhi0waaWmiaMP81Yjii5qxPlPuU/GfTL1Y5E6Jyfiq63qTa39A4J0sOGDgO9WF9bOXl0XfPRbsY2bPNKPy1YrFYrFYmRhhlTIyMjJWJYZHXuCXI8OoXsvfljGLFicNifpp2XunoPiG1wtx3p1Tah+/DD66OnVtVXP9rKbVxOnL0tR/rHtqB5UDErUVcl11D4qqvjpOcxX7armUNJB3LpW6bxVvD08e8h3odKKvyCFZBdSh2FVcST9xV3n3T8t1j7Kr9qgrqXg+13Pt5U7JCvFXVIV1YG5lRhkVYZJYYDDD4KOIMoHCp26WS8GB7uBh2zIdgq/PKyInjV2STShuoapUdCpX1yTwqq/z1VvET7Kh5nVPkO8YyxjLt2MaaMmWTLQvx3qnzltnXW0p2jxgbEtSny/Osv8Y9pLMXYoHVPAhkVdWVeODhR6q9/Sxe2liwwZWMVvFXfRkeIDxAePUPIrdJ4ey6yquzH+PD/bUOWAu05qVHtFd8rrKHSoeNIOUqrYr3FXyToqfYJgwmJdKpXXOwYYegNNGMzfZPp/t3t/DVs4zjNTN61rRqaWaa4NYbRjTa0tWwy2Y2tGN8ZO8ofNKq4j9SL7I+cSm4/6ovLV5HNXLI0jJidwrtk6ynCaP6Z++GjRlWS3tLeW129Mi9evxU9mtz6s5J3Z7M2ngTgnKvmpomxpaLCzPfmx0JWE+m3NLDDGOX47RctdYYNK5jakdqLkRlI39n590T5zctGSwwZZDJj6kW8XSi6ot2MmWWJ0DUT3nuvebBudScjZ79g8cWJ8av0k+/bE5WKd5MdbFpbDVMxu1DVMmtNZGJvq1mtRbn6M+g/kP0FwDwr7quZs7xosNGpbscyxhhd9TyJyFwbLcxlTasg75vW7TsV5K7ji44XPMMrdoj+Y3rT0Hie62nlYV/pwczzOmdLqLhYkzGMzCZWGMQzGMSsZYY6Di1t4nlJ+Em63mJxrVLxPbYxNEdgc1dU2iOKyoYYWjNrEeHTYybVk0atSa7ehuwsWMWTqn1TrnS6hYsi71d1+s+k+ic70e20fzE/VaTdxT9ZtU4GIXdeNx3X77guYYfpHeTQjaMX6brOu4OY4K7Y2d9mbHarI5ox3p4GpJ2Vd/Tst60f7j999pppjR+Q/Qf8J/VaORs3cji7FfFuN61+ui9s8hix1OCh5KGVV23BPXvZfz3CLyHpix+exi8z/KnCnosY2eunor+cxyPO/xJ0vKey9OvE9VjqaYu0x3Z3jd6o2b1T12D+F8l232lwaaacD5LE8LBxu7WTlbWraWpew8Xexjel3E+wWD4APITdNqR8F3R3T0lunCQ4GaE9R37DxeCYfcHi4xci5ovKfxVs55y2hf+65E/Xdp6jR5nrebTmi5incpkyOjs50JvrZwstbbW6kfuuQw+2mykf/EXNFzxfKTrxew929TR6bWnGL//F3JFOFCQT3K4lQ"63 64    kernels = Kernel(65        bz2.decompress(base64.b64decode(quantization_code)),66        [67            "int4WeightCompression",68            "int4WeightExtractionFloat",69            "int4WeightExtractionHalf",70            "int8WeightExtractionFloat",71            "int8WeightExtractionHalf",72        ],73    )74except Exception as exception:75    kernels = None76    logger.warning("Failed to load cpm_kernels:" + str(exception))77 78 79class W8A16Linear(torch.autograd.Function):80    @staticmethod81    def forward(ctx, inp: torch.Tensor, quant_w: torch.Tensor, scale_w: torch.Tensor, weight_bit_width):82        ctx.inp_shape = inp.size()83        ctx.weight_bit_width = weight_bit_width84        out_features = quant_w.size(0)85        inp = inp.contiguous().view(-1, inp.size(-1))86        weight = extract_weight_to_half(quant_w, scale_w, weight_bit_width)87        ctx.weight_shape = weight.size()88        output = inp.mm(weight.t())89        ctx.save_for_backward(inp, quant_w, scale_w)90        return output.view(*(ctx.inp_shape[:-1] + (out_features,)))91 92    @staticmethod93    def backward(ctx, grad_output: torch.Tensor):94        inp, quant_w, scale_w = ctx.saved_tensors95        weight = extract_weight_to_half(quant_w, scale_w, ctx.weight_bit_width)96        grad_output = grad_output.contiguous().view(-1, weight.size(0))97        grad_input = grad_output.mm(weight)98        grad_weight = grad_output.t().mm(inp)99        return grad_input.view(ctx.inp_shape), grad_weight.view(ctx.weight_shape), None, None100 101 102def compress_int4_weight(weight: torch.Tensor):  # (n, m)103    with torch.cuda.device(weight.device):104        n, m = weight.size(0), weight.size(1)105        assert m % 2 == 0106        m = m // 2107        out = torch.empty(n, m, dtype=torch.int8, device="cuda")108        stream = torch.cuda.current_stream()109 110        gridDim = (n, 1, 1)111        blockDim = (min(round_up(m, 32), 1024), 1, 1)112 113        kernels.int4WeightCompression(114            gridDim,115            blockDim,116            0,117            stream,118            [ctypes.c_void_p(weight.data_ptr()), ctypes.c_void_p(out.data_ptr()), ctypes.c_int32(n), ctypes.c_int32(m)],119        )120        return out121 122 123def extract_weight_to_half(weight: torch.Tensor, scale_list: torch.Tensor, source_bit_width: int):124    assert scale_list.dtype in [torch.half, torch.bfloat16]125    assert weight.dtype in [torch.int8]126    if source_bit_width == 8:127        return weight.to(scale_list.dtype) * scale_list[:, None]128    elif source_bit_width == 4:129        func = (130            kernels.int4WeightExtractionHalf if scale_list.dtype == torch.half else kernels.int4WeightExtractionBFloat16131        )132    else:133        assert False, "Unsupported bit-width"134 135    with torch.cuda.device(weight.device):136        n, m = weight.size(0), weight.size(1)137        out = torch.empty(n, m * (8 // source_bit_width), dtype=scale_list.dtype, device="cuda")138        stream = torch.cuda.current_stream()139 140        gridDim = (n, 1, 1)141        blockDim = (min(round_up(m, 32), 1024), 1, 1)142 143        func(144            gridDim,145            blockDim,146            0,147            stream,148            [149                ctypes.c_void_p(weight.data_ptr()),150                ctypes.c_void_p(scale_list.data_ptr()),151                ctypes.c_void_p(out.data_ptr()),152                ctypes.c_int32(n),153                ctypes.c_int32(m),154            ],155        )156        return out157 158 159class QuantizedLinear(torch.nn.Module):160    def __init__(self, weight_bit_width: int, weight, bias=None, device="cuda", dtype=None, empty_init=False):161        super().__init__()162        weight = weight.to(device)  # ensure the weight is on the cuda device163        assert str(weight.device).startswith(164            'cuda'), 'The weights that need to be quantified should be on the CUDA device'165        self.weight_bit_width = weight_bit_width166        shape = weight.shape167 168        if weight is None or empty_init:169            self.weight = torch.empty(shape[0], shape[1] * weight_bit_width // 8, dtype=torch.int8, device=device)170            self.weight_scale = torch.empty(shape[0], dtype=dtype, device=device)171        else:172            self.weight_scale = weight.abs().max(dim=-1).values / ((2 ** (weight_bit_width - 1)) - 1)173            self.weight = torch.round(weight / self.weight_scale[:, None]).to(torch.int8)174            if weight_bit_width == 4:175                self.weight = compress_int4_weight(self.weight)176 177        self.weight = Parameter(self.weight.to(device), requires_grad=False)178        self.weight_scale = Parameter(self.weight_scale.to(device), requires_grad=False)179        self.bias = Parameter(bias.to(device), requires_grad=False) if bias is not None else None180 181    def forward(self, input):182        output = W8A16Linear.apply(input, self.weight, self.weight_scale, self.weight_bit_width)183        if self.bias is not None:184            output = output + self.bias185        return output186 187 188def quantize(model, weight_bit_width, empty_init=False, device=None):189    """Replace fp16 linear with quantized linear"""190    for layer in model.layers:191        layer.self_attention.query_key_value = QuantizedLinear(192            weight_bit_width=weight_bit_width,193            weight=layer.self_attention.query_key_value.weight,194            bias=layer.self_attention.query_key_value.bias,195            dtype=layer.self_attention.query_key_value.weight.dtype,196            device=layer.self_attention.query_key_value.weight.device if device is None else device,197            empty_init=empty_init198        )199        layer.self_attention.dense = QuantizedLinear(200            weight_bit_width=weight_bit_width,201            weight=layer.self_attention.dense.weight,202            bias=layer.self_attention.dense.bias,203            dtype=layer.self_attention.dense.weight.dtype,204            device=layer.self_attention.dense.weight.device if device is None else device,205            empty_init=empty_init206        )207        layer.mlp.dense_h_to_4h = QuantizedLinear(208            weight_bit_width=weight_bit_width,209            weight=layer.mlp.dense_h_to_4h.weight,210            bias=layer.mlp.dense_h_to_4h.bias,211            dtype=layer.mlp.dense_h_to_4h.weight.dtype,212            device=layer.mlp.dense_h_to_4h.weight.device if device is None else device,213            empty_init=empty_init214        )215        layer.mlp.dense_4h_to_h = QuantizedLinear(216            weight_bit_width=weight_bit_width,217            weight=layer.mlp.dense_4h_to_h.weight,218            bias=layer.mlp.dense_4h_to_h.bias,219            dtype=layer.mlp.dense_4h_to_h.weight.dtype,220            device=layer.mlp.dense_4h_to_h.weight.device if device is None else device,221            empty_init=empty_init222        )223 224    return model225 226 227 228class ChatGLMConfig(PretrainedConfig):229    model_type = "chatglm"230    def __init__(231        self,232        num_layers=28,233        padded_vocab_size=65024,234        hidden_size=4096,235        ffn_hidden_size=13696,236        kv_channels=128,237        num_attention_heads=32,238        seq_length=2048,239        hidden_dropout=0.0,240        classifier_dropout=None,241        attention_dropout=0.0,242        layernorm_epsilon=1e-5,243        rmsnorm=True,244        apply_residual_connection_post_layernorm=False,245        post_layer_norm=True,246        add_bias_linear=False,247        add_qkv_bias=False,248        bias_dropout_fusion=True,249        multi_query_attention=False,250        multi_query_group_num=1,251        apply_query_key_layer_scaling=True,252        attention_softmax_in_fp32=True,253        fp32_residual_connection=False,254        quantization_bit=0,255        pre_seq_len=None,256        prefix_projection=False,257        **kwargs258    ):259        self.num_layers = num_layers260        self.vocab_size = padded_vocab_size261        self.padded_vocab_size = padded_vocab_size262        self.hidden_size = hidden_size263        self.ffn_hidden_size = ffn_hidden_size264        self.kv_channels = kv_channels265        self.num_attention_heads = num_attention_heads266        self.seq_length = seq_length267        self.hidden_dropout = hidden_dropout268        self.classifier_dropout = classifier_dropout269        self.attention_dropout = attention_dropout270        self.layernorm_epsilon = layernorm_epsilon271        self.rmsnorm = rmsnorm272        self.apply_residual_connection_post_layernorm = apply_residual_connection_post_layernorm273        self.post_layer_norm = post_layer_norm274        self.add_bias_linear = add_bias_linear275        self.add_qkv_bias = add_qkv_bias276        self.bias_dropout_fusion = bias_dropout_fusion277        self.multi_query_attention = multi_query_attention278        self.multi_query_group_num = multi_query_group_num279        self.apply_query_key_layer_scaling = apply_query_key_layer_scaling280        self.attention_softmax_in_fp32 = attention_softmax_in_fp32281        self.fp32_residual_connection = fp32_residual_connection282        self.quantization_bit = quantization_bit283        self.pre_seq_len = pre_seq_len284        self.prefix_projection = prefix_projection285        super().__init__(**kwargs)286 287 288 289# flags required to enable jit fusion kernels290 291if sys.platform != 'darwin':292    torch._C._jit_set_profiling_mode(False)293    torch._C._jit_set_profiling_executor(False)294    torch._C._jit_override_can_fuse_on_cpu(True)295    torch._C._jit_override_can_fuse_on_gpu(True)296 297logger = logging.get_logger(__name__)298 299_CHECKPOINT_FOR_DOC = "THUDM/ChatGLM"300_CONFIG_FOR_DOC = "ChatGLM6BConfig"301 302CHATGLM_6B_PRETRAINED_MODEL_ARCHIVE_LIST = [303    "THUDM/chatglm3-6b-base",304    # See all ChatGLM models at https://huggingface.co/models?filter=chatglm305]306 307 308def default_init(cls, *args, **kwargs):309    return cls(*args, **kwargs)310 311 312class InvalidScoreLogitsProcessor(LogitsProcessor):313    def __call__(self, input_ids: torch.LongTensor, scores: torch.FloatTensor) -> torch.FloatTensor:314        if torch.isnan(scores).any() or torch.isinf(scores).any():315            scores.zero_()316            scores[..., 5] = 5e4317        return scores318 319 320class PrefixEncoder(torch.nn.Module):321    """322    The torch.nn model to encode the prefix323    Input shape: (batch-size, prefix-length)324    Output shape: (batch-size, prefix-length, 2*layers*hidden)325    """326 327    def __init__(self, config: ChatGLMConfig):328        super().__init__()329        self.prefix_projection = config.prefix_projection330        if self.prefix_projection:331            # Use a two-layer MLP to encode the prefix332            kv_size = config.num_layers * config.kv_channels * config.multi_query_group_num * 2333            self.embedding = torch.nn.Embedding(config.pre_seq_len, kv_size)334            self.trans = torch.nn.Sequential(335                torch.nn.Linear(kv_size, config.hidden_size),336                torch.nn.Tanh(),337                torch.nn.Linear(config.hidden_size, kv_size)338            )339        else:340            self.embedding = torch.nn.Embedding(config.pre_seq_len,341                                                config.num_layers * config.kv_channels * config.multi_query_group_num * 2)342 343    def forward(self, prefix: torch.Tensor):344        if self.prefix_projection:345            prefix_tokens = self.embedding(prefix)346            past_key_values = self.trans(prefix_tokens)347        else:348            past_key_values = self.embedding(prefix)349        return past_key_values350 351 352def split_tensor_along_last_dim(353        tensor: torch.Tensor,354        num_partitions: int,355        contiguous_split_chunks: bool = False,356) -> List[torch.Tensor]:357    """Split a tensor along its last dimension.358 359    Arguments:360        tensor: input tensor.361        num_partitions: number of partitions to split the tensor362        contiguous_split_chunks: If True, make each chunk contiguous363                                 in memory.364 365    Returns:366        A list of Tensors367    """368    # Get the size and dimension.369    last_dim = tensor.dim() - 1370    last_dim_size = tensor.size()[last_dim] // num_partitions371    # Split.372    tensor_list = torch.split(tensor, last_dim_size, dim=last_dim)373    # Note: torch.split does not create contiguous tensors by default.374    if contiguous_split_chunks:375        return tuple(chunk.contiguous() for chunk in tensor_list)376 377    return tensor_list378 379 380class RotaryEmbedding(nn.Module):381    def __init__(self, dim, original_impl=False, device=None, dtype=None):382        super().__init__()383        inv_freq = 1.0 / (10000 ** (torch.arange(0, dim, 2, device=device).to(dtype=dtype) / dim))384        self.register_buffer("inv_freq", inv_freq)385        self.dim = dim386        self.original_impl = original_impl387 388    def forward_impl(389            self, seq_len: int, n_elem: int, dtype: torch.dtype, device: torch.device, base: int = 10000390    ):391        """Enhanced Transformer with Rotary Position Embedding.392 393        Derived from: https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/master/labml_nn/394        transformers/rope/__init__.py. MIT License:395        https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/master/license.396        """397        # $\Theta = {\theta_i = 10000^{\frac{2(i-1)}{d}}, i \in [1, 2, ..., \frac{d}{2}]}$398        theta = 1.0 / (base ** (torch.arange(0, n_elem, 2, dtype=torch.float, device=device) / n_elem))399 400        # Create position indexes `[0, 1, ..., seq_len - 1]`401        seq_idx = torch.arange(seq_len, dtype=torch.float, device=device)402 403        # Calculate the product of position index and $\theta_i$404        idx_theta = torch.outer(seq_idx, theta).float()405 406        cache = torch.stack([torch.cos(idx_theta), torch.sin(idx_theta)], dim=-1)407 408        # this is to mimic the behaviour of complex32, else we will get different results409        if dtype in (torch.float16, torch.bfloat16, torch.int8):410            cache = cache.bfloat16() if dtype == torch.bfloat16 else cache.half()411        return cache412 413    def forward(self, max_seq_len, offset=0):414        return self.forward_impl(415            max_seq_len, self.dim, dtype=self.inv_freq.dtype, device=self.inv_freq.device416        )417 418 419@torch.jit.script420def apply_rotary_pos_emb(x: torch.Tensor, rope_cache: torch.Tensor) -> torch.Tensor:421    # x: [sq, b, np, hn]422    sq, b, np, hn = x.size(0), x.size(1), x.size(2), x.size(3)423    rot_dim = rope_cache.shape[-2] * 2424    x, x_pass = x[..., :rot_dim], x[..., rot_dim:]425    # truncate to support variable sizes426    rope_cache = rope_cache[:sq]427    xshaped = x.reshape(sq, -1, np, rot_dim // 2, 2)428    rope_cache = rope_cache.view(sq, -1, 1, xshaped.size(3), 2)429    x_out2 = torch.stack(430        [431            xshaped[..., 0] * rope_cache[..., 0] - xshaped[..., 1] * rope_cache[..., 1],432            xshaped[..., 1] * rope_cache[..., 0] + xshaped[..., 0] * rope_cache[..., 1],433        ],434        -1,435    )436    x_out2 = x_out2.flatten(3)437    return torch.cat((x_out2, x_pass), dim=-1)438 439 440class RMSNorm(torch.nn.Module):441    def __init__(self, normalized_shape, eps=1e-5, device=None, dtype=None, **kwargs):442        super().__init__()443        self.weight = torch.nn.Parameter(torch.empty(normalized_shape, device=device, dtype=dtype))444        self.eps = eps445 446    def forward(self, hidden_states: torch.Tensor):447        input_dtype = hidden_states.dtype448        variance = hidden_states.to(torch.float32).pow(2).mean(-1, keepdim=True)449        hidden_states = hidden_states * torch.rsqrt(variance + self.eps)450 451        return (self.weight * hidden_states).to(input_dtype)452 453 454class CoreAttention(torch.nn.Module):455    def __init__(self, config: ChatGLMConfig, layer_number):456        super(CoreAttention, self).__init__()457 458        self.apply_query_key_layer_scaling = config.apply_query_key_layer_scaling459        self.attention_softmax_in_fp32 = config.attention_softmax_in_fp32460        if self.apply_query_key_layer_scaling:461            self.attention_softmax_in_fp32 = True462        self.layer_number = max(1, layer_number)463 464        projection_size = config.kv_channels * config.num_attention_heads465 466        # Per attention head and per partition values.467        self.hidden_size_per_partition = projection_size468        self.hidden_size_per_attention_head = projection_size // config.num_attention_heads469        self.num_attention_heads_per_partition = config.num_attention_heads470 471        coeff = None472        self.norm_factor = math.sqrt(self.hidden_size_per_attention_head)473        if self.apply_query_key_layer_scaling:474            coeff = self.layer_number475            self.norm_factor *= coeff476        self.coeff = coeff477 478        self.attention_dropout = torch.nn.Dropout(config.attention_dropout)479 480    def forward(self, query_layer, key_layer, value_layer, attention_mask):481        pytorch_major_version = int(torch.__version__.split('.')[0])482        if pytorch_major_version >= 2:483            query_layer, key_layer, value_layer = [k.permute(1, 2, 0, 3) for k in [query_layer, key_layer, value_layer]]484            if attention_mask is None and query_layer.shape[2] == key_layer.shape[2]:485                context_layer = torch.nn.functional.scaled_dot_product_attention(query_layer, key_layer, value_layer,486                                                                                 is_causal=True)487            else:488                if attention_mask is not None:489                    attention_mask = ~attention_mask490                context_layer = torch.nn.functional.scaled_dot_product_attention(query_layer, key_layer, value_layer,491                                                                                 attention_mask)492            context_layer = context_layer.permute(2, 0, 1, 3)493            new_context_layer_shape = context_layer.size()[:-2] + (self.hidden_size_per_partition,)494            context_layer = context_layer.reshape(*new_context_layer_shape)495        else:496            # Raw attention scores497 498            # [b, np, sq, sk]499            output_size = (query_layer.size(1), query_layer.size(2), query_layer.size(0), key_layer.size(0))500 501            # [sq, b, np, hn] -> [sq, b * np, hn]502            query_layer = query_layer.view(output_size[2], output_size[0] * output_size[1], -1)503            # [sk, b, np, hn] -> [sk, b * np, hn]504            key_layer = key_layer.view(output_size[3], output_size[0] * output_size[1], -1)505 506            # preallocting input tensor: [b * np, sq, sk]507            matmul_input_buffer = torch.empty(508                output_size[0] * output_size[1], output_size[2], output_size[3], dtype=query_layer.dtype,509                device=query_layer.device510            )511 512            # Raw attention scores. [b * np, sq, sk]513            matmul_result = torch.baddbmm(514                matmul_input_buffer,515                query_layer.transpose(0, 1),  # [b * np, sq, hn]516                key_layer.transpose(0, 1).transpose(1, 2),  # [b * np, hn, sk]517                beta=0.0,518                alpha=(1.0 / self.norm_factor),519            )520 521            # change view to [b, np, sq, sk]522            attention_scores = matmul_result.view(*output_size)523 524            # ===========================525            # Attention probs and dropout526            # ===========================527 528            # attention scores and attention mask [b, np, sq, sk]529            if self.attention_softmax_in_fp32:530                attention_scores = attention_scores.float()531            if self.coeff is not None:532                attention_scores = attention_scores * self.coeff533            if attention_mask is None and attention_scores.shape[2] == attention_scores.shape[3]:534                attention_mask = torch.ones(output_size[0], 1, output_size[2], output_size[3],535                                            device=attention_scores.device, dtype=torch.bool)536                attention_mask.tril_()537                attention_mask = ~attention_mask538            if attention_mask is not None:539                attention_scores = attention_scores.masked_fill(attention_mask, float("-inf"))540            attention_probs = F.softmax(attention_scores, dim=-1)541            attention_probs = attention_probs.type_as(value_layer)542 543            # This is actually dropping out entire tokens to attend to, which might544            # seem a bit unusual, but is taken from the original Transformer paper.545            attention_probs = self.attention_dropout(attention_probs)546            # =========================547            # Context layer. [sq, b, hp]548            # =========================549 550            # value_layer -> context layer.551            # [sk, b, np, hn] --> [b, np, sq, hn]552 553            # context layer shape: [b, np, sq, hn]554            output_size = (value_layer.size(1), value_layer.size(2), query_layer.size(0), value_layer.size(3))555            # change view [sk, b * np, hn]556            value_layer = value_layer.view(value_layer.size(0), output_size[0] * output_size[1], -1)557            # change view [b * np, sq, sk]558            attention_probs = attention_probs.view(output_size[0] * output_size[1], output_size[2], -1)559            # matmul: [b * np, sq, hn]560            context_layer = torch.bmm(attention_probs, value_layer.transpose(0, 1))561            # change view [b, np, sq, hn]562            context_layer = context_layer.view(*output_size)563            # [b, np, sq, hn] --> [sq, b, np, hn]564            context_layer = context_layer.permute(2, 0, 1, 3).contiguous()565            # [sq, b, np, hn] --> [sq, b, hp]566            new_context_layer_shape = context_layer.size()[:-2] + (self.hidden_size_per_partition,)567            context_layer = context_layer.view(*new_context_layer_shape)568 569        return context_layer570 571 572class SelfAttention(torch.nn.Module):573    """Parallel self-attention layer abstract class.574 575    Self-attention layer takes input with size [s, b, h]576    and returns output of the same size.577    """578 579    def __init__(self, config: ChatGLMConfig, layer_number, device=None):580        super(SelfAttention, self).__init__()581        self.layer_number = max(1, layer_number)582 583        self.projection_size = config.kv_channels * config.num_attention_heads584 585        # Per attention head and per partition values.586        self.hidden_size_per_attention_head = self.projection_size // config.num_attention_heads587        self.num_attention_heads_per_partition = config.num_attention_heads588 589        self.multi_query_attention = config.multi_query_attention590        self.qkv_hidden_size = 3 * self.projection_size591        if self.multi_query_attention:592            self.num_multi_query_groups_per_partition = config.multi_query_group_num593            self.qkv_hidden_size = (594                    self.projection_size + 2 * self.hidden_size_per_attention_head * config.multi_query_group_num595            )596        self.query_key_value = nn.Linear(config.hidden_size, self.qkv_hidden_size,597                                         bias=config.add_bias_linear or config.add_qkv_bias,598                                         device=device, **_config_to_kwargs(config)599                                         )600 601        self.core_attention = CoreAttention(config, self.layer_number)602 603        # Output.604        self.dense = nn.Linear(self.projection_size, config.hidden_size, bias=config.add_bias_linear,605                               device=device, **_config_to_kwargs(config)606                               )607 608    def _allocate_memory(self, inference_max_sequence_len, batch_size, device=None, dtype=None):609        if self.multi_query_attention:610            num_attention_heads = self.num_multi_query_groups_per_partition611        else:612            num_attention_heads = self.num_attention_heads_per_partition613        return torch.empty(614            inference_max_sequence_len,615            batch_size,616            num_attention_heads,617            self.hidden_size_per_attention_head,618            dtype=dtype,619            device=device,620        )621 622    def forward(623            self, hidden_states, attention_mask, rotary_pos_emb, kv_cache=None, use_cache=True624    ):625        # hidden_states: [sq, b, h]626 627        # =================================================628        # Pre-allocate memory for key-values for inference.629        # =================================================630        # =====================631        # Query, Key, and Value632        # =====================633 634        # Attention heads [sq, b, h] --> [sq, b, (np * 3 * hn)]635        mixed_x_layer = self.query_key_value(hidden_states)636 637        if self.multi_query_attention:638            (query_layer, key_layer, value_layer) = mixed_x_layer.split(639                [640                    self.num_attention_heads_per_partition * self.hidden_size_per_attention_head,641                    self.num_multi_query_groups_per_partition * self.hidden_size_per_attention_head,642                    self.num_multi_query_groups_per_partition * self.hidden_size_per_attention_head,643                ],644                dim=-1,645            )646            query_layer = query_layer.view(647                query_layer.size()[:-1] + (self.num_attention_heads_per_partition, self.hidden_size_per_attention_head)648            )649            key_layer = key_layer.view(650                key_layer.size()[:-1] + (self.num_multi_query_groups_per_partition, self.hidden_size_per_attention_head)651            )652            value_layer = value_layer.view(653                value_layer.size()[:-1]654                + (self.num_multi_query_groups_per_partition, self.hidden_size_per_attention_head)655            )656        else:657            new_tensor_shape = mixed_x_layer.size()[:-1] + \658                               (self.num_attention_heads_per_partition,659                                3 * self.hidden_size_per_attention_head)660            mixed_x_layer = mixed_x_layer.view(*new_tensor_shape)661 662            # [sq, b, np, 3 * hn] --> 3 [sq, b, np, hn]663            (query_layer, key_layer, value_layer) = split_tensor_along_last_dim(mixed_x_layer, 3)664 665        # apply relative positional encoding (rotary embedding)666        if rotary_pos_emb is not None:667            query_layer = apply_rotary_pos_emb(query_layer, rotary_pos_emb)668            key_layer = apply_rotary_pos_emb(key_layer, rotary_pos_emb)669 670        # adjust key and value for inference671        if kv_cache is not None:672            cache_k, cache_v = kv_cache673            key_layer = torch.cat((cache_k, key_layer), dim=0)674            value_layer = torch.cat((cache_v, value_layer), dim=0)675        if use_cache:676            kv_cache = (key_layer, value_layer)677        else:678            kv_cache = None679 680        if self.multi_query_attention:681            key_layer = key_layer.unsqueeze(-2)682            key_layer = key_layer.expand(683                -1, -1, -1, self.num_attention_heads_per_partition // self.num_multi_query_groups_per_partition, -1684            )685            key_layer = key_layer.contiguous().view(686                key_layer.size()[:2] + (self.num_attention_heads_per_partition, self.hidden_size_per_attention_head)687            )688            value_layer = value_layer.unsqueeze(-2)689            value_layer = value_layer.expand(690                -1, -1, -1, self.num_attention_heads_per_partition // self.num_multi_query_groups_per_partition, -1691            )692            value_layer = value_layer.contiguous().view(693                value_layer.size()[:2] + (self.num_attention_heads_per_partition, self.hidden_size_per_attention_head)694            )695 696        # ==================================697        # core attention computation698        # ==================================699 700        context_layer = self.core_attention(query_layer, key_layer, value_layer, attention_mask)701 702        # =================703        # Output. [sq, b, h]704        # =================705 706        output = self.dense(context_layer)707 708        return output, kv_cache709 710 711def _config_to_kwargs(args):712    common_kwargs = {713        "dtype": args.torch_dtype,714    }715    return common_kwargs716 717 718class MLP(torch.nn.Module):719    """MLP.720 721    MLP will take the input with h hidden state, project it to 4*h722    hidden dimension, perform nonlinear transformation, and project the723    state back into h hidden dimension.724    """725 726    def __init__(self, config: ChatGLMConfig, device=None):727        super(MLP, self).__init__()728 729        self.add_bias = config.add_bias_linear730 731        # Project to 4h. If using swiglu double the output width, see https://arxiv.org/pdf/2002.05202.pdf732        self.dense_h_to_4h = nn.Linear(733            config.hidden_size,734            config.ffn_hidden_size * 2,735            bias=self.add_bias,736            device=device,737            **_config_to_kwargs(config)738        )739 740        def swiglu(x):741            x = torch.chunk(x, 2, dim=-1)742            return F.silu(x[0]) * x[1]743 744        self.activation_func = swiglu745 746        # Project back to h.747        self.dense_4h_to_h = nn.Linear(748            config.ffn_hidden_size,749            config.hidden_size,750            bias=self.add_bias,751            device=device,752            **_config_to_kwargs(config)753        )754 755    def forward(self, hidden_states):756        # [s, b, 4hp]757        intermediate_parallel = self.dense_h_to_4h(hidden_states)758        intermediate_parallel = self.activation_func(intermediate_parallel)759        # [s, b, h]760        output = self.dense_4h_to_h(intermediate_parallel)761        return output762 763 764class GLMBlock(torch.nn.Module):765    """A single transformer layer.766 767    Transformer layer takes input with size [s, b, h] and returns an768    output of the same size.769    """770 771    def __init__(self, config: ChatGLMConfig, layer_number, device=None):772        super(GLMBlock, self).__init__()773        self.layer_number = layer_number774 775        self.apply_residual_connection_post_layernorm = config.apply_residual_connection_post_layernorm776 777        self.fp32_residual_connection = config.fp32_residual_connection778 779        LayerNormFunc = RMSNorm if config.rmsnorm else LayerNorm780        # Layernorm on the input data.781        self.input_layernorm = LayerNormFunc(config.hidden_size, eps=config.layernorm_epsilon, device=device,782                                             dtype=config.torch_dtype)783 784        # Self attention.785        self.self_attention = SelfAttention(config, layer_number, device=device)786        self.hidden_dropout = config.hidden_dropout787 788        # Layernorm on the attention output789        self.post_attention_layernorm = LayerNormFunc(config.hidden_size, eps=config.layernorm_epsilon, device=device,790                                                      dtype=config.torch_dtype)791 792        # MLP793        self.mlp = MLP(config, device=device)794 795    def forward(796            self, hidden_states, attention_mask, rotary_pos_emb, kv_cache=None, use_cache=True,797    ):798        # hidden_states: [s, b, h]799 800        # Layer norm at the beginning of the transformer layer.801        layernorm_output = self.input_layernorm(hidden_states)802        # Self attention.803        attention_output, kv_cache = self.self_attention(804            layernorm_output,805            attention_mask,806            rotary_pos_emb,807            kv_cache=kv_cache,808            use_cache=use_cache809        )810 811        # Residual connection.812        if self.apply_residual_connection_post_layernorm:813            residual = layernorm_output814        else:815            residual = hidden_states816 817        layernorm_input = torch.nn.functional.dropout(attention_output, p=self.hidden_dropout, training=self.training)818        layernorm_input = residual + layernorm_input819 820        # Layer norm post the self attention.821        layernorm_output = self.post_attention_layernorm(layernorm_input)822 823        # MLP.824        mlp_output = self.mlp(layernorm_output)825 826        # Second residual connection.827        if self.apply_residual_connection_post_layernorm:828            residual = layernorm_output829        else:830            residual = layernorm_input831 832        output = torch.nn.functional.dropout(mlp_output, p=self.hidden_dropout, training=self.training)833        output = residual + output834 835        return output, kv_cache836 837 838class GLMTransformer(torch.nn.Module):839    """Transformer class."""840 841    def __init__(self, config: ChatGLMConfig, device=None):842        super(GLMTransformer, self).__init__()843 844        self.fp32_residual_connection = config.fp32_residual_connection845        self.post_layer_norm = config.post_layer_norm846 847        # Number of layers.848        self.num_layers = config.num_layers849 850        # Transformer layers.851        def build_layer(layer_number):852            return GLMBlock(config, layer_number, device=device)853 854        self.layers = torch.nn.ModuleList([build_layer(i + 1) for i in range(self.num_layers)])855 856        if self.post_layer_norm:857            LayerNormFunc = RMSNorm if config.rmsnorm else LayerNorm858            # Final layer norm before output.859            self.final_layernorm = LayerNormFunc(config.hidden_size, eps=config.layernorm_epsilon, device=device,860                                                 dtype=config.torch_dtype)861 862        self.gradient_checkpointing = False863 864    def _get_layer(self, layer_number):865        return self.layers[layer_number]866 867    def forward(868            self, hidden_states, attention_mask, rotary_pos_emb, kv_caches=None,869            use_cache: Optional[bool] = True,870            output_hidden_states: Optional[bool] = False,871    ):872        if not kv_caches:873            kv_caches = [None for _ in range(self.num_layers)]874        presents = () if use_cache else None875        if self.gradient_checkpointing and self.training:876            if use_cache:877                logger.warning_once(878                    "`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`..."879                )880                use_cache = False881 882        all_self_attentions = None883        all_hidden_states = () if output_hidden_states else None884        for index in range(self.num_layers):885            if output_hidden_states:886                all_hidden_states = all_hidden_states + (hidden_states,)887 888            layer = self._get_layer(index)889            if self.gradient_checkpointing and self.training:890                layer_ret = torch.utils.checkpoint.checkpoint(891                    layer,892                    hidden_states,893                    attention_mask,894                    rotary_pos_emb,895                    kv_caches[index],896                    use_cache897                )898            else:899                layer_ret = layer(900                    hidden_states,901                    attention_mask,902                    rotary_pos_emb,903                    kv_cache=kv_caches[index],904                    use_cache=use_cache905                )906            hidden_states, kv_cache = layer_ret907            if use_cache:908                presents = presents + (kv_cache,)909 910        if output_hidden_states:911            all_hidden_states = all_hidden_states + (hidden_states,)912 913        # Final layer norm.914        if self.post_layer_norm:915            hidden_states = self.final_layernorm(hidden_states)916 917        return hidden_states, presents, all_hidden_states, all_self_attentions918 919 920class ChatGLMPreTrainedModel(PreTrainedModel):921    """922    An abstract class to handle weights initialization and923    a simple interface for downloading and loading pretrained models.924    """925 926    is_parallelizable = False927    supports_gradient_checkpointing = True928    config_class = ChatGLMConfig929    base_model_prefix = "transformer"930    _no_split_modules = ["GLMBlock"]931 932    def _init_weights(self, module: nn.Module):933        """Initialize the weights."""934        return935 936    def get_masks(self, input_ids, past_key_values, padding_mask=None):937        batch_size, seq_length = input_ids.shape938        full_attention_mask = torch.ones(batch_size, seq_length, seq_length, device=input_ids.device)939        full_attention_mask.tril_()940        past_length = 0941        if past_key_values:942            past_length = past_key_values[0][0].shape[0]943        if past_length:944            full_attention_mask = torch.cat((torch.ones(batch_size, seq_length, past_length,945                                                        device=input_ids.device), full_attention_mask), dim=-1)946        if padding_mask is not None:947            full_attention_mask = full_attention_mask * padding_mask.unsqueeze(1)948        if not past_length and padding_mask is not None:949            full_attention_mask -= padding_mask.unsqueeze(-1) - 1950        full_attention_mask = (full_attention_mask < 0.5).bool()951        full_attention_mask.unsqueeze_(1)952        return full_attention_mask953 954    def get_position_ids(self, input_ids, device):955        batch_size, seq_length = input_ids.shape956        position_ids = torch.arange(seq_length, dtype=torch.long, device=device).unsqueeze(0).repeat(batch_size, 1)957        return position_ids958 959    def _set_gradient_checkpointing(self, module, value=False):960        if isinstance(module, GLMTransformer):961            module.gradient_checkpointing = value962 963 964class Embedding(torch.nn.Module):965    """Language model embeddings."""966 967    def __init__(self, config: ChatGLMConfig, device=None):968        super(Embedding, self).__init__()969 970        self.hidden_size = config.hidden_size971        # Word embeddings (parallel).972        self.word_embeddings = nn.Embedding(973            config.padded_vocab_size,974            self.hidden_size,975            dtype=config.torch_dtype,976            device=device977        )978        self.fp32_residual_connection = config.fp32_residual_connection979 980    def forward(self, input_ids):981        # Embeddings.982        words_embeddings = self.word_embeddings(input_ids)983        embeddings = words_embeddings984        # Data format change to avoid explicit tranposes : [b s h] --> [s b h].985        embeddings = embeddings.transpose(0, 1).contiguous()986        # If the input flag for fp32 residual connection is set, convert for float.987        if self.fp32_residual_connection:988            embeddings = embeddings.float()989        return embeddings990 991 992class ChatGLMModel(ChatGLMPreTrainedModel):993    def __init__(self, config: ChatGLMConfig, device=None, empty_init=True):994        super().__init__(config)995        if empty_init:996            init_method = skip_init997        else:998            init_method = default_init999        init_kwargs = {}1000        if device is not None:1001            init_kwargs["device"] = device1002        self.embedding = init_method(Embedding, config, **init_kwargs)1003        self.num_layers = config.num_layers1004        self.multi_query_group_num = config.multi_query_group_num1005        self.kv_channels = config.kv_channels1006 1007        # Rotary positional embeddings1008        self.seq_length = config.seq_length1009        rotary_dim = (1010            config.hidden_size // config.num_attention_heads if config.kv_channels is None else config.kv_channels1011        )1012 1013        self.rotary_pos_emb = RotaryEmbedding(rotary_dim // 2, original_impl=config.original_rope, device=device,1014                                              dtype=config.torch_dtype)1015        self.encoder = init_method(GLMTransformer, config, **init_kwargs)1016        self.output_layer = init_method(nn.Linear, config.hidden_size, config.padded_vocab_size, bias=False,1017                                        dtype=config.torch_dtype, **init_kwargs)1018        self.pre_seq_len = config.pre_seq_len1019        self.prefix_projection = config.prefix_projection1020        if self.pre_seq_len is not None:1021            for param in self.parameters():1022                param.requires_grad = False1023            self.prefix_tokens = torch.arange(self.pre_seq_len).long()1024            self.prefix_encoder = PrefixEncoder(config)1025            self.dropout = torch.nn.Dropout(0.1)1026 1027    def get_input_embeddings(self):1028        return self.embedding.word_embeddings1029 1030    def get_prompt(self, batch_size, device, dtype=torch.half):1031        prefix_tokens = self.prefix_tokens.unsqueeze(0).expand(batch_size, -1).to(device)1032        past_key_values = self.prefix_encoder(prefix_tokens).type(dtype)1033        past_key_values = past_key_values.view(1034            batch_size,1035            self.pre_seq_len,1036            self.num_layers * 2,1037            self.multi_query_group_num,1038            self.kv_channels1039        )1040        # seq_len, b, nh, hidden_size1041        past_key_values = self.dropout(past_key_values)1042        past_key_values = past_key_values.permute([2, 1, 0, 3, 4]).split(2)1043        return past_key_values1044 1045    def forward(1046            self,1047            input_ids,1048            position_ids: Optional[torch.Tensor] = None,1049            attention_mask: Optional[torch.BoolTensor] = None,1050            full_attention_mask: Optional[torch.BoolTensor] = None,1051            past_key_values: Optional[Tuple[Tuple[torch.Tensor, torch.Tensor], ...]] = None,1052            inputs_embeds: Optional[torch.Tensor] = None,1053            use_cache: Optional[bool] = None,1054            output_hidden_states: Optional[bool] = None,1055            return_dict: Optional[bool] = None,1056    ):1057        output_hidden_states = (1058            output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states1059        )1060        use_cache = use_cache if use_cache is not None else self.config.use_cache1061        return_dict = return_dict if return_dict is not None else self.config.use_return_dict1062 1063        batch_size, seq_length = input_ids.shape1064 1065        if inputs_embeds is None:1066            inputs_embeds = self.embedding(input_ids)1067 1068        if self.pre_seq_len is not None:1069            if past_key_values is None:1070                past_key_values = self.get_prompt(batch_size=batch_size, device=input_ids.device,1071                                                  dtype=inputs_embeds.dtype)1072            if attention_mask is not None:1073                attention_mask = torch.cat([attention_mask.new_ones((batch_size, self.pre_seq_len)),1074                                            attention_mask], dim=-1)1075 1076        if full_attention_mask is None:1077            if (attention_mask is not None and not attention_mask.all()) or (past_key_values and seq_length != 1):1078                full_attention_mask = self.get_masks(input_ids, past_key_values, padding_mask=attention_mask)1079 1080        # Rotary positional embeddings1081        rotary_pos_emb = self.rotary_pos_emb(self.seq_length)1082        if position_ids is not None:1083            rotary_pos_emb = rotary_pos_emb[position_ids]1084        else:1085            rotary_pos_emb = rotary_pos_emb[None, :seq_length]1086        rotary_pos_emb = rotary_pos_emb.transpose(0, 1).contiguous()1087 1088        # Run encoder.1089        hidden_states, presents, all_hidden_states, all_self_attentions = self.encoder(1090            inputs_embeds, full_attention_mask, rotary_pos_emb=rotary_pos_emb,1091            kv_caches=past_key_values, use_cache=use_cache, output_hidden_states=output_hidden_states1092        )1093 1094        if not return_dict:1095            return tuple(v for v in [hidden_states, presents, all_hidden_states, all_self_attentions] if v is not None)1096 1097        return BaseModelOutputWithPast(1098            last_hidden_state=hidden_states,1099            past_key_values=presents,1100            hidden_states=all_hidden_states,1101            attentions=all_self_attentions,1102        )1103 1104    def quantize(self, weight_bit_width: int):1105        # from .quantization import quantize1106        quantize(self.encoder, weight_bit_width)1107        return self1108 1109 1110class ChatGLMForConditionalGeneration(ChatGLMPreTrainedModel):1111    def __init__(self, config: ChatGLMConfig, empty_init=True, device=None):1112        super().__init__(config)1113 1114        self.max_sequence_length = config.max_length1115        self.transformer = ChatGLMModel(config, empty_init=empty_init, device=device)1116        self.config = config1117        self.quantized = False1118 1119        if self.config.quantization_bit:1120            self.quantize(self.config.quantization_bit, empty_init=True)1121 1122    def _update_model_kwargs_for_generation(1123            self,1124            outputs: ModelOutput,1125            model_kwargs: Dict[str, Any],1126            is_encoder_decoder: bool = False,1127            standardize_cache_format: bool = False,1128    ) -> Dict[str, Any]:1129        # update past_key_values1130        model_kwargs["past_key_values"] = self._extract_past_from_model_output(1131            outputs, standardize_cache_format=standardize_cache_format1132        )1133 1134        # update attention mask1135        if "attention_mask" in model_kwargs:1136            attention_mask = model_kwargs["attention_mask"]1137            model_kwargs["attention_mask"] = torch.cat(1138                [attention_mask, attention_mask.new_ones((attention_mask.shape[0], 1))], dim=-11139            )1140 1141        # update position ids1142        if "position_ids" in model_kwargs:1143            position_ids = model_kwargs["position_ids"]1144            new_position_id = position_ids[..., -1:].clone()1145            new_position_id += 11146            model_kwargs["position_ids"] = torch.cat(1147                [position_ids, new_position_id], dim=-11148            )1149 1150        model_kwargs["is_first_forward"] = False1151        return model_kwargs1152 1153    def prepare_inputs_for_generation(1154            self,1155            input_ids: torch.LongTensor,1156            past_key_values: Optional[torch.Tensor] = None,1157            attention_mask: Optional[torch.Tensor] = None,1158            position_ids: Optional[torch.Tensor] = None,1159            use_cache: Optional[bool] = None,1160            is_first_forward: bool = True,1161            **kwargs1162    ) -> dict:1163        # only last token for input_ids if past is not None1164        if position_ids is None:1165            position_ids = self.get_position_ids(input_ids, device=input_ids.device)1166        if not is_first_forward:1167            if past_key_values is not None:1168                position_ids = position_ids[..., -1:]1169                input_ids = input_ids[:, -1:]1170        return {1171            "input_ids": input_ids,1172            "past_key_values": past_key_values,1173            "position_ids": position_ids,1174            "attention_mask": attention_mask,1175            "return_last_logit": True,1176            "use_cache": use_cache1177        }1178 1179    def forward(1180            self,1181            input_ids: Optional[torch.Tensor] = None,1182            position_ids: Optional[torch.Tensor] = None,1183            attention_mask: Optional[torch.Tensor] = None,1184            past_key_values: Optional[Tuple[torch.FloatTensor]] = None,1185            inputs_embeds: Optional[torch.Tensor] = None,1186            labels: Optional[torch.Tensor] = None,1187            use_cache: Optional[bool] = None,1188            output_attentions: Optional[bool] = None,1189            output_hidden_states: Optional[bool] = None,1190            return_dict: Optional[bool] = None,1191            return_last_logit: Optional[bool] = False,1192    ):1193        use_cache = use_cache if use_cache is not None else self.config.use_cache1194        return_dict = return_dict if return_dict is not None else self.config.use_return_dict1195 1196        transformer_outputs = self.transformer(1197            input_ids=input_ids,1198            position_ids=position_ids,1199            attention_mask=attention_mask,1200            past_key_values=past_key_values,

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