hugging-apps/echo-memory
0
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,