Add llama2 support

This commit is contained in:
Kijai
2024-04-11 17:21:29 +03:00
parent 3721278883
commit 6f3bdb8eb6
2 changed files with 865 additions and 121 deletions
+773 -105
View File
@@ -1,66 +1,20 @@
{
"last_node_id": 5,
"last_link_id": 5,
"last_node_id": 23,
"last_link_id": 34,
"nodes": [
{
"id": 3,
"type": "CheckpointLoaderSimple",
"pos": [
373,
314
],
"size": [
320.97000488281265,
98
],
"flags": {},
"order": 0,
"mode": 0,
"outputs": [
{
"name": "MODEL",
"type": "MODEL",
"links": [
2
],
"shape": 3
},
{
"name": "CLIP",
"type": "CLIP",
"links": null,
"shape": 3
},
{
"name": "VAE",
"type": "VAE",
"links": [
3
],
"shape": 3,
"slot_index": 2
}
],
"properties": {
"Node name for S&R": "CheckpointLoaderSimple"
},
"widgets_values": [
"1_5/dreamshaper_8.safetensors"
]
},
{
"id": 2,
"type": "lavibridge_model_loader",
"pos": [
763,
322
741,
313
],
"size": {
"0": 210,
"1": 46
"1": 78
},
"flags": {},
"order": 2,
"order": 9,
"mode": 0,
"inputs": [
{
@@ -80,98 +34,580 @@
"name": "lavibridge",
"type": "LAVIBRIDGE",
"links": [
1
6
],
"shape": 3
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "lavibridge_model_loader"
}
},
"widgets_values": [
"t5_unet"
]
},
{
"id": 4,
"type": "PreviewImage",
"id": 15,
"type": "VAEDecode",
"pos": [
1388,
321
],
"size": [
558.763644131747,
569.452726537531
1066,
1199
],
"size": {
"0": 210,
"1": 46
},
"flags": {},
"order": 4,
"order": 14,
"mode": 0,
"inputs": [
{
"name": "images",
"name": "samples",
"type": "LATENT",
"link": 19
},
{
"name": "vae",
"type": "VAE",
"link": 20,
"slot_index": 1
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"link": 4
"links": [
21
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "PreviewImage"
"Node name for S&R": "VAEDecode"
}
},
{
"id": 5,
"type": "lavi_bridge_t5_encoder",
"id": 19,
"type": "StringConstantMultiline",
"pos": [
376,
484
-57,
643
],
"size": {
"0": 400,
"1": 200
},
"flags": {},
"order": 0,
"mode": 0,
"outputs": [
{
"name": "STRING",
"type": "STRING",
"links": [
24,
25
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "StringConstantMultiline"
},
"widgets_values": [
"Oppenheimer sits on the beach on a chair, watching a nuclear exposition with a huge mushroom cloud, 120mm, best quality, extremely detailed, 4k resolution",
true
]
},
{
"id": 3,
"type": "CheckpointLoaderSimple",
"pos": [
257,
311
],
"size": [
322.0363714044744,
200.2709083557129
405.99999237060547,
98
],
"flags": {},
"order": 1,
"mode": 0,
"outputs": [
{
"name": "t5_embeds",
"type": "T5EMBEDS",
"name": "MODEL",
"type": "MODEL",
"links": [
5
2,
22
],
"shape": 3
},
{
"name": "CLIP",
"type": "CLIP",
"links": [
17,
18
],
"shape": 3,
"slot_index": 1
},
{
"name": "VAE",
"type": "VAE",
"links": [
3,
20
],
"shape": 3,
"slot_index": 2
}
],
"properties": {
"Node name for S&R": "CheckpointLoaderSimple"
},
"widgets_values": [
"1_5/photon_v1.safetensors"
]
},
{
"id": 14,
"type": "CLIPTextEncode",
"pos": [
520,
1100
],
"size": [
335.39999923706057,
117.79999084472661
],
"flags": {},
"order": 8,
"mode": 0,
"inputs": [
{
"name": "clip",
"type": "CLIP",
"link": 18
}
],
"outputs": [
{
"name": "CONDITIONING",
"type": "CONDITIONING",
"links": [
16
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "CLIPTextEncode"
},
"widgets_values": [
"bad quality"
]
},
{
"id": 13,
"type": "CLIPTextEncode",
"pos": [
510,
990
],
"size": [
339.99999237060547,
54
],
"flags": {},
"order": 7,
"mode": 0,
"inputs": [
{
"name": "clip",
"type": "CLIP",
"link": 17
},
{
"name": "text",
"type": "STRING",
"link": 25,
"widget": {
"name": "text"
}
}
],
"outputs": [
{
"name": "CONDITIONING",
"type": "CONDITIONING",
"links": [
15
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "CLIPTextEncode"
},
"widgets_values": [
""
]
},
{
"id": 18,
"type": "EmptyLatentImage",
"pos": [
1010,
1050
],
"size": [
303.4000007629395,
74
],
"flags": {},
"order": 10,
"mode": 0,
"inputs": [
{
"name": "width",
"type": "INT",
"link": 28,
"widget": {
"name": "width"
}
},
{
"name": "height",
"type": "INT",
"link": 29,
"widget": {
"name": "height"
}
},
{
"name": "batch_size",
"type": "INT",
"link": 30,
"widget": {
"name": "batch_size"
}
}
],
"outputs": [
{
"name": "LATENT",
"type": "LATENT",
"links": [
23
],
"shape": 3
}
],
"properties": {
"Node name for S&R": "lavi_bridge_t5_encoder"
"Node name for S&R": "EmptyLatentImage"
},
"widgets_values": [
"Oppenheimer sits on the beach on a chair, watching a nuclear exposition with a huge mushroom cloud, 120mm, best quality, masterpiece, extremely detailed, 4k resolution",
77
512,
512,
4
]
},
{
"id": 1,
"type": "lavibridge_sampler",
"id": 16,
"type": "PreviewImage",
"pos": [
1021,
320
1465,
889
],
"size": {
"0": 315,
"1": 246
"0": 558.763671875,
"1": 569.4526977539062
},
"flags": {},
"order": 15,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 21
}
],
"title": "Preview Image: t5-large",
"properties": {
"Node name for S&R": "PreviewImage"
}
},
{
"id": 4,
"type": "PreviewImage",
"pos": [
1465,
260
],
"size": {
"0": 558.763671875,
"1": 569.4526977539062
},
"flags": {},
"order": 13,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 7
}
],
"title": "Preview Image: t5-large",
"properties": {
"Node name for S&R": "PreviewImage"
}
},
{
"id": 20,
"type": "INTConstant",
"pos": [
637,
607
],
"size": [
200,
58
],
"flags": {},
"order": 2,
"mode": 0,
"outputs": [
{
"name": "value",
"type": "INT",
"links": [
26,
28
],
"shape": 3,
"slot_index": 0
}
],
"title": "Width",
"properties": {
"Node name for S&R": "INTConstant"
},
"widgets_values": [
512
],
"color": "#1b4669",
"bgcolor": "#29699c"
},
{
"id": 21,
"type": "INTConstant",
"pos": [
640,
713
],
"size": [
200,
58
],
"flags": {},
"order": 3,
"mode": 0,
"outputs": [
{
"name": "value",
"type": "INT",
"links": [
27,
29
],
"shape": 3,
"slot_index": 0
}
],
"title": "Height",
"properties": {
"Node name for S&R": "INTConstant"
},
"widgets_values": [
512
],
"color": "#1b4669",
"bgcolor": "#29699c"
},
{
"id": 22,
"type": "INTConstant",
"pos": [
642,
826
],
"size": [
200,
58
],
"flags": {},
"order": 4,
"mode": 0,
"outputs": [
{
"name": "value",
"type": "INT",
"links": [
30,
31
],
"shape": 3,
"slot_index": 0
}
],
"title": "Batch_size",
"properties": {
"Node name for S&R": "INTConstant"
},
"widgets_values": [
4
],
"color": "#1b4669",
"bgcolor": "#29699c"
},
{
"id": 12,
"type": "KSampler",
"pos": [
990,
745
],
"size": [
356.99999237060547,
262
],
"flags": {},
"order": 12,
"mode": 0,
"inputs": [
{
"name": "model",
"type": "MODEL",
"link": 22,
"slot_index": 0
},
{
"name": "positive",
"type": "CONDITIONING",
"link": 15
},
{
"name": "negative",
"type": "CONDITIONING",
"link": 16
},
{
"name": "latent_image",
"type": "LATENT",
"link": 23,
"slot_index": 3
},
{
"name": "seed",
"type": "INT",
"link": 34,
"widget": {
"name": "seed"
}
}
],
"outputs": [
{
"name": "LATENT",
"type": "LATENT",
"links": [
19
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "KSampler"
},
"widgets_values": [
124,
"fixed",
25,
7.5,
"uni_pc",
"normal",
1
]
},
{
"id": 6,
"type": "lavibridge_sampler",
"pos": [
999,
295
],
"size": [
373.99999237060547,
246
],
"flags": {},
"order": 11,
"mode": 0,
"inputs": [
{
"name": "lavibridge_model",
"type": "LAVIBRIDGE",
"link": 1,
"slot_index": 0
"link": 6
},
{
"name": "t5_embeds",
"type": "T5EMBEDS",
"link": 5,
"name": "lavi_embeds",
"type": "LAVIEMBEDS",
"link": 13,
"slot_index": 1
},
{
"name": "width",
"type": "INT",
"link": 26,
"widget": {
"name": "width"
}
},
{
"name": "height",
"type": "INT",
"link": 27,
"widget": {
"name": "height"
}
},
{
"name": "batch_size",
"type": "INT",
"link": 31,
"widget": {
"name": "batch_size"
}
},
{
"name": "seed",
"type": "INT",
"link": 33,
"widget": {
"name": "seed"
},
"slot_index": 5
}
],
"outputs": [
@@ -179,7 +615,7 @@
"name": "images",
"type": "IMAGE",
"links": [
4
7
],
"shape": 3,
"slot_index": 0
@@ -194,21 +630,93 @@
4,
25,
7.5,
0,
124,
"fixed",
"UniPCMultistepScheduler"
]
},
{
"id": 23,
"type": "PrimitiveNode",
"pos": [
991,
607
],
"size": [
272.70000152587886,
82.99999084472654
],
"flags": {},
"order": 5,
"mode": 0,
"outputs": [
{
"name": "INT",
"type": "INT",
"links": [
33,
34
],
"widget": {
"name": "seed"
},
"slot_index": 0
}
],
"title": "seed",
"properties": {
"Run widget replace on values": false
},
"widgets_values": [
124,
"fixed"
]
},
{
"id": 7,
"type": "lavi_bridge_t5_encoder",
"pos": [
273,
458
],
"size": [
360.99999237060547,
70
],
"flags": {},
"order": 6,
"mode": 0,
"inputs": [
{
"name": "prompt",
"type": "STRING",
"link": 24,
"widget": {
"name": "prompt"
}
}
],
"outputs": [
{
"name": "lavi_embeds",
"type": "LAVIEMBEDS",
"links": [
13
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "lavi_bridge_t5_encoder"
},
"widgets_values": [
"Oppenheimer sits on the beach on a chair, watching a nuclear exposition with a huge mushroom cloud, 120mm, best quality, extremely detailed, 4k resolution",
77
]
}
],
"links": [
[
1,
2,
0,
1,
0,
"LAVIBRIDGE"
],
[
2,
3,
@@ -226,20 +734,180 @@
"VAE"
],
[
4,
1,
6,
2,
0,
6,
0,
"LAVIBRIDGE"
],
[
7,
6,
0,
4,
0,
"IMAGE"
],
[
5,
5,
13,
7,
0,
6,
1,
"LAVIEMBEDS"
],
[
15,
13,
0,
12,
1,
"T5EMBEDS"
"CONDITIONING"
],
[
16,
14,
0,
12,
2,
"CONDITIONING"
],
[
17,
3,
1,
13,
0,
"CLIP"
],
[
18,
3,
1,
14,
0,
"CLIP"
],
[
19,
12,
0,
15,
0,
"LATENT"
],
[
20,
3,
2,
15,
1,
"VAE"
],
[
21,
15,
0,
16,
0,
"IMAGE"
],
[
22,
3,
0,
12,
0,
"MODEL"
],
[
23,
18,
0,
12,
3,
"LATENT"
],
[
24,
19,
0,
7,
0,
"STRING"
],
[
25,
19,
0,
13,
1,
"STRING"
],
[
26,
20,
0,
6,
2,
"INT"
],
[
27,
21,
0,
6,
3,
"INT"
],
[
28,
20,
0,
18,
0,
"INT"
],
[
29,
21,
0,
18,
1,
"INT"
],
[
30,
22,
0,
18,
2,
"INT"
],
[
31,
22,
0,
6,
4,
"INT"
],
[
33,
23,
0,
6,
5,
"INT"
],
[
34,
23,
0,
12,
4,
"INT"
]
],
"groups": [],
+92 -16
View File
@@ -1,4 +1,3 @@
import os
from tqdm.auto import tqdm
@@ -44,6 +43,14 @@ class lavibridge_model_loader:
return {"required": {
"model": ("MODEL",),
"vae": ("VAE",),
"lora_type": (
[
'llama2_unet',
't5_unet',
], {
"default": 't5_unet'
}),
},
}
@@ -52,7 +59,7 @@ class lavibridge_model_loader:
FUNCTION = "loadmodel"
CATEGORY = "LaVI-BridgeWrapper"
def loadmodel(self, model, vae):
def loadmodel(self, model, vae, lora_type):
mm.soft_empty_cache()
dtype = mm.unet_dtype()
vae_dtype = mm.vae_dtype()
@@ -68,13 +75,13 @@ class lavibridge_model_loader:
# load models
lavibridge_folder = os.path.join(folder_paths.models_dir,'lavibridge')
lora_vis_path = os.path.join(lavibridge_folder, 't5_unet', 'lora_vis.pt')
lora_vis_path = os.path.join(lavibridge_folder, lora_type, 'lora_vis.pt')
if not os.path.exists(lora_vis_path):
print(f"Downloading LaVi-Bridge from https://huggingface.co/shihaozhao/LaVi-Bridge {lavibridge_folder}")
from huggingface_hub import snapshot_download
snapshot_download(repo_id="shihaozhao/LaVi-Bridge", allow_patterns=["*t5_unet*"],local_dir=lavibridge_folder, local_dir_use_symlinks=False)
snapshot_download(repo_id="shihaozhao/LaVi-Bridge", allow_patterns=[f"*{lora_type}*"],local_dir=lavibridge_folder, local_dir_use_symlinks=False)
print(f"Loaded LaVi-Bridge lora {lora_vis_path}")
pbar.update(1)
# get state dict from comfy models
@@ -119,17 +126,84 @@ class lavibridge_model_loader:
return (lavibridge_model,)
class lavi_bridge_t5_encoder:
class lavi_bridge_llama_encoder:
@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.",}),
"prompt": ("STRING", {"multiline": True, "default": "Oppenheimer sits on the beach on a chair, watching a nuclear exposition with a huge mushroom cloud, 120mm",}),
"max_length": ("INT", {"default": 77, "min": 1, "max": 512, "step": 1}),
},
}
RETURN_TYPES = ("T5EMBEDS",)
RETURN_NAMES = ("t5_embeds",)
RETURN_TYPES = ("LAVIEMBEDS",)
RETURN_NAMES = ("lavi_embeds",)
FUNCTION = "process"
CATEGORY = "LaVI-BridgeWrapper"
def process(self, prompt, max_length):
from transformers import LlamaForCausalLM, LlamaTokenizer
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
mm.soft_empty_cache()
dtype = mm.unet_dtype()
if not hasattr(self, "text_encoder"):
#llama2
llama2_path = os.path.join(folder_paths.models_dir,'llama2', 'Llama-2-7b-hf')
if not os.path.exists(llama2_path):
from huggingface_hub import snapshot_download
snapshot_download(repo_id="NousResearch/Llama-2-7b-hf", local_dir=llama2_path, ignore_patterns=["*.bin"], local_dir_use_symlinks=False)
#adapter
adapter_folder = os.path.join(folder_paths.models_dir,'lavibridge')
adapter_path = os.path.join(adapter_folder, 'llama2_unet','adapter')
if not os.path.exists(adapter_path):
print(f"Downloading LaVi-Bridge from https://huggingface.co/shihaozhao/LaVi-Bridge {adapter_folder}")
from huggingface_hub import snapshot_download
snapshot_download(repo_id="shihaozhao/LaVi-Bridge", allow_patterns=["*llama2_unet*"],local_dir=adapter_folder, local_dir_use_symlinks=False)
lora_text_path = os.path.join(adapter_folder, 'llama2_unet', 'lora_text.pt')
self.adapter = TextAdapter.from_pretrained(adapter_path).eval().to(dtype)
self.tokenizer = LlamaTokenizer.from_pretrained(llama2_path)
self.tokenizer.pad_token = '[PAD]'
self.text_encoder = LlamaForCausalLM.from_pretrained(llama2_path, torch_dtype=dtype)
monkeypatch_or_replace_lora_extended(
self.text_encoder,
torch.load(lora_text_path),
r=32,
target_replace_module = {"LlamaAttention"},
)
self.adapter.to(device)
self.text_encoder.to(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():
text_ids = self.tokenizer(prompt, padding="max_length", max_length=max_length, return_tensors="pt", truncation=True).input_ids.to(device)
text_embeddings = self.text_encoder(input_ids=text_ids, output_hidden_states=True).hidden_states[-1]
text_embeddings = self.adapter(text_embeddings).sample
uncond_input = self.tokenizer([""], padding="max_length", max_length=max_length, return_tensors="pt")
uncond_embeddings = self.text_encoder(uncond_input.input_ids.to(device), output_hidden_states=True).hidden_states[-1]
uncond_embeddings = self.adapter(uncond_embeddings).sample
text_embeddings = torch.cat([uncond_embeddings, text_embeddings])
self.adapter.to(offload_device)
self.text_encoder.to(offload_device)
return (text_embeddings,)
class lavi_bridge_t5_encoder:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"prompt": ("STRING", {"multiline": True, "default": "Oppenheimer sits on the beach on a chair, watching a nuclear exposition with a huge mushroom cloud, 120mm",}),
"max_length": ("INT", {"default": 77, "min": 1, "max": 512, "step": 1}),
},
}
RETURN_TYPES = ("LAVIEMBEDS",)
RETURN_NAMES = ("lavi_embeds",)
FUNCTION = "process"
CATEGORY = "LaVI-BridgeWrapper"
@@ -152,7 +226,7 @@ class lavi_bridge_t5_encoder:
if not os.path.exists(adapter_path):
print(f"Downloading LaVi-Bridge from https://huggingface.co/shihaozhao/LaVi-Bridge {adapter_folder}")
from huggingface_hub import snapshot_download
snapshot_download(repo_id="shihaozhao/LaVi-Bridge", allow_patterns=["t5_unet"],local_dir=adapter_folder, local_dir_use_symlinks=False)
snapshot_download(repo_id="shihaozhao/LaVi-Bridge", allow_patterns=["*t5_unet*"],local_dir=adapter_folder, local_dir_use_symlinks=False)
lora_text_path = os.path.join(adapter_folder, 't5_unet', 'lora_text.pt')
@@ -190,7 +264,7 @@ class lavibridge_sampler:
def INPUT_TYPES(s):
return {"required": {
"lavibridge_model": ("LAVIBRIDGE",),
"t5_embeds": ("T5EMBEDS",),
"lavi_embeds": ("LAVIEMBEDS",),
"width": ("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}),
@@ -220,7 +294,7 @@ class lavibridge_sampler:
FUNCTION = "process"
CATEGORY = "LaVI-BridgeWrapper"
def process(self, lavibridge_model, t5_embeds, width, height, batch_size, steps, guidance_scale, seed, scheduler):
def process(self, lavibridge_model, lavi_embeds, width, height, batch_size, steps, guidance_scale, seed, scheduler):
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
mm.unload_all_models()
@@ -273,14 +347,14 @@ class lavibridge_sampler:
latents = latents * noise_scheduler.init_noise_sigma
vae.to(offload_device)
t5_embeds_repeated = t5_embeds.repeat_interleave(batch_size, dim=0)
lavi_embeds_repeated = lavi_embeds.repeat_interleave(batch_size, dim=0)
# Model prediction
noise_scheduler.set_timesteps(steps)
for t in tqdm(noise_scheduler.timesteps):
latent_model_input = torch.cat([latents] * 2, dim=0)
latent_model_input = noise_scheduler.scale_model_input(latent_model_input, timestep=t)
noise_pred = unet(latent_model_input, t, encoder_hidden_states=t5_embeds_repeated).sample
noise_pred = unet(latent_model_input, t, encoder_hidden_states=lavi_embeds_repeated).sample
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2, dim=0)
noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_text - noise_pred_uncond)
latents = noise_scheduler.step(noise_pred, t, latents).prev_sample
@@ -300,11 +374,13 @@ class lavibridge_sampler:
NODE_CLASS_MAPPINGS = {
"lavibridge_sampler": lavibridge_sampler,
"lavi_bridge_t5_encoder": lavi_bridge_t5_encoder,
"lavibridge_model_loader": lavibridge_model_loader
"lavibridge_model_loader": lavibridge_model_loader,
"lavi_bridge_llama_encoder": lavi_bridge_llama_encoder
}
NODE_DISPLAY_NAME_MAPPINGS = {
"lavibridge_sampler": "LaVi-Bridge Sampler",
"lavi_bridge_t5_encoder": "LaVi-Bridge T5 Encoder",
"lavibridge_model_loader": "LaVi-Bridge Model Loader"
"lavibridge_model_loader": "LaVi-Bridge Model Loader",
"lavi_bridge_llama_encoder": "LaVi-Bridge LLaMA Encoder"
}