update
This commit is contained in:
+155
-115
@@ -1,85 +1,13 @@
|
|||||||
{
|
{
|
||||||
"last_node_id": 6,
|
"last_node_id": 35,
|
||||||
"last_link_id": 10,
|
"last_link_id": 34,
|
||||||
"nodes": [
|
"nodes": [
|
||||||
{
|
{
|
||||||
"id": 4,
|
"id": 29,
|
||||||
"type": "PreviewImage",
|
|
||||||
"pos": [
|
|
||||||
1401,
|
|
||||||
136
|
|
||||||
],
|
|
||||||
"size": {
|
|
||||||
"0": 593,
|
|
||||||
"1": 648
|
|
||||||
},
|
|
||||||
"flags": {},
|
|
||||||
"order": 3,
|
|
||||||
"mode": 0,
|
|
||||||
"inputs": [
|
|
||||||
{
|
|
||||||
"name": "images",
|
|
||||||
"type": "IMAGE",
|
|
||||||
"link": 10
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"properties": {
|
|
||||||
"Node name for S&R": "PreviewImage"
|
|
||||||
}
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"id": 5,
|
|
||||||
"type": "ella_model_loader",
|
|
||||||
"pos": [
|
|
||||||
688,
|
|
||||||
144
|
|
||||||
],
|
|
||||||
"size": {
|
|
||||||
"0": 210,
|
|
||||||
"1": 66
|
|
||||||
},
|
|
||||||
"flags": {},
|
|
||||||
"order": 1,
|
|
||||||
"mode": 0,
|
|
||||||
"inputs": [
|
|
||||||
{
|
|
||||||
"name": "model",
|
|
||||||
"type": "MODEL",
|
|
||||||
"link": 5
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"name": "clip",
|
|
||||||
"type": "CLIP",
|
|
||||||
"link": 6,
|
|
||||||
"slot_index": 1
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"name": "vae",
|
|
||||||
"type": "VAE",
|
|
||||||
"link": 7
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"outputs": [
|
|
||||||
{
|
|
||||||
"name": "ella_model",
|
|
||||||
"type": "ELLAMODEL",
|
|
||||||
"links": [
|
|
||||||
9
|
|
||||||
],
|
|
||||||
"shape": 3,
|
|
||||||
"slot_index": 0
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"properties": {
|
|
||||||
"Node name for S&R": "ella_model_loader"
|
|
||||||
}
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"id": 3,
|
|
||||||
"type": "CheckpointLoaderSimple",
|
"type": "CheckpointLoaderSimple",
|
||||||
"pos": [
|
"pos": [
|
||||||
310,
|
289,
|
||||||
142
|
315
|
||||||
],
|
],
|
||||||
"size": {
|
"size": {
|
||||||
"0": 315,
|
"0": 315,
|
||||||
@@ -93,7 +21,7 @@
|
|||||||
"name": "MODEL",
|
"name": "MODEL",
|
||||||
"type": "MODEL",
|
"type": "MODEL",
|
||||||
"links": [
|
"links": [
|
||||||
5
|
20
|
||||||
],
|
],
|
||||||
"shape": 3,
|
"shape": 3,
|
||||||
"slot_index": 0
|
"slot_index": 0
|
||||||
@@ -102,7 +30,7 @@
|
|||||||
"name": "CLIP",
|
"name": "CLIP",
|
||||||
"type": "CLIP",
|
"type": "CLIP",
|
||||||
"links": [
|
"links": [
|
||||||
6
|
21
|
||||||
],
|
],
|
||||||
"shape": 3
|
"shape": 3
|
||||||
},
|
},
|
||||||
@@ -110,7 +38,7 @@
|
|||||||
"name": "VAE",
|
"name": "VAE",
|
||||||
"type": "VAE",
|
"type": "VAE",
|
||||||
"links": [
|
"links": [
|
||||||
7
|
22
|
||||||
],
|
],
|
||||||
"shape": 3,
|
"shape": 3,
|
||||||
"slot_index": 2
|
"slot_index": 2
|
||||||
@@ -120,28 +48,93 @@
|
|||||||
"Node name for S&R": "CheckpointLoaderSimple"
|
"Node name for S&R": "CheckpointLoaderSimple"
|
||||||
},
|
},
|
||||||
"widgets_values": [
|
"widgets_values": [
|
||||||
"1_5/v1-5-pruned.ckpt"
|
"1_5\\photon_v1.safetensors"
|
||||||
]
|
]
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
"id": 6,
|
"id": 30,
|
||||||
|
"type": "PreviewImage",
|
||||||
|
"pos": [
|
||||||
|
1316,
|
||||||
|
307
|
||||||
|
],
|
||||||
|
"size": {
|
||||||
|
"0": 590.6172485351562,
|
||||||
|
"1": 614.5595092773438
|
||||||
|
},
|
||||||
|
"flags": {},
|
||||||
|
"order": 4,
|
||||||
|
"mode": 0,
|
||||||
|
"inputs": [
|
||||||
|
{
|
||||||
|
"name": "images",
|
||||||
|
"type": "IMAGE",
|
||||||
|
"link": 34
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "PreviewImage"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 34,
|
||||||
|
"type": "ella_t5_embeds",
|
||||||
|
"pos": [
|
||||||
|
512,
|
||||||
|
484
|
||||||
|
],
|
||||||
|
"size": {
|
||||||
|
"0": 400,
|
||||||
|
"1": 200
|
||||||
|
},
|
||||||
|
"flags": {},
|
||||||
|
"order": 1,
|
||||||
|
"mode": 0,
|
||||||
|
"outputs": [
|
||||||
|
{
|
||||||
|
"name": "ella_embeds",
|
||||||
|
"type": "ELLAEMBEDS",
|
||||||
|
"links": [
|
||||||
|
33
|
||||||
|
],
|
||||||
|
"shape": 3
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "ella_t5_embeds"
|
||||||
|
},
|
||||||
|
"widgets_values": [
|
||||||
|
"A vivid red book with a smooth, matte cover lies next to a glossy yellow vase. The vase, with a slightly curved silhouette, stands on a dark wood table with a noticeable grain pattern. The book appears slightly worn at the edges, suggesting frequent use, while the vase holds a fresh array of multicolored wildflowers.",
|
||||||
|
4,
|
||||||
|
128,
|
||||||
|
false
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 35,
|
||||||
"type": "ella_sampler",
|
"type": "ella_sampler",
|
||||||
"pos": [
|
"pos": [
|
||||||
947,
|
962,
|
||||||
142
|
317
|
||||||
],
|
|
||||||
"size": [
|
|
||||||
415,
|
|
||||||
487
|
|
||||||
],
|
],
|
||||||
|
"size": {
|
||||||
|
"0": 315,
|
||||||
|
"1": 222
|
||||||
|
},
|
||||||
"flags": {},
|
"flags": {},
|
||||||
"order": 2,
|
"order": 3,
|
||||||
"mode": 0,
|
"mode": 0,
|
||||||
"inputs": [
|
"inputs": [
|
||||||
{
|
{
|
||||||
"name": "ella_model",
|
"name": "ella_model",
|
||||||
"type": "ELLAMODEL",
|
"type": "ELLAMODEL",
|
||||||
"link": 9
|
"link": 32
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "ella_embeds",
|
||||||
|
"type": "ELLAEMBEDS",
|
||||||
|
"link": 33,
|
||||||
|
"slot_index": 1
|
||||||
}
|
}
|
||||||
],
|
],
|
||||||
"outputs": [
|
"outputs": [
|
||||||
@@ -149,72 +142,119 @@
|
|||||||
"name": "images",
|
"name": "images",
|
||||||
"type": "IMAGE",
|
"type": "IMAGE",
|
||||||
"links": [
|
"links": [
|
||||||
10
|
34
|
||||||
],
|
],
|
||||||
"shape": 3,
|
"shape": 3,
|
||||||
"slot_index": 0
|
"slot_index": 0
|
||||||
},
|
|
||||||
{
|
|
||||||
"name": "last_image",
|
|
||||||
"type": "IMAGE",
|
|
||||||
"links": null,
|
|
||||||
"shape": 3
|
|
||||||
}
|
}
|
||||||
],
|
],
|
||||||
"properties": {
|
"properties": {
|
||||||
"Node name for S&R": "ella_sampler"
|
"Node name for S&R": "ella_sampler"
|
||||||
},
|
},
|
||||||
"widgets_values": [
|
"widgets_values": [
|
||||||
"A vivid red book with a smooth, matte cover lies next to a glossy yellow vase. The vase, with a slightly curved silhouette, stands on a dark wood table with a noticeable grain pattern. The book appears slightly worn at the edges, suggesting frequent use, while the vase holds a fresh array of multicolored wildflowers.\n",
|
|
||||||
512,
|
512,
|
||||||
512,
|
512,
|
||||||
1,
|
|
||||||
25,
|
25,
|
||||||
10,
|
10,
|
||||||
933038223352312,
|
915981713542918,
|
||||||
"randomize",
|
"randomize",
|
||||||
"DDPMScheduler"
|
"DDPMScheduler"
|
||||||
]
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 27,
|
||||||
|
"type": "ella_model_loader",
|
||||||
|
"pos": [
|
||||||
|
680,
|
||||||
|
317
|
||||||
|
],
|
||||||
|
"size": {
|
||||||
|
"0": 210,
|
||||||
|
"1": 66
|
||||||
|
},
|
||||||
|
"flags": {},
|
||||||
|
"order": 2,
|
||||||
|
"mode": 0,
|
||||||
|
"inputs": [
|
||||||
|
{
|
||||||
|
"name": "model",
|
||||||
|
"type": "MODEL",
|
||||||
|
"link": 20
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "clip",
|
||||||
|
"type": "CLIP",
|
||||||
|
"link": 21,
|
||||||
|
"slot_index": 1
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "vae",
|
||||||
|
"type": "VAE",
|
||||||
|
"link": 22
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"outputs": [
|
||||||
|
{
|
||||||
|
"name": "ella_model",
|
||||||
|
"type": "ELLAMODEL",
|
||||||
|
"links": [
|
||||||
|
32
|
||||||
|
],
|
||||||
|
"shape": 3,
|
||||||
|
"slot_index": 0
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "ella_model_loader"
|
||||||
|
}
|
||||||
}
|
}
|
||||||
],
|
],
|
||||||
"links": [
|
"links": [
|
||||||
[
|
[
|
||||||
5,
|
20,
|
||||||
3,
|
29,
|
||||||
0,
|
0,
|
||||||
5,
|
27,
|
||||||
0,
|
0,
|
||||||
"MODEL"
|
"MODEL"
|
||||||
],
|
],
|
||||||
[
|
[
|
||||||
6,
|
21,
|
||||||
3,
|
29,
|
||||||
1,
|
1,
|
||||||
5,
|
27,
|
||||||
1,
|
1,
|
||||||
"CLIP"
|
"CLIP"
|
||||||
],
|
],
|
||||||
[
|
[
|
||||||
7,
|
22,
|
||||||
3,
|
29,
|
||||||
2,
|
2,
|
||||||
5,
|
27,
|
||||||
2,
|
2,
|
||||||
"VAE"
|
"VAE"
|
||||||
],
|
],
|
||||||
[
|
[
|
||||||
9,
|
32,
|
||||||
5,
|
27,
|
||||||
0,
|
0,
|
||||||
6,
|
35,
|
||||||
0,
|
0,
|
||||||
"ELLAMODEL"
|
"ELLAMODEL"
|
||||||
],
|
],
|
||||||
[
|
[
|
||||||
10,
|
33,
|
||||||
6,
|
34,
|
||||||
0,
|
0,
|
||||||
4,
|
35,
|
||||||
|
1,
|
||||||
|
"ELLAEMBEDS"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
34,
|
||||||
|
35,
|
||||||
|
0,
|
||||||
|
30,
|
||||||
0,
|
0,
|
||||||
"IMAGE"
|
"IMAGE"
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -128,7 +128,7 @@ class PerceiverResampler(nn.Module):
|
|||||||
|
|
||||||
|
|
||||||
class T5TextEmbedder(nn.Module):
|
class T5TextEmbedder(nn.Module):
|
||||||
def __init__(self, pretrained_path="google/flan-t5-xl", max_length=None):
|
def __init__(self, pretrained_path="ybelkada/flan-t5-xl-sharded-bf16", max_length=None):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.model = T5EncoderModel.from_pretrained(pretrained_path)
|
self.model = T5EncoderModel.from_pretrained(pretrained_path)
|
||||||
self.tokenizer = T5Tokenizer.from_pretrained(pretrained_path)
|
self.tokenizer = T5Tokenizer.from_pretrained(pretrained_path)
|
||||||
|
|||||||
@@ -73,72 +73,6 @@ class ELLAProxyUNet(torch.nn.Module):
|
|||||||
encoder_attention_mask=encoder_attention_mask,
|
encoder_attention_mask=encoder_attention_mask,
|
||||||
return_dict=return_dict,
|
return_dict=return_dict,
|
||||||
)
|
)
|
||||||
def generate_image_with_flexible_max_length(
|
|
||||||
pipe, t5_encoder, prompt, fixed_negative=False, output_type="pt", **pipe_kwargs
|
|
||||||
):
|
|
||||||
device = pipe.device
|
|
||||||
dtype = pipe.dtype
|
|
||||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
|
||||||
batch_size = len(prompt)
|
|
||||||
|
|
||||||
prompt_embeds = t5_encoder(prompt, max_length=None).to(device, dtype)
|
|
||||||
negative_prompt_embeds = t5_encoder(
|
|
||||||
[""] * batch_size, max_length=128 if fixed_negative else None
|
|
||||||
).to(device, dtype)
|
|
||||||
|
|
||||||
# diffusers pipeline concatenate `prompt_embeds` too early...
|
|
||||||
# https://github.com/huggingface/diffusers/blob/b6d7e31d10df675d86c6fe7838044712c6dca4e9/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion.py#L913
|
|
||||||
pipe.unet.flexible_max_length_workaround = [
|
|
||||||
negative_prompt_embeds.size(1)
|
|
||||||
] * batch_size + [prompt_embeds.size(1)] * batch_size
|
|
||||||
|
|
||||||
max_length = max([prompt_embeds.size(1), negative_prompt_embeds.size(1)])
|
|
||||||
b, _, d = prompt_embeds.shape
|
|
||||||
prompt_embeds = torch.cat(
|
|
||||||
[
|
|
||||||
prompt_embeds,
|
|
||||||
torch.zeros(
|
|
||||||
(b, max_length - prompt_embeds.size(1), d), device=device, dtype=dtype
|
|
||||||
),
|
|
||||||
],
|
|
||||||
dim=1,
|
|
||||||
)
|
|
||||||
negative_prompt_embeds = torch.cat(
|
|
||||||
[
|
|
||||||
negative_prompt_embeds,
|
|
||||||
torch.zeros(
|
|
||||||
(b, max_length - negative_prompt_embeds.size(1), d),
|
|
||||||
device=device,
|
|
||||||
dtype=dtype,
|
|
||||||
),
|
|
||||||
],
|
|
||||||
dim=1,
|
|
||||||
)
|
|
||||||
|
|
||||||
images = pipe(
|
|
||||||
prompt_embeds=prompt_embeds,
|
|
||||||
negative_prompt_embeds=negative_prompt_embeds,
|
|
||||||
**pipe_kwargs,
|
|
||||||
output_type=output_type,
|
|
||||||
).images
|
|
||||||
pipe.unet.flexible_max_length_workaround = None
|
|
||||||
return images
|
|
||||||
|
|
||||||
|
|
||||||
def load_ella(filename, device, dtype):
|
|
||||||
ella = ELLA()
|
|
||||||
safetensors.torch.load_model(ella, filename, strict=True)
|
|
||||||
ella.to(device, dtype=dtype)
|
|
||||||
return ella
|
|
||||||
|
|
||||||
|
|
||||||
def load_ella_for_pipe(pipe, ella):
|
|
||||||
pipe.unet = ELLAProxyUNet(ella, pipe.unet)
|
|
||||||
|
|
||||||
|
|
||||||
def offload_ella_for_pipe(pipe):
|
|
||||||
pipe.unet = pipe.unet.unet
|
|
||||||
|
|
||||||
|
|
||||||
def generate_image_with_fixed_max_length(
|
def generate_image_with_fixed_max_length(
|
||||||
pipe, t5_encoder, prompt, output_type="pt", **pipe_kwargs
|
pipe, t5_encoder, prompt, output_type="pt", **pipe_kwargs
|
||||||
@@ -176,6 +110,7 @@ class ella_model_loader:
|
|||||||
def loadmodel(self, model, clip, vae):
|
def loadmodel(self, model, clip, vae):
|
||||||
mm.soft_empty_cache()
|
mm.soft_empty_cache()
|
||||||
dtype = mm.unet_dtype()
|
dtype = mm.unet_dtype()
|
||||||
|
vae_dtype = mm.vae_dtype()
|
||||||
device = mm.get_torch_device()
|
device = mm.get_torch_device()
|
||||||
|
|
||||||
custom_config = {
|
custom_config = {
|
||||||
@@ -231,24 +166,21 @@ class ella_model_loader:
|
|||||||
'beta_schedule': "linear",
|
'beta_schedule': "linear",
|
||||||
'steps_offset': 1
|
'steps_offset': 1
|
||||||
}
|
}
|
||||||
|
# 4. tokenizer
|
||||||
|
tokenizer_path = os.path.join(script_directory, "configs/tokenizer")
|
||||||
|
tokenizer = CLIPTokenizer.from_pretrained(tokenizer_path)
|
||||||
|
|
||||||
scheduler=DPMSolverMultistepScheduler(**scheduler_config)
|
scheduler=DPMSolverMultistepScheduler(**scheduler_config)
|
||||||
pbar.update(1)
|
pbar.update(1)
|
||||||
del sd
|
del sd
|
||||||
print("loading ELLA")
|
|
||||||
ella_path = os.path.join(script_directory, 'checkpoints', 'ella-sd1.5-tsc-t5xl.safetensors')
|
|
||||||
ella = ELLA()
|
|
||||||
safetensors.torch.load_model(ella, ella_path, strict=True)
|
|
||||||
|
|
||||||
ella.to(device, dtype=dtype)
|
ella.to(device, dtype=dtype)
|
||||||
unet = unet.to(device)
|
unet = unet.to(device)
|
||||||
ella_unet = ELLAProxyUNet(ella, unet)
|
ella_unet = ELLAProxyUNet(ella, unet)
|
||||||
pbar.update(1)
|
pbar.update(1)
|
||||||
print("loading tokenizer")
|
|
||||||
tokenizer_path = os.path.join(script_directory, "configs/tokenizer")
|
|
||||||
tokenizer = CLIPTokenizer.from_pretrained(tokenizer_path)
|
|
||||||
print("creating pipeline")
|
print("creating pipeline")
|
||||||
pipe = StableDiffusionPipeline(
|
self.pipe = StableDiffusionPipeline(
|
||||||
unet=unet,
|
unet=unet,
|
||||||
vae=vae,
|
vae=vae,
|
||||||
text_encoder=text_encoder,
|
text_encoder=text_encoder,
|
||||||
@@ -261,12 +193,10 @@ class ella_model_loader:
|
|||||||
)
|
)
|
||||||
print("pipeline created")
|
print("pipeline created")
|
||||||
pbar.update(1)
|
pbar.update(1)
|
||||||
pipe.unet = ella_unet
|
self.pipe.unet = ella_unet
|
||||||
t5_encoder = T5TextEmbedder().to(pipe.device, dtype=dtype)
|
|
||||||
ella_model = {
|
ella_model = {
|
||||||
'pipe': pipe,
|
'pipe': self.pipe,
|
||||||
'ella': ella,
|
|
||||||
't5_encoder': t5_encoder
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return (ella_model,)
|
return (ella_model,)
|
||||||
@@ -276,10 +206,9 @@ class ella_sampler:
|
|||||||
def INPUT_TYPES(s):
|
def INPUT_TYPES(s):
|
||||||
return {"required": {
|
return {"required": {
|
||||||
"ella_model": ("ELLAMODEL",),
|
"ella_model": ("ELLAMODEL",),
|
||||||
"prompt": ("STRING", {"multiline": True, "default": "A vivid red book with a smooth, matte cover lies next to a glossy yellow vase. The vase, with a slightly curved silhouette, stands on a dark wood table with a noticeable grain pattern. The book appears slightly worn at the edges, suggesting frequent use, while the vase holds a fresh array of multicolored wildflowers.",}),
|
"ella_embeds": ("ELLAEMBEDS",),
|
||||||
"width": ("INT", {"default": 512, "min": 64, "max": 2048, "step": 64}),
|
"width": ("INT", {"default": 512, "min": 64, "max": 2048, "step": 64}),
|
||||||
"height": ("INT", {"default": 512, "min": 64, "max": 2048, "step": 64}),
|
"height": ("INT", {"default": 512, "min": 64, "max": 2048, "step": 64}),
|
||||||
"batch_size": ("INT", {"default": 1, "min": 1, "max": 256, "step": 1}),
|
|
||||||
"steps": ("INT", {"default": 25, "min": 1, "max": 200, "step": 1}),
|
"steps": ("INT", {"default": 25, "min": 1, "max": 200, "step": 1}),
|
||||||
"guidance_scale": ("FLOAT", {"default": 10.0, "min": 0.0, "max": 20.0, "step": 0.01}),
|
"guidance_scale": ("FLOAT", {"default": 10.0, "min": 0.0, "max": 20.0, "step": 0.01}),
|
||||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||||
@@ -291,7 +220,7 @@ class ella_sampler:
|
|||||||
'PNDMScheduler',
|
'PNDMScheduler',
|
||||||
'DEISMultistepScheduler'
|
'DEISMultistepScheduler'
|
||||||
], {
|
], {
|
||||||
"default": 'DDIMScheduler'
|
"default": 'DDPMScheduler'
|
||||||
}),
|
}),
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
@@ -299,14 +228,13 @@ class ella_sampler:
|
|||||||
RETURN_TYPES = ("IMAGE",)
|
RETURN_TYPES = ("IMAGE",)
|
||||||
RETURN_NAMES = ("images",)
|
RETURN_NAMES = ("images",)
|
||||||
FUNCTION = "process"
|
FUNCTION = "process"
|
||||||
CATEGORY = "champWrapper"
|
CATEGORY = "ELLA-Wrapper"
|
||||||
|
|
||||||
def process(self, prompt, batch_size, width, height, steps, guidance_scale, seed, ella_model, scheduler):
|
def process(self, ella_embeds, width, height, steps, guidance_scale, seed, ella_model, scheduler):
|
||||||
device = mm.get_torch_device()
|
device = mm.get_torch_device()
|
||||||
mm.unload_all_models()
|
mm.unload_all_models()
|
||||||
mm.soft_empty_cache()
|
mm.soft_empty_cache()
|
||||||
dtype = mm.unet_dtype()
|
dtype = mm.unet_dtype()
|
||||||
t5_encoder=ella_model['t5_encoder']
|
|
||||||
pipe=ella_model['pipe']
|
pipe=ella_model['pipe']
|
||||||
pipe.to(device, dtype=dtype)
|
pipe.to(device, dtype=dtype)
|
||||||
|
|
||||||
@@ -331,20 +259,64 @@ class ella_sampler:
|
|||||||
|
|
||||||
autocast_condition = (dtype != torch.float32) and not mm.is_device_mps(device)
|
autocast_condition = (dtype != torch.float32) and not mm.is_device_mps(device)
|
||||||
with torch.autocast(mm.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext():
|
with torch.autocast(mm.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext():
|
||||||
|
|
||||||
|
# diffusers pipeline concatenate `prompt_embeds` too early...
|
||||||
|
# https://github.com/huggingface/diffusers/blob/b6d7e31d10df675d86c6fe7838044712c6dca4e9/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion.py#L913
|
||||||
|
pipe.unet.flexible_max_length_workaround = [ella_embeds["negative_prompt_embeds"].size(1)] * ella_embeds["batch_size"] + [ella_embeds["prompt_embeds"].size(1)] * ella_embeds["batch_size"]
|
||||||
|
|
||||||
|
images = pipe(
|
||||||
|
prompt_embeds=ella_embeds["prompt_embeds"],
|
||||||
|
negative_prompt_embeds=ella_embeds["negative_prompt_embeds"],
|
||||||
|
guidance_scale=guidance_scale,
|
||||||
|
num_inference_steps=steps,
|
||||||
|
height=height,
|
||||||
|
width=width,
|
||||||
|
generator=[
|
||||||
|
torch.Generator(device=device).manual_seed(seed + i)
|
||||||
|
for i in range(ella_embeds["batch_size"])
|
||||||
|
],
|
||||||
|
output_type="np.array",
|
||||||
|
).images
|
||||||
|
|
||||||
|
image_out = torch.from_numpy(images).cpu().float()
|
||||||
|
|
||||||
|
return (image_out,)
|
||||||
|
|
||||||
|
class ella_t5_embeds:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {"required": {
|
||||||
|
"prompt": ("STRING", {"multiline": True, "default": "A vivid red book with a smooth, matte cover lies next to a glossy yellow vase. The vase, with a slightly curved silhouette, stands on a dark wood table with a noticeable grain pattern. The book appears slightly worn at the edges, suggesting frequent use, while the vase holds a fresh array of multicolored wildflowers.",}),
|
||||||
|
"batch_size": ("INT", {"default": 1, "min": 1, "max": 256, "step": 1}),
|
||||||
|
"max_length": ("INT", {"default": 128, "min": 1, "max": 256, "step": 1}),
|
||||||
|
"fixed_negative": ("BOOLEAN", {"default": False}),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("ELLAEMBEDS",)
|
||||||
|
RETURN_NAMES = ("ella_embeds",)
|
||||||
|
FUNCTION = "process"
|
||||||
|
CATEGORY = "ELLA-Wrapper"
|
||||||
|
|
||||||
|
def process(self, prompt, batch_size, max_length, fixed_negative):
|
||||||
|
device = mm.get_torch_device()
|
||||||
|
mm.unload_all_models()
|
||||||
|
mm.soft_empty_cache()
|
||||||
|
dtype = mm.unet_dtype()
|
||||||
|
t5_encoder = T5TextEmbedder().to(device, dtype=dtype)
|
||||||
|
|
||||||
|
autocast_condition = (dtype != torch.float32) and not mm.is_device_mps(device)
|
||||||
|
with torch.autocast(mm.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext():
|
||||||
|
print("generating embeds")
|
||||||
|
prompt = [prompt] * batch_size
|
||||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||||
batch_size = len(prompt)
|
#batch_size = len(prompt)
|
||||||
|
|
||||||
fixed_negative = False
|
|
||||||
prompt_embeds = t5_encoder(prompt, max_length=None).to(device, dtype)
|
prompt_embeds = t5_encoder(prompt, max_length=None).to(device, dtype)
|
||||||
negative_prompt_embeds = t5_encoder(
|
negative_prompt_embeds = t5_encoder(
|
||||||
[""] * batch_size, max_length=128 if fixed_negative else None
|
[""] * batch_size, max_length=max_length if fixed_negative else None
|
||||||
).to(device, dtype)
|
).to(device, dtype)
|
||||||
|
|
||||||
pipe.unet.flexible_max_length_workaround = [
|
|
||||||
negative_prompt_embeds.size(1)
|
|
||||||
] * batch_size + [prompt_embeds.size(1)] * batch_size
|
|
||||||
|
|
||||||
max_length = max([prompt_embeds.size(1), negative_prompt_embeds.size(1)])
|
max_length = max([prompt_embeds.size(1), negative_prompt_embeds.size(1)])
|
||||||
b, _, d = prompt_embeds.shape
|
b, _, d = prompt_embeds.shape
|
||||||
prompt_embeds = torch.cat(
|
prompt_embeds = torch.cat(
|
||||||
@@ -367,31 +339,20 @@ class ella_sampler:
|
|||||||
],
|
],
|
||||||
dim=1,
|
dim=1,
|
||||||
)
|
)
|
||||||
|
embeds = {
|
||||||
images = pipe(
|
"prompt_embeds": prompt_embeds,
|
||||||
prompt_embeds=prompt_embeds,
|
"negative_prompt_embeds": negative_prompt_embeds,
|
||||||
negative_prompt_embeds=negative_prompt_embeds,
|
"batch_size": batch_size
|
||||||
guidance_scale=guidance_scale,
|
}
|
||||||
num_inference_steps=steps,
|
return (embeds,)
|
||||||
height=height,
|
|
||||||
width=width,
|
|
||||||
generator=[
|
|
||||||
torch.Generator(device=device).manual_seed(seed + i)
|
|
||||||
for i in range(batch_size)
|
|
||||||
],
|
|
||||||
output_type="np.array",
|
|
||||||
).images
|
|
||||||
|
|
||||||
tensor = torch.from_numpy(images).cpu().float()
|
|
||||||
|
|
||||||
return (tensor,)
|
|
||||||
|
|
||||||
|
|
||||||
NODE_CLASS_MAPPINGS = {
|
NODE_CLASS_MAPPINGS = {
|
||||||
"ella_model_loader": ella_model_loader,
|
"ella_model_loader": ella_model_loader,
|
||||||
"ella_sampler": ella_sampler,
|
"ella_sampler": ella_sampler,
|
||||||
|
"ella_t5_embeds": ella_t5_embeds
|
||||||
}
|
}
|
||||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||||
"ella_model_loader": "ELLA Model Loader",
|
"ella_model_loader": "ELLA Model Loader",
|
||||||
"ella_sampler": "ELLA Sampler",
|
"ella_sampler": "ELLA Sampler",
|
||||||
|
"ella_t5_embeds": "ELLA T5 Embeds"
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user