Store lora diffs in buffers for GGUF as well
This commit is contained in:
+3
-22
@@ -8,15 +8,15 @@ def fp8_linear_forward(cls, base_dtype, input):
|
||||
if weight_dtype in [torch.float8_e4m3fn, torch.float8_e5m2]:
|
||||
if len(input.shape) == 3:
|
||||
input_shape = input.shape
|
||||
|
||||
|
||||
scale_weight = getattr(cls, 'scale_weight', None)
|
||||
if scale_weight is None:
|
||||
scale_weight = torch.ones((), device=input.device, dtype=torch.float32)
|
||||
else:
|
||||
scale_weight = scale_weight.to(input.device).squeeze()
|
||||
|
||||
|
||||
scale_input = torch.ones((), device=input.device, dtype=torch.float32)
|
||||
|
||||
|
||||
input = torch.clamp(input, min=-448, max=448, out=input)
|
||||
inn = input.reshape(-1, input_shape[2]).to(torch.float8_e4m3fn).contiguous() #always e4m3fn because e5m2 * e5m2 is not supported
|
||||
|
||||
@@ -31,24 +31,6 @@ def fp8_linear_forward(cls, base_dtype, input):
|
||||
return cls.original_forward(input)
|
||||
|
||||
|
||||
@torch.compiler.disable()
|
||||
def apply_lora(weight, lora, step=None):
|
||||
for lora_diff, lora_strength in zip(lora[0], lora[1]):
|
||||
if isinstance(lora_strength, list):
|
||||
lora_strength = lora_strength[step]
|
||||
if lora_strength == 0.0:
|
||||
continue
|
||||
elif lora_strength == 0.0:
|
||||
continue
|
||||
patch_diff = torch.mm(
|
||||
lora_diff[0].flatten(start_dim=1).to(weight.device),
|
||||
lora_diff[1].flatten(start_dim=1).to(weight.device)
|
||||
).reshape(weight.shape)
|
||||
alpha = lora_diff[2] / lora_diff[1].shape[0] if lora_diff[2] is not None else 1.0
|
||||
scale = lora_strength * alpha
|
||||
weight = weight.add(patch_diff, alpha=scale)
|
||||
return weight
|
||||
|
||||
def convert_fp8_linear(module, base_dtype, params_to_keep={}, scale_weight_keys=None):
|
||||
log.info("FP8 matmul enabled")
|
||||
for name, submodule in module.named_modules():
|
||||
@@ -61,4 +43,3 @@ def convert_fp8_linear(module, base_dtype, params_to_keep={}, scale_weight_keys=
|
||||
original_forward = submodule.forward
|
||||
setattr(submodule, "original_forward", original_forward)
|
||||
setattr(submodule, "forward", lambda input, m=submodule: fp8_linear_forward(m, base_dtype, input))
|
||||
|
||||
|
||||
+50
-36
@@ -1,13 +1,11 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import numpy as np
|
||||
from diffusers.quantizers.gguf.utils import GGUFParameter, dequantize_gguf_tensor
|
||||
import gguf
|
||||
from diffusers.utils import is_accelerate_available
|
||||
from contextlib import nullcontext
|
||||
from accelerate import init_empty_weights
|
||||
|
||||
from .gguf_utils import GGUFParameter, dequantize_gguf_tensor
|
||||
from ..utils import log
|
||||
if is_accelerate_available():
|
||||
from accelerate import init_empty_weights
|
||||
|
||||
def load_gguf(model_path):
|
||||
from gguf import GGUFReader
|
||||
@@ -43,8 +41,7 @@ def _replace_with_gguf_linear(model, compute_dtype, state_dict, prefix="", modul
|
||||
in_features = state_dict[module_prefix + "weight"].shape[1]
|
||||
out_features = state_dict[module_prefix + "weight"].shape[0]
|
||||
|
||||
ctx = init_empty_weights if is_accelerate_available() else nullcontext
|
||||
with ctx():
|
||||
with init_empty_weights():
|
||||
model._modules[name] = GGUFLinear(
|
||||
in_features,
|
||||
out_features,
|
||||
@@ -53,19 +50,26 @@ def _replace_with_gguf_linear(model, compute_dtype, state_dict, prefix="", modul
|
||||
)
|
||||
|
||||
model._modules[name].source_cls = type(module)
|
||||
# Force requires_grad to False to avoid unexpected errors
|
||||
model._modules[name].requires_grad_(False)
|
||||
return model
|
||||
|
||||
def set_lora_params_gguf(module, patches, module_prefix=""):
|
||||
def set_lora_params_gguf(module, patches, module_prefix="", device=torch.device("cpu")):
|
||||
# Recursively set lora_diffs and lora_strengths for all GGUFLinear layers
|
||||
for name, child in module.named_children():
|
||||
params = list(child.parameters())
|
||||
if params:
|
||||
device = params[0].device
|
||||
else:
|
||||
device = torch.device("cpu")
|
||||
child_prefix = (f"{module_prefix}{name}.")
|
||||
set_lora_params_gguf(child, patches, child_prefix)
|
||||
set_lora_params_gguf(child, patches, child_prefix, device)
|
||||
if isinstance(module, GGUFLinear):
|
||||
key = f"diffusion_model.{module_prefix}weight"
|
||||
patch = patches.get(key, [])
|
||||
#print(f"Processing LoRA patches for {key}: {len(patch)} patches found")
|
||||
if len(patch) == 0:
|
||||
key = key.replace("_orig_mod.", "")
|
||||
patch = patches.get(key, [])
|
||||
if len(patch) != 0:
|
||||
lora_diffs = []
|
||||
for p in patch:
|
||||
@@ -78,8 +82,8 @@ def set_lora_params_gguf(module, patches, module_prefix=""):
|
||||
lora_diffs.append(lora_obj[1])
|
||||
else:
|
||||
continue
|
||||
lora_strengths = [p[0] for p in patch]
|
||||
module.lora = (lora_diffs, lora_strengths)
|
||||
module.lora_strengths = [p[0] for p in patch]
|
||||
module.set_lora_diffs(lora_diffs, device=device)
|
||||
module.step = 0 # Initialize step for LoRA scheduling
|
||||
|
||||
|
||||
@@ -94,41 +98,51 @@ class GGUFLinear(nn.Linear):
|
||||
) -> None:
|
||||
super().__init__(in_features, out_features, bias, device)
|
||||
self.compute_dtype = compute_dtype
|
||||
self.lora = None
|
||||
self.lora_diffs = []
|
||||
self.lora_strengths = []
|
||||
self.step = 0
|
||||
|
||||
def forward(self, inputs):
|
||||
weight = self.dequantize_without_compile()
|
||||
weight = weight.to(self.compute_dtype)
|
||||
weight = dequantize_gguf_tensor(self.weight).to(self.compute_dtype)
|
||||
bias = self.bias.to(self.compute_dtype) if self.bias is not None else None
|
||||
|
||||
if hasattr(self, "lora") and self.lora is not None:
|
||||
weight = self.apply_lora(weight, self.step).to(self.compute_dtype)
|
||||
if hasattr(self, f"lora_diff_0_0"):
|
||||
weight = self.apply_lora(weight).to(self.compute_dtype)
|
||||
|
||||
output = torch.nn.functional.linear(inputs, weight, bias)
|
||||
return output
|
||||
return torch.nn.functional.linear(inputs, weight, bias)
|
||||
|
||||
@torch.compiler.disable()
|
||||
def dequantize_without_compile(self):
|
||||
return dequantize_gguf_tensor(self.weight)
|
||||
def set_lora_diffs(self, lora_diffs, device=torch.device("cpu")):
|
||||
self.lora_diffs = []
|
||||
for i, diff in enumerate(lora_diffs):
|
||||
if isinstance(diff, tuple):
|
||||
self.register_buffer(f"lora_diff_{i}_0", diff[0].to(device))
|
||||
self.register_buffer(f"lora_diff_{i}_1", diff[1].to(device))
|
||||
setattr(self, f"lora_diff_{i}_2", diff[2])
|
||||
self.lora_diffs.append((f"lora_diff_{i}_0", f"lora_diff_{i}_1", f"lora_diff_{i}_2"))
|
||||
else:
|
||||
self.register_buffer(f"lora_diff_{i}", diff.to(device))
|
||||
self.lora_diffs.append(f"lora_diff_{i}")
|
||||
|
||||
@torch.compiler.disable()
|
||||
def apply_lora(self, weight, step=None):
|
||||
for lora_diff, lora_strength in zip(self.lora[0], self.lora[1]):
|
||||
def apply_lora(self, weight):
|
||||
for lora_diff_names, lora_strength in zip(self.lora_diffs, self.lora_strengths):
|
||||
if isinstance(lora_strength, list):
|
||||
lora_strength = lora_strength[step]
|
||||
lora_strength = lora_strength[self.step]
|
||||
if lora_strength == 0.0:
|
||||
continue
|
||||
elif lora_strength == 0.0:
|
||||
continue
|
||||
if len(lora_diff) == 1:
|
||||
weight = weight.add(lora_diff[0].to(weight.device), alpha=lora_strength)
|
||||
continue
|
||||
patch_diff = torch.mm(
|
||||
lora_diff[0].flatten(start_dim=1).to(weight.device),
|
||||
lora_diff[1].flatten(start_dim=1).to(weight.device)
|
||||
).reshape(weight.shape)
|
||||
alpha = lora_diff[2] / lora_diff[1].shape[0] if lora_diff[2] is not None else 1.0
|
||||
scale = lora_strength * alpha
|
||||
weight = weight.add(patch_diff, alpha=scale)
|
||||
if isinstance(lora_diff_names, tuple):
|
||||
lora_diff_0 = getattr(self, lora_diff_names[0])
|
||||
lora_diff_1 = getattr(self, lora_diff_names[1])
|
||||
lora_diff_2 = getattr(self, lora_diff_names[2])
|
||||
patch_diff = torch.mm(
|
||||
lora_diff_0.flatten(start_dim=1),
|
||||
lora_diff_1.flatten(start_dim=1)
|
||||
).reshape(weight.shape) + 0
|
||||
alpha = lora_diff_2 / lora_diff_1.shape[0] if lora_diff_2 is not None else 1.0
|
||||
scale = lora_strength * alpha
|
||||
weight = weight.add(patch_diff, alpha=scale)
|
||||
else:
|
||||
lora_diff = getattr(self, lora_diff_names)
|
||||
weight = weight.add(lora_diff, alpha=lora_strength)
|
||||
return weight
|
||||
@@ -0,0 +1,338 @@
|
||||
# Copyright 2024 The HuggingFace Team and City96. All rights reserved.
|
||||
# #
|
||||
# # Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# # you may not use this file except in compliance with the License.
|
||||
# # You may obtain a copy of the License at
|
||||
# #
|
||||
# # http://www.apache.org/licenses/LICENSE-2.0
|
||||
# #
|
||||
# # Unless required by applicable law or agreed to in writing, software
|
||||
# # distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# # See the License for the specific language governing permissions and
|
||||
# # limitations under the License.
|
||||
|
||||
|
||||
import gguf
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
# dequantize operations based on torch ports of GGUF dequantize_functions
|
||||
# from City96
|
||||
# more info: https://github.com/city96/ComfyUI-GGUF/blob/main/dequant.py
|
||||
|
||||
|
||||
QK_K = 256
|
||||
K_SCALE_SIZE = 12
|
||||
|
||||
|
||||
def to_uint32(x):
|
||||
x = x.view(torch.uint8).to(torch.int32)
|
||||
return (x[:, 0] | x[:, 1] << 8 | x[:, 2] << 16 | x[:, 3] << 24).unsqueeze(1)
|
||||
|
||||
|
||||
def split_block_dims(blocks, *args):
|
||||
n_max = blocks.shape[1]
|
||||
dims = list(args) + [n_max - sum(args)]
|
||||
return torch.split(blocks, dims, dim=1)
|
||||
|
||||
|
||||
def get_scale_min(scales):
|
||||
n_blocks = scales.shape[0]
|
||||
scales = scales.view(torch.uint8)
|
||||
scales = scales.reshape((n_blocks, 3, 4))
|
||||
|
||||
d, m, m_d = torch.split(scales, scales.shape[-2] // 3, dim=-2)
|
||||
|
||||
sc = torch.cat([d & 0x3F, (m_d & 0x0F) | ((d >> 2) & 0x30)], dim=-1)
|
||||
min = torch.cat([m & 0x3F, (m_d >> 4) | ((m >> 2) & 0x30)], dim=-1)
|
||||
|
||||
return (sc.reshape((n_blocks, 8)), min.reshape((n_blocks, 8)))
|
||||
|
||||
|
||||
def dequantize_blocks_Q8_0(blocks, block_size, type_size, dtype=None):
|
||||
d, x = split_block_dims(blocks, 2)
|
||||
d = d.view(torch.float16).to(dtype)
|
||||
x = x.view(torch.int8)
|
||||
return d * x
|
||||
|
||||
|
||||
def dequantize_blocks_Q5_1(blocks, block_size, type_size, dtype=None):
|
||||
n_blocks = blocks.shape[0]
|
||||
|
||||
d, m, qh, qs = split_block_dims(blocks, 2, 2, 4)
|
||||
d = d.view(torch.float16).to(dtype)
|
||||
m = m.view(torch.float16).to(dtype)
|
||||
qh = to_uint32(qh)
|
||||
|
||||
qh = qh.reshape((n_blocks, 1)) >> torch.arange(32, device=d.device, dtype=torch.int32).reshape(1, 32)
|
||||
ql = qs.reshape((n_blocks, -1, 1, block_size // 2)) >> torch.tensor(
|
||||
[0, 4], device=d.device, dtype=torch.uint8
|
||||
).reshape(1, 1, 2, 1)
|
||||
qh = (qh & 1).to(torch.uint8)
|
||||
ql = (ql & 0x0F).reshape((n_blocks, -1))
|
||||
|
||||
qs = ql | (qh << 4)
|
||||
return (d * qs) + m
|
||||
|
||||
|
||||
def dequantize_blocks_Q5_0(blocks, block_size, type_size, dtype=None):
|
||||
n_blocks = blocks.shape[0]
|
||||
|
||||
d, qh, qs = split_block_dims(blocks, 2, 4)
|
||||
d = d.view(torch.float16).to(dtype)
|
||||
qh = to_uint32(qh)
|
||||
|
||||
qh = qh.reshape(n_blocks, 1) >> torch.arange(32, device=d.device, dtype=torch.int32).reshape(1, 32)
|
||||
ql = qs.reshape(n_blocks, -1, 1, block_size // 2) >> torch.tensor(
|
||||
[0, 4], device=d.device, dtype=torch.uint8
|
||||
).reshape(1, 1, 2, 1)
|
||||
|
||||
qh = (qh & 1).to(torch.uint8)
|
||||
ql = (ql & 0x0F).reshape(n_blocks, -1)
|
||||
|
||||
qs = (ql | (qh << 4)).to(torch.int8) - 16
|
||||
return d * qs
|
||||
|
||||
|
||||
def dequantize_blocks_Q4_1(blocks, block_size, type_size, dtype=None):
|
||||
n_blocks = blocks.shape[0]
|
||||
|
||||
d, m, qs = split_block_dims(blocks, 2, 2)
|
||||
d = d.view(torch.float16).to(dtype)
|
||||
m = m.view(torch.float16).to(dtype)
|
||||
|
||||
qs = qs.reshape((n_blocks, -1, 1, block_size // 2)) >> torch.tensor(
|
||||
[0, 4], device=d.device, dtype=torch.uint8
|
||||
).reshape(1, 1, 2, 1)
|
||||
qs = (qs & 0x0F).reshape(n_blocks, -1)
|
||||
|
||||
return (d * qs) + m
|
||||
|
||||
|
||||
def dequantize_blocks_Q4_0(blocks, block_size, type_size, dtype=None):
|
||||
n_blocks = blocks.shape[0]
|
||||
|
||||
d, qs = split_block_dims(blocks, 2)
|
||||
d = d.view(torch.float16).to(dtype)
|
||||
|
||||
qs = qs.reshape((n_blocks, -1, 1, block_size // 2)) >> torch.tensor(
|
||||
[0, 4], device=d.device, dtype=torch.uint8
|
||||
).reshape((1, 1, 2, 1))
|
||||
qs = (qs & 0x0F).reshape((n_blocks, -1)).to(torch.int8) - 8
|
||||
return d * qs
|
||||
|
||||
|
||||
def dequantize_blocks_Q6_K(blocks, block_size, type_size, dtype=None):
|
||||
n_blocks = blocks.shape[0]
|
||||
|
||||
(
|
||||
ql,
|
||||
qh,
|
||||
scales,
|
||||
d,
|
||||
) = split_block_dims(blocks, QK_K // 2, QK_K // 4, QK_K // 16)
|
||||
|
||||
scales = scales.view(torch.int8).to(dtype)
|
||||
d = d.view(torch.float16).to(dtype)
|
||||
d = (d * scales).reshape((n_blocks, QK_K // 16, 1))
|
||||
|
||||
ql = ql.reshape((n_blocks, -1, 1, 64)) >> torch.tensor([0, 4], device=d.device, dtype=torch.uint8).reshape(
|
||||
(1, 1, 2, 1)
|
||||
)
|
||||
ql = (ql & 0x0F).reshape((n_blocks, -1, 32))
|
||||
qh = qh.reshape((n_blocks, -1, 1, 32)) >> torch.tensor([0, 2, 4, 6], device=d.device, dtype=torch.uint8).reshape(
|
||||
(1, 1, 4, 1)
|
||||
)
|
||||
qh = (qh & 0x03).reshape((n_blocks, -1, 32))
|
||||
q = (ql | (qh << 4)).to(torch.int8) - 32
|
||||
q = q.reshape((n_blocks, QK_K // 16, -1))
|
||||
|
||||
return (d * q).reshape((n_blocks, QK_K))
|
||||
|
||||
|
||||
def dequantize_blocks_Q5_K(blocks, block_size, type_size, dtype=None):
|
||||
n_blocks = blocks.shape[0]
|
||||
|
||||
d, dmin, scales, qh, qs = split_block_dims(blocks, 2, 2, K_SCALE_SIZE, QK_K // 8)
|
||||
|
||||
d = d.view(torch.float16).to(dtype)
|
||||
dmin = dmin.view(torch.float16).to(dtype)
|
||||
|
||||
sc, m = get_scale_min(scales)
|
||||
|
||||
d = (d * sc).reshape((n_blocks, -1, 1))
|
||||
dm = (dmin * m).reshape((n_blocks, -1, 1))
|
||||
|
||||
ql = qs.reshape((n_blocks, -1, 1, 32)) >> torch.tensor([0, 4], device=d.device, dtype=torch.uint8).reshape(
|
||||
(1, 1, 2, 1)
|
||||
)
|
||||
qh = qh.reshape((n_blocks, -1, 1, 32)) >> torch.arange(0, 8, device=d.device, dtype=torch.uint8).reshape(
|
||||
(1, 1, 8, 1)
|
||||
)
|
||||
ql = (ql & 0x0F).reshape((n_blocks, -1, 32))
|
||||
qh = (qh & 0x01).reshape((n_blocks, -1, 32))
|
||||
q = ql | (qh << 4)
|
||||
|
||||
return (d * q - dm).reshape((n_blocks, QK_K))
|
||||
|
||||
|
||||
def dequantize_blocks_Q4_K(blocks, block_size, type_size, dtype=None):
|
||||
n_blocks = blocks.shape[0]
|
||||
|
||||
d, dmin, scales, qs = split_block_dims(blocks, 2, 2, K_SCALE_SIZE)
|
||||
d = d.view(torch.float16).to(dtype)
|
||||
dmin = dmin.view(torch.float16).to(dtype)
|
||||
|
||||
sc, m = get_scale_min(scales)
|
||||
|
||||
d = (d * sc).reshape((n_blocks, -1, 1))
|
||||
dm = (dmin * m).reshape((n_blocks, -1, 1))
|
||||
|
||||
qs = qs.reshape((n_blocks, -1, 1, 32)) >> torch.tensor([0, 4], device=d.device, dtype=torch.uint8).reshape(
|
||||
(1, 1, 2, 1)
|
||||
)
|
||||
qs = (qs & 0x0F).reshape((n_blocks, -1, 32))
|
||||
|
||||
return (d * qs - dm).reshape((n_blocks, QK_K))
|
||||
|
||||
|
||||
def dequantize_blocks_Q3_K(blocks, block_size, type_size, dtype=None):
|
||||
n_blocks = blocks.shape[0]
|
||||
|
||||
hmask, qs, scales, d = split_block_dims(blocks, QK_K // 8, QK_K // 4, 12)
|
||||
d = d.view(torch.float16).to(dtype)
|
||||
|
||||
lscales, hscales = scales[:, :8], scales[:, 8:]
|
||||
lscales = lscales.reshape((n_blocks, 1, 8)) >> torch.tensor([0, 4], device=d.device, dtype=torch.uint8).reshape(
|
||||
(1, 2, 1)
|
||||
)
|
||||
lscales = lscales.reshape((n_blocks, 16))
|
||||
hscales = hscales.reshape((n_blocks, 1, 4)) >> torch.tensor(
|
||||
[0, 2, 4, 6], device=d.device, dtype=torch.uint8
|
||||
).reshape((1, 4, 1))
|
||||
hscales = hscales.reshape((n_blocks, 16))
|
||||
scales = (lscales & 0x0F) | ((hscales & 0x03) << 4)
|
||||
scales = scales.to(torch.int8) - 32
|
||||
|
||||
dl = (d * scales).reshape((n_blocks, 16, 1))
|
||||
|
||||
ql = qs.reshape((n_blocks, -1, 1, 32)) >> torch.tensor([0, 2, 4, 6], device=d.device, dtype=torch.uint8).reshape(
|
||||
(1, 1, 4, 1)
|
||||
)
|
||||
qh = hmask.reshape(n_blocks, -1, 1, 32) >> torch.arange(0, 8, device=d.device, dtype=torch.uint8).reshape(
|
||||
(1, 1, 8, 1)
|
||||
)
|
||||
ql = ql.reshape((n_blocks, 16, QK_K // 16)) & 3
|
||||
qh = (qh.reshape((n_blocks, 16, QK_K // 16)) & 1) ^ 1
|
||||
q = ql.to(torch.int8) - (qh << 2).to(torch.int8)
|
||||
|
||||
return (dl * q).reshape((n_blocks, QK_K))
|
||||
|
||||
|
||||
def dequantize_blocks_Q2_K(blocks, block_size, type_size, dtype=None):
|
||||
n_blocks = blocks.shape[0]
|
||||
|
||||
scales, qs, d, dmin = split_block_dims(blocks, QK_K // 16, QK_K // 4, 2)
|
||||
d = d.view(torch.float16).to(dtype)
|
||||
dmin = dmin.view(torch.float16).to(dtype)
|
||||
|
||||
# (n_blocks, 16, 1)
|
||||
dl = (d * (scales & 0xF)).reshape((n_blocks, QK_K // 16, 1))
|
||||
ml = (dmin * (scales >> 4)).reshape((n_blocks, QK_K // 16, 1))
|
||||
|
||||
shift = torch.tensor([0, 2, 4, 6], device=d.device, dtype=torch.uint8).reshape((1, 1, 4, 1))
|
||||
|
||||
qs = (qs.reshape((n_blocks, -1, 1, 32)) >> shift) & 3
|
||||
qs = qs.reshape((n_blocks, QK_K // 16, 16))
|
||||
qs = dl * qs - ml
|
||||
|
||||
return qs.reshape((n_blocks, -1))
|
||||
|
||||
|
||||
def dequantize_blocks_BF16(blocks, block_size, type_size, dtype=None):
|
||||
return (blocks.view(torch.int16).to(torch.int32) << 16).view(torch.float32)
|
||||
|
||||
|
||||
GGML_QUANT_SIZES = gguf.GGML_QUANT_SIZES
|
||||
dequantize_functions = {
|
||||
gguf.GGMLQuantizationType.BF16: dequantize_blocks_BF16,
|
||||
gguf.GGMLQuantizationType.Q8_0: dequantize_blocks_Q8_0,
|
||||
gguf.GGMLQuantizationType.Q5_1: dequantize_blocks_Q5_1,
|
||||
gguf.GGMLQuantizationType.Q5_0: dequantize_blocks_Q5_0,
|
||||
gguf.GGMLQuantizationType.Q4_1: dequantize_blocks_Q4_1,
|
||||
gguf.GGMLQuantizationType.Q4_0: dequantize_blocks_Q4_0,
|
||||
gguf.GGMLQuantizationType.Q6_K: dequantize_blocks_Q6_K,
|
||||
gguf.GGMLQuantizationType.Q5_K: dequantize_blocks_Q5_K,
|
||||
gguf.GGMLQuantizationType.Q4_K: dequantize_blocks_Q4_K,
|
||||
gguf.GGMLQuantizationType.Q3_K: dequantize_blocks_Q3_K,
|
||||
gguf.GGMLQuantizationType.Q2_K: dequantize_blocks_Q2_K,
|
||||
}
|
||||
SUPPORTED_GGUF_QUANT_TYPES = list(dequantize_functions.keys())
|
||||
|
||||
|
||||
def _quant_shape_from_byte_shape(shape, type_size, block_size):
|
||||
return (*shape[:-1], shape[-1] // type_size * block_size)
|
||||
|
||||
|
||||
def dequantize_gguf_tensor(tensor):
|
||||
if not hasattr(tensor, "quant_type"):
|
||||
return tensor
|
||||
|
||||
quant_type = tensor.quant_type
|
||||
dequant_fn = dequantize_functions[quant_type]
|
||||
|
||||
block_size, type_size = GGML_QUANT_SIZES[quant_type]
|
||||
|
||||
tensor = tensor.view(torch.uint8)
|
||||
shape = _quant_shape_from_byte_shape(tensor.shape, type_size, block_size)
|
||||
|
||||
n_blocks = tensor.numel() // type_size
|
||||
blocks = tensor.reshape((n_blocks, type_size))
|
||||
|
||||
dequant = dequant_fn(blocks, block_size, type_size)
|
||||
dequant = dequant.reshape(shape)
|
||||
|
||||
return dequant.as_tensor()
|
||||
|
||||
|
||||
class GGUFParameter(torch.nn.Parameter):
|
||||
def __new__(cls, data, requires_grad=False, quant_type=None):
|
||||
data = data if data is not None else torch.empty(0)
|
||||
self = torch.Tensor._make_subclass(cls, data, requires_grad)
|
||||
self.quant_type = quant_type
|
||||
block_size, type_size = GGML_QUANT_SIZES[quant_type]
|
||||
self.quant_shape = _quant_shape_from_byte_shape(self.shape, type_size, block_size)
|
||||
|
||||
return self
|
||||
|
||||
def as_tensor(self):
|
||||
return torch.Tensor._make_subclass(torch.Tensor, self, self.requires_grad)
|
||||
|
||||
@classmethod
|
||||
def __torch_function__(cls, func, types, args=(), kwargs=None):
|
||||
if kwargs is None:
|
||||
kwargs = {}
|
||||
|
||||
result = super().__torch_function__(func, types, args, kwargs)
|
||||
|
||||
# When converting from original format checkpoints we often use splits, cats etc on tensors
|
||||
# this method ensures that the returned tensor type from those operations remains GGUFParameter
|
||||
# so that we preserve quant_type information
|
||||
quant_type = None
|
||||
for arg in args:
|
||||
if isinstance(arg, list) and isinstance(arg[0], GGUFParameter):
|
||||
quant_type = arg[0].quant_type
|
||||
break
|
||||
if isinstance(arg, GGUFParameter):
|
||||
quant_type = arg.quant_type
|
||||
break
|
||||
if isinstance(result, torch.Tensor):
|
||||
return cls(result, quant_type=quant_type)
|
||||
# Handle tuples and lists
|
||||
elif isinstance(result, (tuple, list)):
|
||||
# Preserve the original type (tuple or list)
|
||||
wrapped = [cls(x, quant_type=quant_type) if isinstance(x, torch.Tensor) else x for x in result]
|
||||
return type(result)(wrapped)
|
||||
else:
|
||||
return result
|
||||
+28
-34
@@ -24,6 +24,28 @@ from ...echoshot.echoshot import rope_apply_z, rope_apply_c, rope_apply_echoshot
|
||||
|
||||
from ...MTV.mtv import apply_rotary_emb
|
||||
|
||||
from comfy import model_management as mm
|
||||
|
||||
__all__ = ['WanModel']
|
||||
|
||||
class AdaLayerNorm(nn.Module):
|
||||
def __init__(self, embedding_dim, output_dim=None, norm_elementwise_affine=False, norm_eps=1e-5, dtype=None, device=None, operations=None):
|
||||
super().__init__()
|
||||
|
||||
output_dim = output_dim or embedding_dim * 2
|
||||
|
||||
self.silu = nn.SiLU()
|
||||
self.linear = operations.Linear(embedding_dim, output_dim, dtype=dtype, device=device)
|
||||
self.norm = operations.LayerNorm(output_dim // 2, norm_eps, norm_elementwise_affine, dtype=dtype, device=device)
|
||||
|
||||
def forward(self, x, temb):
|
||||
temb = self.linear(self.silu(temb))
|
||||
shift, scale = temb.chunk(2, dim=1)
|
||||
shift = shift[:, None, :]
|
||||
scale = scale[:, None, :]
|
||||
x = self.norm(x) * (1 + scale) + shift
|
||||
return x
|
||||
|
||||
class FramePackMotioner(nn.Module):#from comfy.ldm.wan.model
|
||||
def __init__(
|
||||
self,
|
||||
@@ -77,22 +99,11 @@ class FramePackMotioner(nn.Module):#from comfy.ldm.wan.model
|
||||
rope = torch.cat([rope_post, rope_2x, rope_4x], dim=1)
|
||||
return motion_lat, rope
|
||||
|
||||
from diffusers.models.attention import AdaLayerNorm
|
||||
|
||||
__all__ = ['WanModel']
|
||||
|
||||
from comfy import model_management as mm
|
||||
|
||||
|
||||
def zero_module(module):
|
||||
"""
|
||||
Zero out the parameters of a module and return it.
|
||||
"""
|
||||
for p in module.parameters():
|
||||
p.detach().zero_()
|
||||
return module
|
||||
|
||||
|
||||
def torch_dfs(model: nn.Module, parent_name='root'):
|
||||
module_names, modules = [], []
|
||||
current_name = parent_name if parent_name else 'root'
|
||||
@@ -404,7 +415,7 @@ class WanLayerNorm(nn.LayerNorm):
|
||||
"""
|
||||
return super().forward(x)
|
||||
|
||||
|
||||
#region selfattn
|
||||
class WanSelfAttention(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
@@ -883,29 +894,12 @@ WAN_CROSSATTENTION_CLASSES = {
|
||||
class WanAttentionBlock(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
cross_attn_type,
|
||||
in_features,
|
||||
out_features,
|
||||
ffn_dim,
|
||||
ffn2_dim,
|
||||
num_heads,
|
||||
qk_norm=True,
|
||||
cross_attn_norm=False,
|
||||
eps=1e-6,
|
||||
attention_mode="sdpa",
|
||||
rope_func="comfy",
|
||||
rms_norm_function="default",
|
||||
use_motion_attn=False,
|
||||
use_humo_audio_attn=False,
|
||||
face_fuser_block=False,
|
||||
lynx_ip_layers=None,
|
||||
lynx_ref_layers=None,
|
||||
block_idx=0,
|
||||
# long cat
|
||||
is_longcat = False,
|
||||
):
|
||||
cross_attn_type, in_features, out_features, ffn_dim, ffn2_dim, num_heads,
|
||||
qk_norm=True, cross_attn_norm=False, eps=1e-6, attention_mode="sdpa", rope_func="comfy", rms_norm_function="default",
|
||||
use_motion_attn=False, use_humo_audio_attn=False, face_fuser_block=False, lynx_ip_layers=None, lynx_ref_layers=None,
|
||||
block_idx=0, is_longcat=False):
|
||||
super().__init__()
|
||||
self.dim = min(out_features, in_features)
|
||||
self.dim = out_features
|
||||
self.ffn_dim = ffn_dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = out_features // num_heads
|
||||
|
||||
Reference in New Issue
Block a user