118 lines
5.2 KiB
Python
118 lines
5.2 KiB
Python
import torch
|
|
import math
|
|
import types
|
|
|
|
EPSILON = 1e-4
|
|
|
|
SD_layer_dims = {
|
|
"SD1" : {"input_1": 4096,"input_2": 4096,"input_4": 1024,"input_5": 1024,"input_7": 256,"input_8": 256,"middle_0": 64,"output_3": 256,"output_4": 256,"output_5": 256,"output_6": 1024,"output_7": 1024,"output_8": 1024,"output_9": 4096,"output_10": 4096,"output_11": 4096},
|
|
"SDXL": {"input_4": 4096,"input_5": 4096,"input_7": 1024,"input_8": 1024,"middle_0": 1024,"output_0": 1024,"output_1": 1024,"output_2": 1024,"output_3": 4096,"output_4": 4096,"output_5": 4096},
|
|
"Disabled":{}
|
|
}
|
|
|
|
def should_scale(mname,lname,q2):
|
|
if mname != "Disabled" and lname in SD_layer_dims[mname]:
|
|
return q2 != SD_layer_dims[mname][lname]
|
|
return False
|
|
|
|
class temperature_patcher():
|
|
def __init__(self, temperature, layer_name = "", model_name=""):
|
|
self.temperature = max(temperature,EPSILON)
|
|
self.layer_name = layer_name
|
|
self.model_name = model_name
|
|
|
|
def pytorch_attention_with_temperature(self, q, k, v, extra_options, mask=None, attn_precision=None):
|
|
heads = extra_options if isinstance(extra_options, int) else extra_options['n_heads']
|
|
b, _, dim_head = q.shape
|
|
dim_head //= heads
|
|
q, k, v = map(
|
|
lambda t: t.view(b, -1, heads, dim_head).transpose(1, 2),
|
|
(q, k, v),
|
|
)
|
|
|
|
scale = 1 / (math.sqrt(q.size(-1)) * self.temperature)
|
|
|
|
out = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=mask, dropout_p=0.0, is_causal=False, scale=scale)
|
|
|
|
if should_scale(self.model_name, self.layer_name,q.size(-2)):
|
|
out *= math.log(q.size(-2) ** 0.5) / math.log(SD_layer_dims[self.model_name][self.layer_name] ** 0.5)
|
|
|
|
out = (
|
|
out.transpose(1, 2).reshape(b, -1, heads * dim_head)
|
|
)
|
|
|
|
return out
|
|
|
|
class UnetTemperaturePatch:
|
|
@classmethod
|
|
def INPUT_TYPES(s):
|
|
required_inputs = {}
|
|
required_inputs["model"] = ("MODEL",)
|
|
required_inputs["Temperature"] = ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "round": 0.01})
|
|
required_inputs["Attention"] = (["self","cross","both"],)
|
|
required_inputs["Dynamic_Scale_Attention"] = (["Disabled","SDXL","SD1"],)
|
|
return {"required": required_inputs}
|
|
|
|
TOGGLES = {}
|
|
RETURN_TYPES = ("MODEL","STRING",)
|
|
RETURN_NAMES = ("Model","String",)
|
|
FUNCTION = "patch"
|
|
|
|
CATEGORY = "model_patches/Temperature"
|
|
|
|
def patch(self, model, Temperature, Attention, Dynamic_Scale_Attention, **kwargs):
|
|
m = model.clone()
|
|
levels = ["input","middle","output"]
|
|
layer_names = {f"{l}_{n}": True for l in levels for n in range(12)}
|
|
|
|
for key, toggle in layer_names.items():
|
|
current_level = key.split("_")[0]
|
|
b_number = int(key.split("_")[1])
|
|
|
|
if Attention in ["both","self"]:
|
|
patcher = temperature_patcher(Temperature,layer_name=key,model_name=Dynamic_Scale_Attention)
|
|
m.set_model_attn1_replace(patcher.pytorch_attention_with_temperature, current_level, b_number)
|
|
if Attention in ["both","cross"]:
|
|
patcher = temperature_patcher(Temperature,layer_name=key,model_name=Dynamic_Scale_Attention)
|
|
m.set_model_attn2_replace(patcher.pytorch_attention_with_temperature, current_level, b_number)
|
|
|
|
parameters_as_string = f"Temperature: {Temperature}\nAttention: {Attention}\nDynamic scale: {Dynamic_Scale_Attention}"
|
|
return (m, parameters_as_string,)
|
|
|
|
class CLIPTemperaturePatch:
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
return {"required": { "clip": ("CLIP",),
|
|
"Temperature": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}),
|
|
}}
|
|
|
|
RETURN_TYPES = ("CLIP",)
|
|
FUNCTION = "patch"
|
|
CATEGORY = "model_patches/Temperature"
|
|
|
|
def patch(self, clip, Temperature):
|
|
def custom_optimized_attention(device, mask=None, small_input=True):
|
|
return temperature_patcher(Temperature).pytorch_attention_with_temperature
|
|
|
|
def new_forward(self, x, mask=None, intermediate_output=None):
|
|
optimized_attention = custom_optimized_attention(x.device, mask=mask is not None, small_input=True)
|
|
|
|
if intermediate_output is not None:
|
|
if intermediate_output < 0:
|
|
intermediate_output = len(self.layers) + intermediate_output
|
|
|
|
intermediate = None
|
|
for i, l in enumerate(self.layers):
|
|
x = l(x, mask, optimized_attention)
|
|
if i == intermediate_output:
|
|
intermediate = x.clone()
|
|
return x, intermediate
|
|
|
|
clip_encoder_instance = clip.cond_stage_model.clip_l.transformer.text_model.encoder
|
|
clip_encoder_instance.forward = types.MethodType(new_forward, clip_encoder_instance)
|
|
|
|
if getattr(clip.cond_stage_model, f"clip_g", None) is not None:
|
|
clip_encoder_instance_g = clip.cond_stage_model.clip_g.transformer.text_model.encoder
|
|
clip_encoder_instance_g.forward = types.MethodType(new_forward, clip_encoder_instance_g)
|
|
|
|
return (clip,) |