Team Ai
Apppublic

hugging-apps/echo-memory

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
kolors_text_encoder.py1552 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 77 78class W8A16Linear(torch.autograd.Function):79    @staticmethod80    def forward(ctx, inp: torch.Tensor, quant_w: torch.Tensor, scale_w: torch.Tensor, weight_bit_width):81        ctx.inp_shape = inp.size()82        ctx.weight_bit_width = weight_bit_width83        out_features = quant_w.size(0)84        inp = inp.contiguous().view(-1, inp.size(-1))85        weight = extract_weight_to_half(quant_w, scale_w, weight_bit_width)86        ctx.weight_shape = weight.size()87        output = inp.mm(weight.t())88        ctx.save_for_backward(inp, quant_w, scale_w)89        return output.view(*(ctx.inp_shape[:-1] + (out_features,)))90 91    @staticmethod92    def backward(ctx, grad_output: torch.Tensor):93        inp, quant_w, scale_w = ctx.saved_tensors94        weight = extract_weight_to_half(quant_w, scale_w, ctx.weight_bit_width)95        grad_output = grad_output.contiguous().view(-1, weight.size(0))96        grad_input = grad_output.mm(weight)97        grad_weight = grad_output.t().mm(inp)98        return grad_input.view(ctx.inp_shape), grad_weight.view(ctx.weight_shape), None, None99 100 101def compress_int4_weight(weight: torch.Tensor):  # (n, m)102    with torch.cuda.device(weight.device):103        n, m = weight.size(0), weight.size(1)104        assert m % 2 == 0105        m = m // 2106        out = torch.empty(n, m, dtype=torch.int8, device="cuda")107        stream = torch.cuda.current_stream()108 109        gridDim = (n, 1, 1)110        blockDim = (min(round_up(m, 32), 1024), 1, 1)111 112        kernels.int4WeightCompression(113            gridDim,114            blockDim,115            0,116            stream,117            [ctypes.c_void_p(weight.data_ptr()), ctypes.c_void_p(out.data_ptr()), ctypes.c_int32(n), ctypes.c_int32(m)],118        )119        return out120 121 122def extract_weight_to_half(weight: torch.Tensor, scale_list: torch.Tensor, source_bit_width: int):123    assert scale_list.dtype in [torch.half, torch.bfloat16]124    assert weight.dtype in [torch.int8]125    if source_bit_width == 8:126        return weight.to(scale_list.dtype) * scale_list[:, None]127    elif source_bit_width == 4:128        func = (129            kernels.int4WeightExtractionHalf if scale_list.dtype == torch.half else kernels.int4WeightExtractionBFloat16130        )131    else:132        assert False, "Unsupported bit-width"133 134    with torch.cuda.device(weight.device):135        n, m = weight.size(0), weight.size(1)136        out = torch.empty(n, m * (8 // source_bit_width), dtype=scale_list.dtype, device="cuda")137        stream = torch.cuda.current_stream()138 139        gridDim = (n, 1, 1)140        blockDim = (min(round_up(m, 32), 1024), 1, 1)141 142        func(143            gridDim,144            blockDim,145            0,146            stream,147            [148                ctypes.c_void_p(weight.data_ptr()),149                ctypes.c_void_p(scale_list.data_ptr()),150                ctypes.c_void_p(out.data_ptr()),151                ctypes.c_int32(n),152                ctypes.c_int32(m),153            ],154        )155        return out156 157 158class QuantizedLinear(torch.nn.Module):159    def __init__(self, weight_bit_width: int, weight, bias=None, device="cuda", dtype=None, empty_init=False):160        super().__init__()161        weight = weight.to(device)  # ensure the weight is on the cuda device162        assert str(weight.device).startswith(163            'cuda'), 'The weights that need to be quantified should be on the CUDA device'164        self.weight_bit_width = weight_bit_width165        shape = weight.shape166 167        if weight is None or empty_init:168            self.weight = torch.empty(shape[0], shape[1] * weight_bit_width // 8, dtype=torch.int8, device=device)169            self.weight_scale = torch.empty(shape[0], dtype=dtype, device=device)170        else:171            self.weight_scale = weight.abs().max(dim=-1).values / ((2 ** (weight_bit_width - 1)) - 1)172            self.weight = torch.round(weight / self.weight_scale[:, None]).to(torch.int8)173            if weight_bit_width == 4:174                self.weight = compress_int4_weight(self.weight)175 176        self.weight = Parameter(self.weight.to(device), requires_grad=False)177        self.weight_scale = Parameter(self.weight_scale.to(device), requires_grad=False)178        self.bias = Parameter(bias.to(device), requires_grad=False) if bias is not None else None179 180    def forward(self, input):181        output = W8A16Linear.apply(input, self.weight, self.weight_scale, self.weight_bit_width)182        if self.bias is not None:183            output = output + self.bias184        return output185 186 187def quantize(model, weight_bit_width, empty_init=False, device=None):188    """Replace fp16 linear with quantized linear"""189    for layer in model.layers:190        layer.self_attention.query_key_value = QuantizedLinear(191            weight_bit_width=weight_bit_width,192            weight=layer.self_attention.query_key_value.weight,193            bias=layer.self_attention.query_key_value.bias,194            dtype=layer.self_attention.query_key_value.weight.dtype,195            device=layer.self_attention.query_key_value.weight.device if device is None else device,196            empty_init=empty_init197        )198        layer.self_attention.dense = QuantizedLinear(199            weight_bit_width=weight_bit_width,200            weight=layer.self_attention.dense.weight,201            bias=layer.self_attention.dense.bias,202            dtype=layer.self_attention.dense.weight.dtype,203            device=layer.self_attention.dense.weight.device if device is None else device,204            empty_init=empty_init205        )206        layer.mlp.dense_h_to_4h = QuantizedLinear(207            weight_bit_width=weight_bit_width,208            weight=layer.mlp.dense_h_to_4h.weight,209            bias=layer.mlp.dense_h_to_4h.bias,210            dtype=layer.mlp.dense_h_to_4h.weight.dtype,211            device=layer.mlp.dense_h_to_4h.weight.device if device is None else device,212            empty_init=empty_init213        )214        layer.mlp.dense_4h_to_h = QuantizedLinear(215            weight_bit_width=weight_bit_width,216            weight=layer.mlp.dense_4h_to_h.weight,217            bias=layer.mlp.dense_4h_to_h.bias,218            dtype=layer.mlp.dense_4h_to_h.weight.dtype,219            device=layer.mlp.dense_4h_to_h.weight.device if device is None else device,220            empty_init=empty_init221        )222 223    return model224 225 226 227class ChatGLMConfig(PretrainedConfig):228    model_type = "chatglm"229    def __init__(230        self,231        num_layers=28,232        padded_vocab_size=65024,233        hidden_size=4096,234        ffn_hidden_size=13696,235        kv_channels=128,236        num_attention_heads=32,237        seq_length=2048,238        hidden_dropout=0.0,239        classifier_dropout=None,240        attention_dropout=0.0,241        layernorm_epsilon=1e-5,242        rmsnorm=True,243        apply_residual_connection_post_layernorm=False,244        post_layer_norm=True,245        add_bias_linear=False,246        add_qkv_bias=False,247        bias_dropout_fusion=True,248        multi_query_attention=False,249        multi_query_group_num=1,250        apply_query_key_layer_scaling=True,251        attention_softmax_in_fp32=True,252        fp32_residual_connection=False,253        quantization_bit=0,254        pre_seq_len=None,255        prefix_projection=False,256        **kwargs257    ):258        self.num_layers = num_layers259        self.vocab_size = padded_vocab_size260        self.padded_vocab_size = padded_vocab_size261        self.hidden_size = hidden_size262        self.ffn_hidden_size = ffn_hidden_size263        self.kv_channels = kv_channels264        self.num_attention_heads = num_attention_heads265        self.seq_length = seq_length266        self.hidden_dropout = hidden_dropout267        self.classifier_dropout = classifier_dropout268        self.attention_dropout = attention_dropout269        self.layernorm_epsilon = layernorm_epsilon270        self.rmsnorm = rmsnorm271        self.apply_residual_connection_post_layernorm = apply_residual_connection_post_layernorm272        self.post_layer_norm = post_layer_norm273        self.add_bias_linear = add_bias_linear274        self.add_qkv_bias = add_qkv_bias275        self.bias_dropout_fusion = bias_dropout_fusion276        self.multi_query_attention = multi_query_attention277        self.multi_query_group_num = multi_query_group_num278        self.apply_query_key_layer_scaling = apply_query_key_layer_scaling279        self.attention_softmax_in_fp32 = attention_softmax_in_fp32280        self.fp32_residual_connection = fp32_residual_connection281        self.quantization_bit = quantization_bit282        self.pre_seq_len = pre_seq_len283        self.prefix_projection = prefix_projection284        super().__init__(**kwargs)285 286 287 288# flags required to enable jit fusion kernels289 290if sys.platform != 'darwin':291    torch._C._jit_set_profiling_mode(False)292    torch._C._jit_set_profiling_executor(False)293    torch._C._jit_override_can_fuse_on_cpu(True)294    torch._C._jit_override_can_fuse_on_gpu(True)295 296logger = logging.get_logger(__name__)297 298_CHECKPOINT_FOR_DOC = "THUDM/ChatGLM"299_CONFIG_FOR_DOC = "ChatGLM6BConfig"300 301CHATGLM_6B_PRETRAINED_MODEL_ARCHIVE_LIST = [302    "THUDM/chatglm3-6b-base",303    # See all ChatGLM models at https://huggingface.co/models?filter=chatglm304]305 306 307def default_init(cls, *args, **kwargs):308    return cls(*args, **kwargs)309 310 311class InvalidScoreLogitsProcessor(LogitsProcessor):312    def __call__(self, input_ids: torch.LongTensor, scores: torch.FloatTensor) -> torch.FloatTensor:313        if torch.isnan(scores).any() or torch.isinf(scores).any():314            scores.zero_()315            scores[..., 5] = 5e4316        return scores317 318 319class PrefixEncoder(torch.nn.Module):320    """321    The torch.nn model to encode the prefix322    Input shape: (batch-size, prefix-length)323    Output shape: (batch-size, prefix-length, 2*layers*hidden)324    """325 326    def __init__(self, config: ChatGLMConfig):327        super().__init__()328        self.prefix_projection = config.prefix_projection329        if self.prefix_projection:330            # Use a two-layer MLP to encode the prefix331            kv_size = config.num_layers * config.kv_channels * config.multi_query_group_num * 2332            self.embedding = torch.nn.Embedding(config.pre_seq_len, kv_size)333            self.trans = torch.nn.Sequential(334                torch.nn.Linear(kv_size, config.hidden_size),335                torch.nn.Tanh(),336                torch.nn.Linear(config.hidden_size, kv_size)337            )338        else:339            self.embedding = torch.nn.Embedding(config.pre_seq_len,340                                                config.num_layers * config.kv_channels * config.multi_query_group_num * 2)341 342    def forward(self, prefix: torch.Tensor):343        if self.prefix_projection:344            prefix_tokens = self.embedding(prefix)345            past_key_values = self.trans(prefix_tokens)346        else:347            past_key_values = self.embedding(prefix)348        return past_key_values349 350 351def split_tensor_along_last_dim(352        tensor: torch.Tensor,353        num_partitions: int,354        contiguous_split_chunks: bool = False,355) -> List[torch.Tensor]:356    """Split a tensor along its last dimension.357 358    Arguments:359        tensor: input tensor.360        num_partitions: number of partitions to split the tensor361        contiguous_split_chunks: If True, make each chunk contiguous362                                 in memory.363 364    Returns:365        A list of Tensors366    """367    # Get the size and dimension.368    last_dim = tensor.dim() - 1369    last_dim_size = tensor.size()[last_dim] // num_partitions370    # Split.371    tensor_list = torch.split(tensor, last_dim_size, dim=last_dim)372    # Note: torch.split does not create contiguous tensors by default.373    if contiguous_split_chunks:374        return tuple(chunk.contiguous() for chunk in tensor_list)375 376    return tensor_list377 378 379class RotaryEmbedding(nn.Module):380    def __init__(self, dim, original_impl=False, device=None, dtype=None):381        super().__init__()382        inv_freq = 1.0 / (10000 ** (torch.arange(0, dim, 2, device=device).to(dtype=dtype) / dim))383        self.register_buffer("inv_freq", inv_freq)384        self.dim = dim385        self.original_impl = original_impl386 387    def forward_impl(388            self, seq_len: int, n_elem: int, dtype: torch.dtype, device: torch.device, base: int = 10000389    ):390        """Enhanced Transformer with Rotary Position Embedding.391 392        Derived from: https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/master/labml_nn/393        transformers/rope/__init__.py. MIT License:394        https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/master/license.395        """396        # $\Theta = {\theta_i = 10000^{\frac{2(i-1)}{d}}, i \in [1, 2, ..., \frac{d}{2}]}$397        theta = 1.0 / (base ** (torch.arange(0, n_elem, 2, dtype=torch.float, device=device) / n_elem))398 399        # Create position indexes `[0, 1, ..., seq_len - 1]`400        seq_idx = torch.arange(seq_len, dtype=torch.float, device=device)401 402        # Calculate the product of position index and $\theta_i$403        idx_theta = torch.outer(seq_idx, theta).float()404 405        cache = torch.stack([torch.cos(idx_theta), torch.sin(idx_theta)], dim=-1)406 407        # this is to mimic the behaviour of complex32, else we will get different results408        if dtype in (torch.float16, torch.bfloat16, torch.int8):409            cache = cache.bfloat16() if dtype == torch.bfloat16 else cache.half()410        return cache411 412    def forward(self, max_seq_len, offset=0):413        return self.forward_impl(414            max_seq_len, self.dim, dtype=self.inv_freq.dtype, device=self.inv_freq.device415        )416 417 418@torch.jit.script419def apply_rotary_pos_emb(x: torch.Tensor, rope_cache: torch.Tensor) -> torch.Tensor:420    # x: [sq, b, np, hn]421    sq, b, np, hn = x.size(0), x.size(1), x.size(2), x.size(3)422    rot_dim = rope_cache.shape[-2] * 2423    x, x_pass = x[..., :rot_dim], x[..., rot_dim:]424    # truncate to support variable sizes425    rope_cache = rope_cache[:sq]426    xshaped = x.reshape(sq, -1, np, rot_dim // 2, 2)427    rope_cache = rope_cache.view(sq, -1, 1, xshaped.size(3), 2)428    x_out2 = torch.stack(429        [430            xshaped[..., 0] * rope_cache[..., 0] - xshaped[..., 1] * rope_cache[..., 1],431            xshaped[..., 1] * rope_cache[..., 0] + xshaped[..., 0] * rope_cache[..., 1],432        ],433        -1,434    )435    x_out2 = x_out2.flatten(3)436    return torch.cat((x_out2, x_pass), dim=-1)437 438 439class RMSNorm(torch.nn.Module):440    def __init__(self, normalized_shape, eps=1e-5, device=None, dtype=None, **kwargs):441        super().__init__()442        self.weight = torch.nn.Parameter(torch.empty(normalized_shape, device=device, dtype=dtype))443        self.eps = eps444 445    def forward(self, hidden_states: torch.Tensor):446        input_dtype = hidden_states.dtype447        variance = hidden_states.to(torch.float32).pow(2).mean(-1, keepdim=True)448        hidden_states = hidden_states * torch.rsqrt(variance + self.eps)449 450        return (self.weight * hidden_states).to(input_dtype)451 452 453class CoreAttention(torch.nn.Module):454    def __init__(self, config: ChatGLMConfig, layer_number):455        super(CoreAttention, self).__init__()456 457        self.apply_query_key_layer_scaling = config.apply_query_key_layer_scaling458        self.attention_softmax_in_fp32 = config.attention_softmax_in_fp32459        if self.apply_query_key_layer_scaling:460            self.attention_softmax_in_fp32 = True461        self.layer_number = max(1, layer_number)462 463        projection_size = config.kv_channels * config.num_attention_heads464 465        # Per attention head and per partition values.466        self.hidden_size_per_partition = projection_size467        self.hidden_size_per_attention_head = projection_size // config.num_attention_heads468        self.num_attention_heads_per_partition = config.num_attention_heads469 470        coeff = None471        self.norm_factor = math.sqrt(self.hidden_size_per_attention_head)472        if self.apply_query_key_layer_scaling:473            coeff = self.layer_number474            self.norm_factor *= coeff475        self.coeff = coeff476 477        self.attention_dropout = torch.nn.Dropout(config.attention_dropout)478 479    def forward(self, query_layer, key_layer, value_layer, attention_mask):480        pytorch_major_version = int(torch.__version__.split('.')[0])481        if pytorch_major_version >= 2:482            query_layer, key_layer, value_layer = [k.permute(1, 2, 0, 3) for k in [query_layer, key_layer, value_layer]]483            if attention_mask is None and query_layer.shape[2] == key_layer.shape[2]:484                context_layer = torch.nn.functional.scaled_dot_product_attention(query_layer, key_layer, value_layer,485                                                                                 is_causal=True)486            else:487                if attention_mask is not None:488                    attention_mask = ~attention_mask489                context_layer = torch.nn.functional.scaled_dot_product_attention(query_layer, key_layer, value_layer,490                                                                                 attention_mask)491            context_layer = context_layer.permute(2, 0, 1, 3)492            new_context_layer_shape = context_layer.size()[:-2] + (self.hidden_size_per_partition,)493            context_layer = context_layer.reshape(*new_context_layer_shape)494        else:495            # Raw attention scores496 497            # [b, np, sq, sk]498            output_size = (query_layer.size(1), query_layer.size(2), query_layer.size(0), key_layer.size(0))499 500            # [sq, b, np, hn] -> [sq, b * np, hn]501            query_layer = query_layer.view(output_size[2], output_size[0] * output_size[1], -1)502            # [sk, b, np, hn] -> [sk, b * np, hn]503            key_layer = key_layer.view(output_size[3], output_size[0] * output_size[1], -1)504 505            # preallocting input tensor: [b * np, sq, sk]506            matmul_input_buffer = torch.empty(507                output_size[0] * output_size[1], output_size[2], output_size[3], dtype=query_layer.dtype,508                device=query_layer.device509            )510 511            # Raw attention scores. [b * np, sq, sk]512            matmul_result = torch.baddbmm(513                matmul_input_buffer,514                query_layer.transpose(0, 1),  # [b * np, sq, hn]515                key_layer.transpose(0, 1).transpose(1, 2),  # [b * np, hn, sk]516                beta=0.0,517                alpha=(1.0 / self.norm_factor),518            )519 520            # change view to [b, np, sq, sk]521            attention_scores = matmul_result.view(*output_size)522 523            # ===========================524            # Attention probs and dropout525            # ===========================526 527            # attention scores and attention mask [b, np, sq, sk]528            if self.attention_softmax_in_fp32:529                attention_scores = attention_scores.float()530            if self.coeff is not None:531                attention_scores = attention_scores * self.coeff532            if attention_mask is None and attention_scores.shape[2] == attention_scores.shape[3]:533                attention_mask = torch.ones(output_size[0], 1, output_size[2], output_size[3],534                                            device=attention_scores.device, dtype=torch.bool)535                attention_mask.tril_()536                attention_mask = ~attention_mask537            if attention_mask is not None:538                attention_scores = attention_scores.masked_fill(attention_mask, float("-inf"))539            attention_probs = F.softmax(attention_scores, dim=-1)540            attention_probs = attention_probs.type_as(value_layer)541 542            # This is actually dropping out entire tokens to attend to, which might543            # seem a bit unusual, but is taken from the original Transformer paper.544            attention_probs = self.attention_dropout(attention_probs)545            # =========================546            # Context layer. [sq, b, hp]547            # =========================548 549            # value_layer -> context layer.550            # [sk, b, np, hn] --> [b, np, sq, hn]551 552            # context layer shape: [b, np, sq, hn]553            output_size = (value_layer.size(1), value_layer.size(2), query_layer.size(0), value_layer.size(3))554            # change view [sk, b * np, hn]555            value_layer = value_layer.view(value_layer.size(0), output_size[0] * output_size[1], -1)556            # change view [b * np, sq, sk]557            attention_probs = attention_probs.view(output_size[0] * output_size[1], output_size[2], -1)558            # matmul: [b * np, sq, hn]559            context_layer = torch.bmm(attention_probs, value_layer.transpose(0, 1))560            # change view [b, np, sq, hn]561            context_layer = context_layer.view(*output_size)562            # [b, np, sq, hn] --> [sq, b, np, hn]563            context_layer = context_layer.permute(2, 0, 1, 3).contiguous()564            # [sq, b, np, hn] --> [sq, b, hp]565            new_context_layer_shape = context_layer.size()[:-2] + (self.hidden_size_per_partition,)566            context_layer = context_layer.view(*new_context_layer_shape)567 568        return context_layer569 570 571class SelfAttention(torch.nn.Module):572    """Parallel self-attention layer abstract class.573 574    Self-attention layer takes input with size [s, b, h]575    and returns output of the same size.576    """577 578    def __init__(self, config: ChatGLMConfig, layer_number, device=None):579        super(SelfAttention, self).__init__()580        self.layer_number = max(1, layer_number)581 582        self.projection_size = config.kv_channels * config.num_attention_heads583 584        # Per attention head and per partition values.585        self.hidden_size_per_attention_head = self.projection_size // config.num_attention_heads586        self.num_attention_heads_per_partition = config.num_attention_heads587 588        self.multi_query_attention = config.multi_query_attention589        self.qkv_hidden_size = 3 * self.projection_size590        if self.multi_query_attention:591            self.num_multi_query_groups_per_partition = config.multi_query_group_num592            self.qkv_hidden_size = (593                    self.projection_size + 2 * self.hidden_size_per_attention_head * config.multi_query_group_num594            )595        self.query_key_value = nn.Linear(config.hidden_size, self.qkv_hidden_size,596                                         bias=config.add_bias_linear or config.add_qkv_bias,597                                         device=device, **_config_to_kwargs(config)598                                         )599 600        self.core_attention = CoreAttention(config, self.layer_number)601 602        # Output.603        self.dense = nn.Linear(self.projection_size, config.hidden_size, bias=config.add_bias_linear,604                               device=device, **_config_to_kwargs(config)605                               )606 607    def _allocate_memory(self, inference_max_sequence_len, batch_size, device=None, dtype=None):608        if self.multi_query_attention:609            num_attention_heads = self.num_multi_query_groups_per_partition610        else:611            num_attention_heads = self.num_attention_heads_per_partition612        return torch.empty(613            inference_max_sequence_len,614            batch_size,615            num_attention_heads,616            self.hidden_size_per_attention_head,617            dtype=dtype,618            device=device,619        )620 621    def forward(622            self, hidden_states, attention_mask, rotary_pos_emb, kv_cache=None, use_cache=True623    ):624        # hidden_states: [sq, b, h]625 626        # =================================================627        # Pre-allocate memory for key-values for inference.628        # =================================================629        # =====================630        # Query, Key, and Value631        # =====================632 633        # Attention heads [sq, b, h] --> [sq, b, (np * 3 * hn)]634        mixed_x_layer = self.query_key_value(hidden_states)635 636        if self.multi_query_attention:637            (query_layer, key_layer, value_layer) = mixed_x_layer.split(638                [639                    self.num_attention_heads_per_partition * self.hidden_size_per_attention_head,640                    self.num_multi_query_groups_per_partition * self.hidden_size_per_attention_head,641                    self.num_multi_query_groups_per_partition * self.hidden_size_per_attention_head,642                ],643                dim=-1,644            )645            query_layer = query_layer.view(646                query_layer.size()[:-1] + (self.num_attention_heads_per_partition, self.hidden_size_per_attention_head)647            )648            key_layer = key_layer.view(649                key_layer.size()[:-1] + (self.num_multi_query_groups_per_partition, self.hidden_size_per_attention_head)650            )651            value_layer = value_layer.view(652                value_layer.size()[:-1]653                + (self.num_multi_query_groups_per_partition, self.hidden_size_per_attention_head)654            )655        else:656            new_tensor_shape = mixed_x_layer.size()[:-1] + \657                               (self.num_attention_heads_per_partition,658                                3 * self.hidden_size_per_attention_head)659            mixed_x_layer = mixed_x_layer.view(*new_tensor_shape)660 661            # [sq, b, np, 3 * hn] --> 3 [sq, b, np, hn]662            (query_layer, key_layer, value_layer) = split_tensor_along_last_dim(mixed_x_layer, 3)663 664        # apply relative positional encoding (rotary embedding)665        if rotary_pos_emb is not None:666            query_layer = apply_rotary_pos_emb(query_layer, rotary_pos_emb)667            key_layer = apply_rotary_pos_emb(key_layer, rotary_pos_emb)668 669        # adjust key and value for inference670        if kv_cache is not None:671            cache_k, cache_v = kv_cache672            key_layer = torch.cat((cache_k, key_layer), dim=0)673            value_layer = torch.cat((cache_v, value_layer), dim=0)674        if use_cache:675            kv_cache = (key_layer, value_layer)676        else:677            kv_cache = None678 679        if self.multi_query_attention:680            key_layer = key_layer.unsqueeze(-2)681            key_layer = key_layer.expand(682                -1, -1, -1, self.num_attention_heads_per_partition // self.num_multi_query_groups_per_partition, -1683            )684            key_layer = key_layer.contiguous().view(685                key_layer.size()[:2] + (self.num_attention_heads_per_partition, self.hidden_size_per_attention_head)686            )687            value_layer = value_layer.unsqueeze(-2)688            value_layer = value_layer.expand(689                -1, -1, -1, self.num_attention_heads_per_partition // self.num_multi_query_groups_per_partition, -1690            )691            value_layer = value_layer.contiguous().view(692                value_layer.size()[:2] + (self.num_attention_heads_per_partition, self.hidden_size_per_attention_head)693            )694 695        # ==================================696        # core attention computation697        # ==================================698 699        context_layer = self.core_attention(query_layer, key_layer, value_layer, attention_mask)700 701        # =================702        # Output. [sq, b, h]703        # =================704 705        output = self.dense(context_layer)706 707        return output, kv_cache708 709 710def _config_to_kwargs(args):711    common_kwargs = {712        "dtype": args.torch_dtype,713    }714    return common_kwargs715 716 717class MLP(torch.nn.Module):718    """MLP.719 720    MLP will take the input with h hidden state, project it to 4*h721    hidden dimension, perform nonlinear transformation, and project the722    state back into h hidden dimension.723    """724 725    def __init__(self, config: ChatGLMConfig, device=None):726        super(MLP, self).__init__()727 728        self.add_bias = config.add_bias_linear729 730        # Project to 4h. If using swiglu double the output width, see https://arxiv.org/pdf/2002.05202.pdf731        self.dense_h_to_4h = nn.Linear(732            config.hidden_size,733            config.ffn_hidden_size * 2,734            bias=self.add_bias,735            device=device,736            **_config_to_kwargs(config)737        )738 739        def swiglu(x):740            x = torch.chunk(x, 2, dim=-1)741            return F.silu(x[0]) * x[1]742 743        self.activation_func = swiglu744 745        # Project back to h.746        self.dense_4h_to_h = nn.Linear(747            config.ffn_hidden_size,748            config.hidden_size,749            bias=self.add_bias,750            device=device,751            **_config_to_kwargs(config)752        )753 754    def forward(self, hidden_states):755        # [s, b, 4hp]756        intermediate_parallel = self.dense_h_to_4h(hidden_states)757        intermediate_parallel = self.activation_func(intermediate_parallel)758        # [s, b, h]759        output = self.dense_4h_to_h(intermediate_parallel)760        return output761 762 763class GLMBlock(torch.nn.Module):764    """A single transformer layer.765 766    Transformer layer takes input with size [s, b, h] and returns an767    output of the same size.768    """769 770    def __init__(self, config: ChatGLMConfig, layer_number, device=None):771        super(GLMBlock, self).__init__()772        self.layer_number = layer_number773 774        self.apply_residual_connection_post_layernorm = config.apply_residual_connection_post_layernorm775 776        self.fp32_residual_connection = config.fp32_residual_connection777 778        LayerNormFunc = RMSNorm if config.rmsnorm else LayerNorm779        # Layernorm on the input data.780        self.input_layernorm = LayerNormFunc(config.hidden_size, eps=config.layernorm_epsilon, device=device,781                                             dtype=config.torch_dtype)782 783        # Self attention.784        self.self_attention = SelfAttention(config, layer_number, device=device)785        self.hidden_dropout = config.hidden_dropout786 787        # Layernorm on the attention output788        self.post_attention_layernorm = LayerNormFunc(config.hidden_size, eps=config.layernorm_epsilon, device=device,789                                                      dtype=config.torch_dtype)790 791        # MLP792        self.mlp = MLP(config, device=device)793 794    def forward(795            self, hidden_states, attention_mask, rotary_pos_emb, kv_cache=None, use_cache=True,796    ):797        # hidden_states: [s, b, h]798 799        # Layer norm at the beginning of the transformer layer.800        layernorm_output = self.input_layernorm(hidden_states)801        # Self attention.802        attention_output, kv_cache = self.self_attention(803            layernorm_output,804            attention_mask,805            rotary_pos_emb,806            kv_cache=kv_cache,807            use_cache=use_cache808        )809 810        # Residual connection.811        if self.apply_residual_connection_post_layernorm:812            residual = layernorm_output813        else:814            residual = hidden_states815 816        layernorm_input = torch.nn.functional.dropout(attention_output, p=self.hidden_dropout, training=self.training)817        layernorm_input = residual + layernorm_input818 819        # Layer norm post the self attention.820        layernorm_output = self.post_attention_layernorm(layernorm_input)821 822        # MLP.823        mlp_output = self.mlp(layernorm_output)824 825        # Second residual connection.826        if self.apply_residual_connection_post_layernorm:827            residual = layernorm_output828        else:829            residual = layernorm_input830 831        output = torch.nn.functional.dropout(mlp_output, p=self.hidden_dropout, training=self.training)832        output = residual + output833 834        return output, kv_cache835 836 837class GLMTransformer(torch.nn.Module):838    """Transformer class."""839 840    def __init__(self, config: ChatGLMConfig, device=None):841        super(GLMTransformer, self).__init__()842 843        self.fp32_residual_connection = config.fp32_residual_connection844        self.post_layer_norm = config.post_layer_norm845 846        # Number of layers.847        self.num_layers = config.num_layers848 849        # Transformer layers.850        def build_layer(layer_number):851            return GLMBlock(config, layer_number, device=device)852 853        self.layers = torch.nn.ModuleList([build_layer(i + 1) for i in range(self.num_layers)])854 855        if self.post_layer_norm:856            LayerNormFunc = RMSNorm if config.rmsnorm else LayerNorm857            # Final layer norm before output.858            self.final_layernorm = LayerNormFunc(config.hidden_size, eps=config.layernorm_epsilon, device=device,859                                                 dtype=config.torch_dtype)860 861        self.gradient_checkpointing = False862 863    def _get_layer(self, layer_number):864        return self.layers[layer_number]865 866    def forward(867            self, hidden_states, attention_mask, rotary_pos_emb, kv_caches=None,868            use_cache: Optional[bool] = True,869            output_hidden_states: Optional[bool] = False,870    ):871        if not kv_caches:872            kv_caches = [None for _ in range(self.num_layers)]873        presents = () if use_cache else None874        if self.gradient_checkpointing and self.training:875            if use_cache:876                logger.warning_once(877                    "`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`..."878                )879                use_cache = False880 881        all_self_attentions = None882        all_hidden_states = () if output_hidden_states else None883        for index in range(self.num_layers):884            if output_hidden_states:885                all_hidden_states = all_hidden_states + (hidden_states,)886 887            layer = self._get_layer(index)888            if self.gradient_checkpointing and self.training:889                layer_ret = torch.utils.checkpoint.checkpoint(890                    layer,891                    hidden_states,892                    attention_mask,893                    rotary_pos_emb,894                    kv_caches[index],895                    use_cache896                )897            else:898                layer_ret = layer(899                    hidden_states,900                    attention_mask,901                    rotary_pos_emb,902                    kv_cache=kv_caches[index],903                    use_cache=use_cache904                )905            hidden_states, kv_cache = layer_ret906            if use_cache:907                presents = presents + (kv_cache,)908 909        if output_hidden_states:910            all_hidden_states = all_hidden_states + (hidden_states,)911 912        # Final layer norm.913        if self.post_layer_norm:914            hidden_states = self.final_layernorm(hidden_states)915 916        return hidden_states, presents, all_hidden_states, all_self_attentions917 918 919class ChatGLMPreTrainedModel(PreTrainedModel):920    """921    An abstract class to handle weights initialization and922    a simple interface for downloading and loading pretrained models.923    """924 925    is_parallelizable = False926    supports_gradient_checkpointing = True927    config_class = ChatGLMConfig928    base_model_prefix = "transformer"929    _no_split_modules = ["GLMBlock"]930 931    def _init_weights(self, module: nn.Module):932        """Initialize the weights."""933        return934 935    def get_masks(self, input_ids, past_key_values, padding_mask=None):936        batch_size, seq_length = input_ids.shape937        full_attention_mask = torch.ones(batch_size, seq_length, seq_length, device=input_ids.device)938        full_attention_mask.tril_()939        past_length = 0940        if past_key_values:941            past_length = past_key_values[0][0].shape[0]942        if past_length:943            full_attention_mask = torch.cat((torch.ones(batch_size, seq_length, past_length,944                                                        device=input_ids.device), full_attention_mask), dim=-1)945        if padding_mask is not None:946            full_attention_mask = full_attention_mask * padding_mask.unsqueeze(1)947        if not past_length and padding_mask is not None:948            full_attention_mask -= padding_mask.unsqueeze(-1) - 1949        full_attention_mask = (full_attention_mask < 0.5).bool()950        full_attention_mask.unsqueeze_(1)951        return full_attention_mask952 953    def get_position_ids(self, input_ids, device):954        batch_size, seq_length = input_ids.shape955        position_ids = torch.arange(seq_length, dtype=torch.long, device=device).unsqueeze(0).repeat(batch_size, 1)956        return position_ids957 958    def _set_gradient_checkpointing(self, module, value=False):959        if isinstance(module, GLMTransformer):960            module.gradient_checkpointing = value961 962 963class Embedding(torch.nn.Module):964    """Language model embeddings."""965 966    def __init__(self, config: ChatGLMConfig, device=None):967        super(Embedding, self).__init__()968 969        self.hidden_size = config.hidden_size970        # Word embeddings (parallel).971        self.word_embeddings = nn.Embedding(972            config.padded_vocab_size,973            self.hidden_size,974            dtype=config.torch_dtype,975            device=device976        )977        self.fp32_residual_connection = config.fp32_residual_connection978 979    def forward(self, input_ids):980        # Embeddings.981        words_embeddings = self.word_embeddings(input_ids)982        embeddings = words_embeddings983        # Data format change to avoid explicit transposes : [b s h] --> [s b h].984        embeddings = embeddings.transpose(0, 1).contiguous()985        # If the input flag for fp32 residual connection is set, convert for float.986        if self.fp32_residual_connection:987            embeddings = embeddings.float()988        return embeddings989 990 991class ChatGLMModel(ChatGLMPreTrainedModel):992    def __init__(self, config: ChatGLMConfig, device=None, empty_init=True):993        super().__init__(config)994        if empty_init:995            init_method = skip_init996        else:997            init_method = default_init998        init_kwargs = {}999        if device is not None:1000            init_kwargs["device"] = device1001        self.embedding = init_method(Embedding, config, **init_kwargs)1002        self.num_layers = config.num_layers1003        self.multi_query_group_num = config.multi_query_group_num1004        self.kv_channels = config.kv_channels1005 1006        # Rotary positional embeddings1007        self.seq_length = config.seq_length1008        rotary_dim = (1009            config.hidden_size // config.num_attention_heads if config.kv_channels is None else config.kv_channels1010        )1011 1012        self.rotary_pos_emb = RotaryEmbedding(rotary_dim // 2, original_impl=config.original_rope, device=device,1013                                              dtype=config.torch_dtype)1014        self.encoder = init_method(GLMTransformer, config, **init_kwargs)1015        self.output_layer = init_method(nn.Linear, config.hidden_size, config.padded_vocab_size, bias=False,1016                                        dtype=config.torch_dtype, **init_kwargs)1017        self.pre_seq_len = config.pre_seq_len1018        self.prefix_projection = config.prefix_projection1019        if self.pre_seq_len is not None:1020            for param in self.parameters():1021                param.requires_grad = False1022            self.prefix_tokens = torch.arange(self.pre_seq_len).long()1023            self.prefix_encoder = PrefixEncoder(config)1024            self.dropout = torch.nn.Dropout(0.1)1025 1026    def get_input_embeddings(self):1027        return self.embedding.word_embeddings1028 1029    def get_prompt(self, batch_size, device, dtype=torch.half):1030        prefix_tokens = self.prefix_tokens.unsqueeze(0).expand(batch_size, -1).to(device)1031        past_key_values = self.prefix_encoder(prefix_tokens).type(dtype)1032        past_key_values = past_key_values.view(1033            batch_size,1034            self.pre_seq_len,1035            self.num_layers * 2,1036            self.multi_query_group_num,1037            self.kv_channels1038        )1039        # seq_len, b, nh, hidden_size1040        past_key_values = self.dropout(past_key_values)1041        past_key_values = past_key_values.permute([2, 1, 0, 3, 4]).split(2)1042        return past_key_values1043 1044    def forward(1045            self,1046            input_ids,1047            position_ids: Optional[torch.Tensor] = None,1048            attention_mask: Optional[torch.BoolTensor] = None,1049            full_attention_mask: Optional[torch.BoolTensor] = None,1050            past_key_values: Optional[Tuple[Tuple[torch.Tensor, torch.Tensor], ...]] = None,1051            inputs_embeds: Optional[torch.Tensor] = None,1052            use_cache: Optional[bool] = None,1053            output_hidden_states: Optional[bool] = None,1054            return_dict: Optional[bool] = None,1055    ):1056        output_hidden_states = (1057            output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states1058        )1059        use_cache = use_cache if use_cache is not None else self.config.use_cache1060        return_dict = return_dict if return_dict is not None else self.config.use_return_dict1061 1062        batch_size, seq_length = input_ids.shape1063 1064        if inputs_embeds is None:1065            inputs_embeds = self.embedding(input_ids)1066 1067        if self.pre_seq_len is not None:1068            if past_key_values is None:1069                past_key_values = self.get_prompt(batch_size=batch_size, device=input_ids.device,1070                                                  dtype=inputs_embeds.dtype)1071            if attention_mask is not None:1072                attention_mask = torch.cat([attention_mask.new_ones((batch_size, self.pre_seq_len)),1073                                            attention_mask], dim=-1)1074 1075        if full_attention_mask is None:1076            if (attention_mask is not None and not attention_mask.all()) or (past_key_values and seq_length != 1):1077                full_attention_mask = self.get_masks(input_ids, past_key_values, padding_mask=attention_mask)1078 1079        # Rotary positional embeddings1080        rotary_pos_emb = self.rotary_pos_emb(self.seq_length)1081        if position_ids is not None:1082            rotary_pos_emb = rotary_pos_emb[position_ids]1083        else:1084            rotary_pos_emb = rotary_pos_emb[None, :seq_length]1085        rotary_pos_emb = rotary_pos_emb.transpose(0, 1).contiguous()1086 1087        # Run encoder.1088        hidden_states, presents, all_hidden_states, all_self_attentions = self.encoder(1089            inputs_embeds, full_attention_mask, rotary_pos_emb=rotary_pos_emb,1090            kv_caches=past_key_values, use_cache=use_cache, output_hidden_states=output_hidden_states1091        )1092 1093        if not return_dict:1094            return tuple(v for v in [hidden_states, presents, all_hidden_states, all_self_attentions] if v is not None)1095 1096        return BaseModelOutputWithPast(1097            last_hidden_state=hidden_states,1098            past_key_values=presents,1099            hidden_states=all_hidden_states,1100            attentions=all_self_attentions,1101        )1102 1103    def quantize(self, weight_bit_width: int):1104        # from .quantization import quantize1105        quantize(self.encoder, weight_bit_width)1106        return self1107 1108 1109class ChatGLMForConditionalGeneration(ChatGLMPreTrainedModel):1110    def __init__(self, config: ChatGLMConfig, empty_init=True, device=None):1111        super().__init__(config)1112 1113        self.max_sequence_length = config.max_length1114        self.transformer = ChatGLMModel(config, empty_init=empty_init, device=device)1115        self.config = config1116        self.quantized = False1117 1118        if self.config.quantization_bit:1119            self.quantize(self.config.quantization_bit, empty_init=True)1120 1121    def _update_model_kwargs_for_generation(1122            self,1123            outputs: ModelOutput,1124            model_kwargs: Dict[str, Any],1125            is_encoder_decoder: bool = False,1126            standardize_cache_format: bool = False,1127    ) -> Dict[str, Any]:1128        # update past_key_values1129        model_kwargs["past_key_values"] = self._extract_past_from_model_output(1130            outputs, standardize_cache_format=standardize_cache_format1131        )1132 1133        # update attention mask1134        if "attention_mask" in model_kwargs:1135            attention_mask = model_kwargs["attention_mask"]1136            model_kwargs["attention_mask"] = torch.cat(1137                [attention_mask, attention_mask.new_ones((attention_mask.shape[0], 1))], dim=-11138            )1139 1140        # update position ids1141        if "position_ids" in model_kwargs:1142            position_ids = model_kwargs["position_ids"]1143            new_position_id = position_ids[..., -1:].clone()1144            new_position_id += 11145            model_kwargs["position_ids"] = torch.cat(1146                [position_ids, new_position_id], dim=-11147            )1148 1149        model_kwargs["is_first_forward"] = False1150        return model_kwargs1151 1152    def prepare_inputs_for_generation(1153            self,1154            input_ids: torch.LongTensor,1155            past_key_values: Optional[torch.Tensor] = None,1156            attention_mask: Optional[torch.Tensor] = None,1157            position_ids: Optional[torch.Tensor] = None,1158            use_cache: Optional[bool] = None,1159            is_first_forward: bool = True,1160            **kwargs1161    ) -> dict:1162        # only last token for input_ids if past is not None1163        if position_ids is None:1164            position_ids = self.get_position_ids(input_ids, device=input_ids.device)1165        if not is_first_forward:1166            if past_key_values is not None:1167                position_ids = position_ids[..., -1:]1168                input_ids = input_ids[:, -1:]1169        return {1170            "input_ids": input_ids,1171            "past_key_values": past_key_values,1172            "position_ids": position_ids,1173            "attention_mask": attention_mask,1174            "return_last_logit": True,1175            "use_cache": use_cache1176        }1177 1178    def forward(1179            self,1180            input_ids: Optional[torch.Tensor] = None,1181            position_ids: Optional[torch.Tensor] = None,1182            attention_mask: Optional[torch.Tensor] = None,1183            past_key_values: Optional[Tuple[torch.FloatTensor]] = None,1184            inputs_embeds: Optional[torch.Tensor] = None,1185            labels: Optional[torch.Tensor] = None,1186            use_cache: Optional[bool] = None,1187            output_attentions: Optional[bool] = None,1188            output_hidden_states: Optional[bool] = None,1189            return_dict: Optional[bool] = None,1190            return_last_logit: Optional[bool] = False,1191    ):1192        use_cache = use_cache if use_cache is not None else self.config.use_cache1193        return_dict = return_dict if return_dict is not None else self.config.use_return_dict1194 1195        transformer_outputs = self.transformer(1196            input_ids=input_ids,1197            position_ids=position_ids,1198            attention_mask=attention_mask,1199            past_key_values=past_key_values,1200            inputs_embeds=inputs_embeds,

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