From 39df37422f530e701ad3168c4825250e9b0c7442 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 7 Jul 2024 02:25:26 +0300 Subject: [PATCH] text encoder quantization --- examples/kolors_example.json | 432 ++++++++++-------- kolors/models/quantization.py | 188 ++++++++ ...ipeline_stable_diffusion_xl_chatglm_256.py | 8 +- nodes.py | 87 +++- requirements.txt | 3 +- 5 files changed, 506 insertions(+), 212 deletions(-) create mode 100644 kolors/models/quantization.py diff --git a/examples/kolors_example.json b/examples/kolors_example.json index 9843c63..286bc4e 100644 --- a/examples/kolors_example.json +++ b/examples/kolors_example.json @@ -1,90 +1,7 @@ { - "last_node_id": 11, - "last_link_id": 13, + "last_node_id": 15, + "last_link_id": 18, "nodes": [ - { - "id": 6, - "type": "DownloadAndLoadKolorsModel", - "pos": [ - 547, - 372 - ], - "size": { - "0": 315, - "1": 82 - }, - "flags": {}, - "order": 0, - "mode": 0, - "outputs": [ - { - "name": "kolors_model", - "type": "KOLORSMODEL", - "links": [ - 7, - 9 - ], - "shape": 3, - "slot_index": 0 - } - ], - "properties": { - "Node name for S&R": "DownloadAndLoadKolorsModel" - }, - "widgets_values": [ - "Kwai-Kolors/Kolors", - "fp16" - ] - }, - { - "id": 9, - "type": "KolorsSampler", - "pos": [ - 1011, - 371 - ], - "size": { - "0": 315, - "1": 198 - }, - "flags": {}, - "order": 3, - "mode": 0, - "inputs": [ - { - "name": "kolors_model", - "type": "KOLORSMODEL", - "link": 9 - }, - { - "name": "kolors_embeds", - "type": "KOLORS_EMBEDS", - "link": 10 - } - ], - "outputs": [ - { - "name": "latent", - "type": "LATENT", - "links": [ - 11 - ], - "shape": 3, - "slot_index": 0 - } - ], - "properties": { - "Node name for S&R": "KolorsSampler" - }, - "widgets_values": [ - 1024, - 1024, - 243590846571465, - "randomize", - 25, - 5 - ] - }, { "id": 11, "type": "VAELoader", @@ -97,7 +14,7 @@ "1": 58 }, "flags": {}, - "order": 1, + "order": 0, "mode": 0, "outputs": [ { @@ -116,71 +33,6 @@ "sdxl.vae.safetensors" ] }, - { - "id": 8, - "type": "KolorsTextEncode", - "pos": [ - 549, - 522 - ], - "size": { - "0": 400, - "1": 200 - }, - "flags": {}, - "order": 2, - "mode": 0, - "inputs": [ - { - "name": "kolors_model", - "type": "KOLORSMODEL", - "link": 7 - } - ], - "outputs": [ - { - "name": "kolors_embeds", - "type": "KOLORS_EMBEDS", - "links": [ - 10 - ], - "shape": 3 - } - ], - "properties": { - "Node name for S&R": "KolorsTextEncode" - }, - "widgets_values": [ - "cinematic photograph monkey giving two thumbs up", - "nsfw, naked", - 4 - ] - }, - { - "id": 3, - "type": "PreviewImage", - "pos": [ - 1367, - 467 - ], - "size": [ - 670, - 646.6666259765625 - ], - "flags": {}, - "order": 5, - "mode": 0, - "inputs": [ - { - "name": "images", - "type": "IMAGE", - "link": 13 - } - ], - "properties": { - "Node name for S&R": "PreviewImage" - } - }, { "id": 10, "type": "VAEDecode", @@ -193,13 +45,13 @@ "1": 46 }, "flags": {}, - "order": 4, + "order": 6, "mode": 0, "inputs": [ { "name": "samples", "type": "LATENT", - "link": 11 + "link": 18 }, { "name": "vae", @@ -222,41 +74,213 @@ "properties": { "Node name for S&R": "VAEDecode" } + }, + { + "id": 14, + "type": "KolorsSampler", + "pos": [ + 1011, + 371 + ], + "size": { + "0": 315, + "1": 222 + }, + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [ + { + "name": "kolors_model", + "type": "KOLORSMODEL", + "link": 16 + }, + { + "name": "kolors_embeds", + "type": "KOLORS_EMBEDS", + "link": 17 + } + ], + "outputs": [ + { + "name": "latent", + "type": "LATENT", + "links": [ + 18 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "KolorsSampler" + }, + "widgets_values": [ + 1024, + 1024, + 1000102404233412, + "fixed", + 25, + 5, + "EulerDiscreteScheduler" + ] + }, + { + "id": 6, + "type": "DownloadAndLoadKolorsModel", + "pos": [ + 201, + 368 + ], + "size": { + "0": 315, + "1": 82 + }, + "flags": {}, + "order": 1, + "mode": 0, + "outputs": [ + { + "name": "kolors_model", + "type": "KOLORSMODEL", + "links": [ + 16 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "DownloadAndLoadKolorsModel" + }, + "widgets_values": [ + "Kwai-Kolors/Kolors", + "fp16" + ] + }, + { + "id": 3, + "type": "PreviewImage", + "pos": [ + 1366, + 468 + ], + "size": [ + 535.4001724243165, + 562.2001106262207 + ], + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 13 + } + ], + "properties": { + "Node name for S&R": "PreviewImage" + } + }, + { + "id": 12, + "type": "KolorsTextEncode", + "pos": [ + 519, + 529 + ], + "size": [ + 457.2893696934723, + 225.28656056301645 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [ + { + "name": "chatglm3_model", + "type": "CHATGLM3MODEL", + "link": 14, + "slot_index": 0 + } + ], + "outputs": [ + { + "name": "kolors_embeds", + "type": "KOLORS_EMBEDS", + "links": [ + 17 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "KolorsTextEncode" + }, + "widgets_values": [ + "cinematic photograph of an astronaut riding a horse in space |\nillustration of a cat wearing a top hat and a scarf |\nphotograph of a goldfish in a bowl |\nanime screencap of a red haired girl", + "", + 1 + ] + }, + { + "id": 15, + "type": "Note", + "pos": [ + 200, + 636 + ], + "size": [ + 273.5273818969726, + 149.55464588512064 + ], + "flags": {}, + "order": 2, + "mode": 0, + "properties": { + "text": "" + }, + "widgets_values": [ + "Text encoding takes the most VRAM, quantization can reduce that a lot.\n\nApproximate values I have observed:\nfp16 - 12 GB\nquant8 - 8-9 GB\nquant4 - 4-5 GB\n\nquant4 reduces the quality quite a bit, 8 seems fine" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 13, + "type": "DownloadAndLoadChatGLM3", + "pos": [ + 206, + 522 + ], + "size": [ + 274.5334274291992, + 58 + ], + "flags": {}, + "order": 3, + "mode": 0, + "outputs": [ + { + "name": "chatglm3_model", + "type": "CHATGLM3MODEL", + "links": [ + 14 + ], + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "DownloadAndLoadChatGLM3" + }, + "widgets_values": [ + "fp16" + ] } ], "links": [ - [ - 7, - 6, - 0, - 8, - 0, - "KOLORSMODEL" - ], - [ - 9, - 6, - 0, - 9, - 0, - "KOLORSMODEL" - ], - [ - 10, - 8, - 0, - 9, - 1, - "KOLORS_EMBEDS" - ], - [ - 11, - 9, - 0, - 10, - 0, - "LATENT" - ], [ 12, 11, @@ -272,16 +296,48 @@ 3, 0, "IMAGE" + ], + [ + 14, + 13, + 0, + 12, + 0, + "CHATGLM3MODEL" + ], + [ + 16, + 6, + 0, + 14, + 0, + "KOLORSMODEL" + ], + [ + 17, + 12, + 0, + 14, + 1, + "KOLORS_EMBEDS" + ], + [ + 18, + 14, + 0, + 10, + 0, + "LATENT" ] ], "groups": [], "config": {}, "extra": { "ds": { - "scale": 1, + "scale": 1.1, "offset": { - "0": -404.6999816894531, - "1": -111.86663818359375 + "0": -114.73954010009766, + "1": -139.79705810546875 } } }, diff --git a/kolors/models/quantization.py b/kolors/models/quantization.py new file mode 100644 index 0000000..cb95bfe --- /dev/null +++ b/kolors/models/quantization.py @@ -0,0 +1,188 @@ +from torch.nn import Linear +from torch.nn.parameter import Parameter + +import bz2 +import torch +import base64 +import ctypes +from transformers.utils import logging + +from typing import List +from functools import partial + +logger = logging.get_logger(__name__) + +try: + from cpm_kernels.kernels.base import LazyKernelCModule, KernelFunction, round_up + + class Kernel: + def __init__(self, code: bytes, function_names: List[str]): + self.code = code + self._function_names = function_names + self._cmodule = LazyKernelCModule(self.code) + + for name in self._function_names: + setattr(self, name, KernelFunction(self._cmodule, name)) + + 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" + + kernels = Kernel( + bz2.decompress(base64.b64decode(quantization_code)), + [ + "int4WeightCompression", + "int4WeightExtractionFloat", + "int4WeightExtractionHalf", + "int8WeightExtractionFloat", + "int8WeightExtractionHalf", + ], + ) +except Exception as exception: + kernels = None + logger.warning("Failed to load cpm_kernels:" + str(exception)) + + +class W8A16Linear(torch.autograd.Function): + @staticmethod + def forward(ctx, inp: torch.Tensor, quant_w: torch.Tensor, scale_w: torch.Tensor, weight_bit_width): + ctx.inp_shape = inp.size() + ctx.weight_bit_width = weight_bit_width + out_features = quant_w.size(0) + inp = inp.contiguous().view(-1, inp.size(-1)) + weight = extract_weight_to_half(quant_w, scale_w, weight_bit_width) + ctx.weight_shape = weight.size() + output = inp.mm(weight.t()) + ctx.save_for_backward(inp, quant_w, scale_w) + return output.view(*(ctx.inp_shape[:-1] + (out_features,))) + + @staticmethod + def backward(ctx, grad_output: torch.Tensor): + inp, quant_w, scale_w = ctx.saved_tensors + weight = extract_weight_to_half(quant_w, scale_w, ctx.weight_bit_width) + grad_output = grad_output.contiguous().view(-1, weight.size(0)) + grad_input = grad_output.mm(weight) + grad_weight = grad_output.t().mm(inp) + return grad_input.view(ctx.inp_shape), grad_weight.view(ctx.weight_shape), None, None + + +def compress_int4_weight(weight: torch.Tensor): # (n, m) + with torch.cuda.device(weight.device): + n, m = weight.size(0), weight.size(1) + assert m % 2 == 0 + m = m // 2 + out = torch.empty(n, m, dtype=torch.int8, device="cuda") + stream = torch.cuda.current_stream() + + gridDim = (n, 1, 1) + blockDim = (min(round_up(m, 32), 1024), 1, 1) + + kernels.int4WeightCompression( + gridDim, + blockDim, + 0, + stream, + [ctypes.c_void_p(weight.data_ptr()), ctypes.c_void_p(out.data_ptr()), ctypes.c_int32(n), ctypes.c_int32(m)], + ) + return out + + +def extract_weight_to_half(weight: torch.Tensor, scale_list: torch.Tensor, source_bit_width: int): + assert scale_list.dtype in [torch.half, torch.bfloat16] + assert weight.dtype in [torch.int8] + if source_bit_width == 8: + return weight.to(scale_list.dtype) * scale_list[:, None] + elif source_bit_width == 4: + func = ( + kernels.int4WeightExtractionHalf if scale_list.dtype == torch.half else kernels.int4WeightExtractionBFloat16 + ) + else: + assert False, "Unsupported bit-width" + + with torch.cuda.device(weight.device): + n, m = weight.size(0), weight.size(1) + out = torch.empty(n, m * (8 // source_bit_width), dtype=scale_list.dtype, device="cuda") + stream = torch.cuda.current_stream() + + gridDim = (n, 1, 1) + blockDim = (min(round_up(m, 32), 1024), 1, 1) + + func( + gridDim, + blockDim, + 0, + stream, + [ + ctypes.c_void_p(weight.data_ptr()), + ctypes.c_void_p(scale_list.data_ptr()), + ctypes.c_void_p(out.data_ptr()), + ctypes.c_int32(n), + ctypes.c_int32(m), + ], + ) + return out + + +class QuantizedLinear(torch.nn.Module): + def __init__(self, weight_bit_width: int, weight, bias=None, device="cpu", dtype=None, empty_init=False, *args, + **kwargs): + super().__init__() + self.weight_bit_width = weight_bit_width + + shape = weight.shape + + if weight is None or empty_init: + self.weight = torch.empty(shape[0], shape[1] * weight_bit_width // 8, dtype=torch.int8, device=device) + self.weight_scale = torch.empty(shape[0], dtype=dtype, device=device) + else: + self.weight_scale = weight.abs().max(dim=-1).values / ((2 ** (weight_bit_width - 1)) - 1) + self.weight = torch.round(weight / self.weight_scale[:, None]).to(torch.int8) + if weight_bit_width == 4: + self.weight = compress_int4_weight(self.weight) + + self.weight = Parameter(self.weight.to(device), requires_grad=False) + self.weight_scale = Parameter(self.weight_scale.to(device), requires_grad=False) + self.bias = Parameter(bias.to(device), requires_grad=False) if bias is not None else None + + def forward(self, input): + output = W8A16Linear.apply(input, self.weight, self.weight_scale, self.weight_bit_width) + if self.bias is not None: + output = output + self.bias + return output + + +def quantize(model, weight_bit_width, empty_init=False, device=None): + """Replace fp16 linear with quantized linear""" + for layer in model.layers: + layer.self_attention.query_key_value = QuantizedLinear( + weight_bit_width=weight_bit_width, + weight=layer.self_attention.query_key_value.weight.to(torch.cuda.current_device()), + bias=layer.self_attention.query_key_value.bias, + dtype=layer.self_attention.query_key_value.weight.dtype, + device=layer.self_attention.query_key_value.weight.device if device is None else device, + empty_init=empty_init + ) + layer.self_attention.dense = QuantizedLinear( + weight_bit_width=weight_bit_width, + weight=layer.self_attention.dense.weight.to(torch.cuda.current_device()), + bias=layer.self_attention.dense.bias, + dtype=layer.self_attention.dense.weight.dtype, + device=layer.self_attention.dense.weight.device if device is None else device, + empty_init=empty_init + ) + layer.mlp.dense_h_to_4h = QuantizedLinear( + weight_bit_width=weight_bit_width, + weight=layer.mlp.dense_h_to_4h.weight.to(torch.cuda.current_device()), + bias=layer.mlp.dense_h_to_4h.bias, + dtype=layer.mlp.dense_h_to_4h.weight.dtype, + device=layer.mlp.dense_h_to_4h.weight.device if device is None else device, + empty_init=empty_init + ) + layer.mlp.dense_4h_to_h = QuantizedLinear( + weight_bit_width=weight_bit_width, + weight=layer.mlp.dense_4h_to_h.weight.to(torch.cuda.current_device()), + bias=layer.mlp.dense_4h_to_h.bias, + dtype=layer.mlp.dense_4h_to_h.weight.dtype, + device=layer.mlp.dense_4h_to_h.weight.device if device is None else device, + empty_init=empty_init + ) + + return model diff --git a/kolors/pipelines/pipeline_stable_diffusion_xl_chatglm_256.py b/kolors/pipelines/pipeline_stable_diffusion_xl_chatglm_256.py index 9e576ef..e1c9a69 100755 --- a/kolors/pipelines/pipeline_stable_diffusion_xl_chatglm_256.py +++ b/kolors/pipelines/pipeline_stable_diffusion_xl_chatglm_256.py @@ -107,8 +107,8 @@ class StableDiffusionXLPipeline(DiffusionPipeline, FromSingleFileMixin, LoraLoad def __init__( self, - text_encoder: ChatGLMModel, - tokenizer: ChatGLMTokenizer, + # text_encoder: ChatGLMModel, + # tokenizer: ChatGLMTokenizer, unet: UNet2DConditionModel, scheduler: KarrasDiffusionSchedulers, force_zeros_for_empty_prompt: bool = True, @@ -117,8 +117,8 @@ class StableDiffusionXLPipeline(DiffusionPipeline, FromSingleFileMixin, LoraLoad self.register_modules( #vae=vae, - text_encoder=text_encoder, - tokenizer=tokenizer, + #text_encoder=text_encoder, + #tokenizer=tokenizer, unet=unet, scheduler=scheduler, ) diff --git a/nodes.py b/nodes.py index 3e345ca..a875980 100755 --- a/nodes.py +++ b/nodes.py @@ -3,13 +3,14 @@ import os import random import re import gc - +import sys import comfy.model_management as mm from comfy.utils import ProgressBar, load_torch_file import folder_paths script_directory = os.path.dirname(os.path.abspath(__file__)) +sys.path.append(script_directory) from .kolors.pipelines.pipeline_stable_diffusion_xl_chatglm_256 import StableDiffusionXLPipeline from .kolors.models.modeling_chatglm import ChatGLMModel @@ -59,7 +60,8 @@ class DownloadAndLoadKolorsModel: print(f"Downloading Kolor model to: {model_path}") from huggingface_hub import snapshot_download snapshot_download(repo_id=model, - allow_patterns=['*fp16.safetensors*', '*.json', 'text_encoder/*', 'tokenizer/*'], + allow_patterns=['*fp16.safetensors*', '*.json'], + ignore_patthers=['text_encoder/*', 'tokenizer/*'], local_dir=model_path, local_dir_use_symlinks=False) pbar.update(1) @@ -67,28 +69,19 @@ class DownloadAndLoadKolorsModel: scheduler = EulerDiscreteScheduler.from_pretrained(model_path, subfolder= 'scheduler') print("Load UNET...") - unet = UNet2DConditionModel.from_pretrained(model_path, subfolder= 'unet', variant="fp16", revision=None, low_cpu_mem_usage=True).to(dtype).eval() - print("Load TEXT_ENCODER...") - pbar.update(1) + unet = UNet2DConditionModel.from_pretrained(model_path, subfolder= 'unet', variant="fp16", revision=None, low_cpu_mem_usage=True).to(dtype).eval() - text_encoder_path = os.path.join(model_path, "text_encoder") - text_encoder = ChatGLMModel.from_pretrained( - text_encoder_path, - torch_dtype=dtype, - ) - tokenizer = ChatGLMTokenizer.from_pretrained(text_encoder_path) - pbar.update(1) pipeline = StableDiffusionXLPipeline( #vae=None, - text_encoder=text_encoder, - tokenizer=tokenizer, + #text_encoder=None, + #tokenizer=None, unet=unet, scheduler=scheduler, force_zeros_for_empty_prompt=False ) #pipeline = pipeline.to(device) - pipeline.enable_model_cpu_offload() + #pipeline.enable_model_cpu_offload() kolors_model = { 'pipeline': pipeline, @@ -97,13 +90,67 @@ class DownloadAndLoadKolorsModel: return (kolors_model,) +class DownloadAndLoadChatGLM3: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "precision": ([ 'fp16', 'quant4', 'quant8'], + { + "default": 'fp16' + }), + }, + } + + RETURN_TYPES = ("CHATGLM3MODEL",) + RETURN_NAMES = ("chatglm3_model",) + FUNCTION = "loadmodel" + CATEGORY = "KwaiKolorsWrapper" + + def loadmodel(self, precision): + + pbar = ProgressBar(2) + model = "Kwai-Kolors/Kolors" + model_name = model.rsplit('/', 1)[-1] + model_path = os.path.join(folder_paths.models_dir, "diffusers", model_name) + text_encoder_path = os.path.join(model_path, "text_encoder") + + if not os.path.exists(text_encoder_path): + print(f"Downloading Kolor model to: {model_path}") + from huggingface_hub import snapshot_download + snapshot_download(repo_id=model, + allow_patterns=['text_encoder/*', 'tokenizer/*'], + local_dir=model_path, + local_dir_use_symlinks=False) + pbar.update(1) + + print("Load TEXT_ENCODER...") + + text_encoder_path = os.path.join(model_path, "text_encoder") + text_encoder = ChatGLMModel.from_pretrained( + text_encoder_path, + torch_dtype=torch.float16, + ) + if precision == 'quant8': + text_encoder.quantize(8) + elif precision == 'quant4': + text_encoder.quantize(4) + + tokenizer = ChatGLMTokenizer.from_pretrained(text_encoder_path) + pbar.update(1) + chatglm3_model = { + 'text_encoder': text_encoder, + 'tokenizer': tokenizer + } + + return (chatglm3_model,) + class KolorsTextEncode: @classmethod def INPUT_TYPES(s): return { "required": { - "kolors_model": ("KOLORSMODEL", ), + "chatglm3_model": ("CHATGLM3MODEL", ), "prompt": ("STRING", {"multiline": True, "default": "",}), "negative_prompt": ("STRING", {"multiline": True, "default": "",}), "num_images_per_prompt": ("INT", {"default": 1, "min": 1, "max": 128, "step": 1}), @@ -115,7 +162,7 @@ class KolorsTextEncode: FUNCTION = "encode" CATEGORY = "KwaiKolorsWrapper" - def encode(self, kolors_model, prompt, negative_prompt, num_images_per_prompt): + def encode(self, chatglm3_model, prompt, negative_prompt, num_images_per_prompt): device = mm.get_torch_device() offload_device = mm.unet_offload_device() mm.unload_all_models() @@ -143,8 +190,8 @@ class KolorsTextEncode: batch_size = len(prompt) # Define tokenizers and text encoders - tokenizer = kolors_model['pipeline'].tokenizer - text_encoder = kolors_model['pipeline'].text_encoder + tokenizer = chatglm3_model['tokenizer'] + text_encoder = chatglm3_model['text_encoder'] text_encoder.to(device) @@ -341,11 +388,13 @@ class KolorsSampler: NODE_CLASS_MAPPINGS = { "DownloadAndLoadKolorsModel": DownloadAndLoadKolorsModel, + "DownloadAndLoadChatGLM3": DownloadAndLoadChatGLM3, "KolorsSampler": KolorsSampler, "KolorsTextEncode": KolorsTextEncode } NODE_DISPLAY_NAME_MAPPINGS = { "DownloadAndLoadKolorsModel": "(Down)load Kolors Model", + "DownloadAndLoadChatGLM3": "(Down)load ChatGLM3 Model", "KolorsSampler": "Kolors Sampler", "KolorsTextEncode": "Kolors Text Encode" } \ No newline at end of file diff --git a/requirements.txt b/requirements.txt index 592f3b8..5eb41f2 100755 --- a/requirements.txt +++ b/requirements.txt @@ -1,4 +1,5 @@ diffusers>=0.28.2 transformers>=4.26.1 sentencepiece -accelerate \ No newline at end of file +accelerate +cpm-kernels \ No newline at end of file