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