lens-demo / diffusers_wrapper.py
Guy24's picture
add primary_tokenizer_key to Flux+SDXL forward; SD2 per-token support
d64ba11 verified
Raw
History Blame Contribute Delete
63.2 kB
import os
import argparse
from tqdm.auto import tqdm
from PIL import Image
import torch
import re
# from diffusers import AutoencoderKL, UNet2DConditionModel, UniPCMultistepScheduler
# from transformers import LlamaForCausalLM, LlamaTokenizer
# from huggingface_hub import hf_hub_download
from box import Box
import pandas as pd
# from torchvision.transforms import functional as TF
from scipy.spatial.distance import cosine
def sanitize_filename(name):
"""Remove or replace characters that are invalid in filenames."""
# Replace problematic characters with underscores
name = re.sub(r'[<>:"/\\|?*]', '_', name)
# Remove leading/trailing spaces and dots
name = name.strip('. ')
# Replace multiple underscores with single underscore
name = re.sub(r'_+', '_', name)
return name
class TextToImage:
def __init__(self, model_name, ckpt_dir, num_images=1, device="cuda", seed=None):
self.model_name = model_name
self.ckpt_dir = ckpt_dir
self.device = device
self.num_images = num_images
self.seed = seed
self.load_model_components()
def load_model_components(self):
raise NotImplementedError("This method should be implemented by subclasses")
def forward(self, prompt, num_images=1, ranges_to_keep=None, skip_layers=0, return_grid=False, **kwargs):
raise NotImplementedError("This method should be implemented by subclasses")
def get_complementary_range(self, range_to_keep, max_length):
if range_to_keep == 'full':
return [] # For full range, we don't want to skip any tokens
return list(set(range(max_length)) - set(range_to_keep))
def get_ranges_single_tokenizer(self, prompt, tokenizer, max_length,
ranges_to_keep=None, specific_tokens=None,
specific_token_idx_to_keep_per_prompt_lists=None):
ranges = {}
tokens = tokenizer(prompt, return_tensors="pt")['input_ids'][0]
if ranges_to_keep is None:
ranges_to_keep = ['full']
for range_to_keep in ranges_to_keep:
if range_to_keep == "full":
ranges[range_to_keep] = self.get_complementary_range(range_to_keep, max_length)
elif range_to_keep == "tokens":
for word_idx, word in enumerate(specific_tokens):
word_token_ids = tokenizer(word, return_tensors="pt")['input_ids'][0]
word_token_ids = word_token_ids[1:-1] # Remove special tokens
matched_indices = []
for i in range(len(tokens) - len(word_token_ids) + 1):
if torch.all(tokens[i:i + len(word_token_ids)] == word_token_ids):
matched_indices.append(list(range(i, i + len(word_token_ids))))
if matched_indices:
range_name = f"st_{word_idx}_{word}"
ranges[range_name] = [token_index for token_index in range(max_length) if token_index not in matched_indices]
else:
print(f"Word {word} not found in prompt")
elif range_to_keep == "specific_token_idx_to_keep_per_prompt":
# Keep only specific token indices (by index in THIS prompt's tokenization)
# The provided indices MUST correspond to this prompt, not another prompt
for specific_token_idx_to_keep_per_prompt in specific_token_idx_to_keep_per_prompt_lists:
# Build a readable name from the actual tokens at those indices
tokens_to_keep = []
valid_indices = []
for token_idx in specific_token_idx_to_keep_per_prompt:
if token_idx >= len(tokens):
print(f"Warning: Token index {token_idx} is out of bounds for tokens of length {len(tokens)}")
continue
token_to_keep = tokens[token_idx].item()
decoded_token = tokenizer.decode(token_to_keep)
tokens_to_keep.append(decoded_token.strip())
valid_indices.append(token_idx)
if valid_indices:
# Name includes positions to avoid confusion
positions_str = "-".join(str(i) for i in valid_indices)
token_str = "_".join(t for t in tokens_to_keep if t)
# Sanitize token_str to avoid invalid filename characters
token_str = sanitize_filename(token_str) if token_str else 'tokens'
range_name = f"st_pos{positions_str}_{token_str}"
# Compute complement based on actual tokenized length (not max_length padding)
complement = [i for i in range(len(tokens)) if i not in valid_indices]
# For consistency with rest of pipeline, expand to max_length with padding if needed
if len(complement) < max_length:
complement = complement + list(range(len(tokens), max_length))
ranges[range_name] = complement
print(f"{range_name}: {ranges[range_name]}")
else:
print(f"Range {range_to_keep} not recognized")
return ranges
def get_ranges_all_tokenizers(self, prompt, tokenizers, max_lengths, ranges_to_keep=None,
specific_tokens=None):
ranges = None
for (tokenizer_key, tokenizer), max_len in zip(tokenizers.items(), max_lengths):
print(f"Getting ranges for {tokenizer_key}")
updated_ranges = self.get_ranges_single_tokenizer(prompt, tokenizer, max_len, ranges_to_keep, specific_tokens)
if ranges is None:
ranges = {key: [value] for key, value in updated_ranges.items()}
else:
for key, value in updated_ranges.items():
if key in ranges:
ranges[key].append(value)
else:
ranges[key] = [value]
bad_keys = []
for key in ranges:
if len(ranges[key]) < len(tokenizers):
bad_keys.append(key)
for key in bad_keys:
print(f"Removing key {key} from ranges")
ranges.pop(key)
return ranges
def create_image_grid(self, images, output_path=None, number_of_images_per_row=2, do_save=False, do_return=False):
print(f'Creating image grid for {len(images)} images')
widths, heights = zip(*(i.size for i in images))
total_width = sum(widths[:number_of_images_per_row])
total_height = sum(heights[i] for i in range(0, len(heights), number_of_images_per_row))
new_img = Image.new('RGB', (total_width, total_height))
x_offset = 0
y_offset = 0
for i, img in enumerate(images):
new_img.paste(img, (x_offset, y_offset))
x_offset += img.width
if (i + 1) % number_of_images_per_row == 0:
x_offset = 0
y_offset += img.height
if do_save:
if not os.path.exists(os.path.dirname(output_path)):
os.makedirs(os.path.dirname(output_path))
new_img.save(output_path)
print(f'Image grid saved to {output_path}')
if do_return:
return new_img
def save_images(self, images, output_path, skip_tokens_name, save_grid, save_per_image, return_grids,
skip_layers, composite_token=None, sufficient_token=None):
# Sanitize skip_tokens_name to ensure valid filename
skip_tokens_name = sanitize_filename(skip_tokens_name)
unique_token_case = composite_token if composite_token else skip_tokens_name
if composite_token:
suffix= f"st_{sanitize_filename(unique_token_case)}"
prompt_path = os.path.join(output_path, suffix)
else:
prompt_path = os.path.join(output_path, skip_tokens_name)
# Ensure directory exists (makedirs creates all parent directories)
os.makedirs(prompt_path, exist_ok=True)
if save_grid and os.path.exists(os.path.join(prompt_path, f"{skip_tokens_name}_{skip_layers}.png")):
# print(f"Image grid for {skip_tokens_name} already exists, skipping")
skip_tokens_name+='_2'
# return
if save_per_image and os.path.exists(os.path.join(prompt_path, f"{skip_tokens_name}_{skip_layers}_0.png")):
print(f"Image for {skip_tokens_name} already exists, skipping")
return
if save_per_image:
for i, image in enumerate(images):
# Sanitize skip_layers representation for filename
skip_layers_str = str(skip_layers).replace(' ', '')
if unique_token_case:
filename_base = f"composite_{unique_token_case}_{skip_layers_str}_{i}"
filename = sanitize_filename(filename_base) + ".png"
curr_image_path = os.path.join(prompt_path, filename)
else:
filename_base = f"{skip_tokens_name}_{skip_layers_str}_{i}"
filename = sanitize_filename(filename_base) + ".png"
curr_image_path = os.path.join(prompt_path, filename)
image.save(curr_image_path)
print(f"Saved images to path: {prompt_path}")
if save_grid or return_grids:
# Sanitize skip_layers representation for filename
skip_layers_str = str(skip_layers).replace(' ', '')
filename_base = f"{skip_tokens_name}_{skip_layers_str}"
grid_output_path = os.path.join(prompt_path, sanitize_filename(filename_base) + ".png")
grid = self.create_image_grid(images=images, output_path=grid_output_path,
number_of_images_per_row=(self.num_images + 1) // 2, do_save=save_grid, do_return=return_grids)
if return_grids:
return grid
print(f"Saved the image grid to path: {grid_output_path}")
def validate_skip_layers(self, skip_layers, num_tokenizers):
if isinstance(skip_layers, int):
skip_layers = [skip_layers] * num_tokenizers
elif isinstance(skip_layers, list) and len(skip_layers) == 1:
skip_layers = skip_layers * num_tokenizers
elif not (isinstance(skip_layers, list) and len(skip_layers) == num_tokenizers):
raise ValueError(f"skip_layers must be an int or a list of length {num_tokenizers}, got {skip_layers}")
# print(f"Validated skip_layers: {skip_layers}")
return skip_layers
def get_tokenizers(self):
return {
'tokenizer': self.pipe.tokenizer,
}
def get_token_indices_for_entity(self, prompt, entity, tokenizer):
"""Get the exact token indices for an entity within a prompt.
Args:
prompt (str): The full prompt text
entity (str): The entity to find within the prompt
tokenizer: The tokenizer to use
Returns:
tuple: (indices, token_info) where:
- indices is a list of token indices in the prompt that match the entity
- token_info is a dict containing token details for debugging
"""
# Get full prompt tokens
full_tokens = tokenizer(prompt, return_tensors="pt")['input_ids'][0]
full_token_ids = full_tokens.tolist()
full_token_texts = [tokenizer.decode([t]) for t in full_token_ids]
# Get entity tokens
entity_tokens = tokenizer(entity, return_tensors="pt")['input_ids'][0]
entity_token_ids = entity_tokens.tolist()[1:] # Remove BOS token
entity_token_texts = [tokenizer.decode([t]) for t in entity_token_ids]
# Find all occurrences of the entity in the prompt
matches = []
# Match entity tokens in the prompt, handling leading spaces
for i in range(len(full_token_ids) - len(entity_token_ids) + 1):
# Get the tokens we're comparing
prompt_tokens = full_token_ids[i:i + len(entity_token_ids)]
prompt_token_texts = [tokenizer.decode([t]).strip() for t in prompt_tokens]
entity_token_texts_stripped = [t.strip() for t in entity_token_texts]
# Compare the stripped tokens
if prompt_token_texts == entity_token_texts_stripped:
matches.append(list(range(i, i + len(entity_token_ids))))
if not matches:
print(f"Warning: Could not find exact match for entity '{entity}' in prompt")
print(f"Entity tokens: {entity_token_texts}")
print(f"Full prompt tokens: {full_token_texts}")
return None, None
# Return the first match and token information
token_info = {
'entity_tokens': entity_token_texts,
'matched_tokens': [full_token_texts[i] for i in matches[0]],
'full_prompt_tokens': full_token_texts
}
return matches[0], token_info
class StableDiffusion3TextToImage(TextToImage):
def __init__(self, model_name, ckpt_dir, num_images, device="cuda", seed=42, max_sequence_length=256):
super().__init__(model_name, ckpt_dir, num_images, device, seed)
self.max_sequence_length = max_sequence_length
def load_model_components(self):
from diffusers import StableDiffusion3Pipeline
generator = torch.manual_seed(self.seed)
self.generator = generator
pipe = StableDiffusion3Pipeline.from_pretrained("stabilityai/stable-diffusion-3-medium-diffusers", torch_dtype=torch.float16)
self.pipe = pipe.to(self.device)
def forward(self, prompt, num_images, output_path, save_grid=False, save_per_image=True, return_grids=False, skip_layers=0,
ranges_to_keep=None, specific_tokens=None, pad_encoders=[], specific_token_idx_to_keep_per_prompt=None):
tokenizers = {
'tokenizer': self.pipe.tokenizer,
'tokenizer_2': self.pipe.tokenizer_2,
'tokenizer_3': self.pipe.tokenizer_3
}
skip_layers = self.validate_skip_layers(skip_layers, len(tokenizers))
if specific_token_idx_to_keep_per_prompt:
ranges_to_try = self.get_ranges_single_tokenizer(prompt=prompt, tokenizer=self.pipe.tokenizer_3, max_length=self.max_sequence_length,
ranges_to_keep=ranges_to_keep, specific_tokens=None,
specific_token_idx_to_keep_per_prompt=specific_token_idx_to_keep_per_prompt)
for key, value in ranges_to_try.items():
ranges_to_try[key] = [value, value]
else:
ranges_to_try = self.get_ranges_all_tokenizers(prompt=prompt, tokenizers=tokenizers, max_lengths=[self.max_sequence_length] * 3, ranges_to_keep=ranges_to_keep, specific_tokens=specific_tokens)
grids = []
for skip_tokens_name, skip_tokens in ranges_to_try.items():
lens_kwargs = {
'clip_skip': skip_layers,
'skip_tokens': skip_tokens,
'pad_encoders': pad_encoders,
}
pipe_output = self.pipe(prompt, num_images_per_prompt=num_images,
generator=self.generator, num_inference_steps=50,
lens_kwargs=lens_kwargs)
images = pipe_output.images
grid = self.save_images(images, output_path, skip_tokens_name, save_grid, save_per_image, return_grids, skip_layers)
if return_grids and grid is not None:
grids.append(grid)
if return_grids:
return grids
def get_tokenizers(self):
return {
'tokenizer': self.pipe.tokenizer,
'tokenizer_2': self.pipe.tokenizer_2,
'tokenizer_3': self.pipe.tokenizer_3
}
class FluxTextToImage(TextToImage):
def __init__(self, model_name, ckpt_dir, num_images, device="cuda", max_sequence_length=512, seed=42, mask_diffusion=False):
super().__init__(model_name, ckpt_dir, num_images, device)
self.max_sequence_length = max_sequence_length
self.seed = seed
self.mask_diffusion = mask_diffusion
def load_model_components(self):
from diffusers import FluxPipeline
if self.model_name == "flux-schnell":
pipe = FluxPipeline.from_pretrained("black-forest-labs/FLUX.1-schnell", torch_dtype=torch.bfloat16)
self.num_inference_steps=4
elif self.model_name == "flux-dev":
pipe = FluxPipeline.from_pretrained("black-forest-labs/FLUX.1-dev", torch_dtype=torch.bfloat16)
self.num_inference_steps=50
else:
raise ValueError(f"Model name {self.model_name} not recognized")
self.pipe = pipe.to(self.device)
def forward(self, prompt, num_images, output_path, save_grid=False, save_per_image=True, skip_layers=0, ranges_to_keep=None,
specific_tokens=None, return_grids=False, specific_token_idx_to_keep_per_prompt=None, composite_token=None,
mix_tokens=False, clean_prompt=False, clean_tokens_ids=None, sufficient_token=None, merge_ranges=None,
specific_token_idx_to_keep_per_prompt_lists=None, encode_separate_indices=None, clean_tokens=None,
primary_tokenizer_key="tokenizer_2"):
tokenizers = {
'tokenizer': self.pipe.tokenizer,
'tokenizer_2': self.pipe.tokenizer_2,
}
skip_layers = self.validate_skip_layers(skip_layers, len(tokenizers))
clip_max = getattr(self.pipe.tokenizer, "model_max_length", 77)
if specific_token_idx_to_keep_per_prompt_lists:
if primary_tokenizer_key == "tokenizer_2":
# T5-primary: compute T5 complement, align CLIP by surface text
primary_tok = self.pipe.tokenizer_2
secondary_tok = self.pipe.tokenizer
primary_max = self.max_sequence_length
secondary_max = clip_max
else:
# CLIP-primary: compute CLIP complement, align T5 by surface text
primary_tok = self.pipe.tokenizer
secondary_tok = self.pipe.tokenizer_2
primary_max = clip_max
secondary_max = self.max_sequence_length
ranges_primary = self.get_ranges_single_tokenizer(
prompt=prompt, tokenizer=primary_tok,
max_length=primary_max,
ranges_to_keep=ranges_to_keep, specific_tokens=None,
specific_token_idx_to_keep_per_prompt_lists=specific_token_idx_to_keep_per_prompt_lists)
primary_tokens = primary_tok(prompt, return_tensors="pt")["input_ids"][0]
ranges_to_try = {}
for key, primary_complement in ranges_primary.items():
primary_complement_set = set(primary_complement)
kept_indices = [i for i in range(len(primary_tokens)) if i not in primary_complement_set]
kept_text = primary_tok.decode(
[primary_tokens[i].item() for i in kept_indices]).strip()
secondary_ranges = self.get_ranges_single_tokenizer(
prompt=prompt, tokenizer=secondary_tok,
max_length=secondary_max,
ranges_to_keep=["tokens"] if kept_text else ["full"],
specific_tokens=[kept_text] if kept_text else None)
secondary_complement = next(iter(secondary_ranges.values()), [])
# always store as [clip_complement, t5_complement]
if primary_tokenizer_key == "tokenizer_2":
ranges_to_try[key] = [secondary_complement, primary_complement]
else:
ranges_to_try[key] = [primary_complement, secondary_complement]
else:
ranges_to_try = self.get_ranges_all_tokenizers(prompt=prompt, tokenizers=tokenizers, max_lengths=[self.max_sequence_length] * 2, ranges_to_keep=ranges_to_keep, specific_tokens=specific_tokens)
grids = []
for skip_tokens_name, skip_tokens in ranges_to_try.items():
print("skip_tokens", skip_tokens)
lens_kwargs = {
'clip_skip': skip_layers,
'skip_tokens': skip_tokens,
'merge_ranges': merge_ranges,
'encode_separate_indices': encode_separate_indices,
'clean_tokens': clean_tokens
}
with torch.no_grad():
images = self.pipe(
prompt=prompt, #
guidance_scale=0.,
height=512,
width=512,
max_sequence_length=self.max_sequence_length,
generator=torch.Generator("cpu").manual_seed(self.seed),
lens_kwargs=lens_kwargs,
num_inference_steps=self.num_inference_steps,
num_images_per_prompt=num_images,
).images
grid = self.save_images(images, output_path, skip_tokens_name, save_grid, save_per_image,
return_grids=return_grids, skip_layers=skip_layers)
if return_grids and grid is not None:
grids.append(grid)
if return_grids:
return grids
def get_tokenizers(self):
return {
'tokenizer': self.pipe.tokenizer,
'tokenizer_2': self.pipe.tokenizer_2,
}
class StableDiffusion2TextToImage(TextToImage):
def __init__(self, model_name, ckpt_dir, num_images, device="cuda", seed=42):
super().__init__(model_name, ckpt_dir, num_images, device, seed)
def load_model_components(self):
from diffusers import StableDiffusionPipeline
generator = torch.manual_seed(self.seed)
pipe = StableDiffusionPipeline.from_pretrained(
'stabilityai/stable-diffusion-2-1',
torch_dtype=torch.float16,
)
pipe.to(self.device)
self.pipe = pipe
self.generator = generator
def forward(self, prompt, num_images, output_path, save_grid=False, save_per_image=True,
skip_layers=0, ranges_to_keep=None, specific_tokens=None, return_grids=False,
specific_token_idx_to_keep_per_prompt_lists=None, **kwargs):
if specific_token_idx_to_keep_per_prompt_lists:
ranges_to_try = self.get_ranges_single_tokenizer(
prompt=prompt, tokenizer=self.pipe.tokenizer, max_length=77,
ranges_to_keep=ranges_to_keep, specific_tokens=None,
specific_token_idx_to_keep_per_prompt_lists=specific_token_idx_to_keep_per_prompt_lists)
else:
ranges_to_try = self.get_ranges_single_tokenizer(
prompt=prompt, tokenizer=self.pipe.tokenizer, max_length=77,
ranges_to_keep=ranges_to_keep, specific_tokens=specific_tokens)
skip_layers = self.validate_skip_layers(skip_layers, 1)
grids = []
for skip_tokens_name, skip_tokens in ranges_to_try.items():
pipe_output = self.pipe(prompt, num_images_per_prompt=num_images, generator=self.generator, skip_tokens=skip_tokens,
num_inference_steps=20, clip_skip=skip_layers)
images = pipe_output.images
grid = self.save_images(images, output_path, skip_tokens_name, save_grid, save_per_image, return_grids, skip_layers)
if return_grids and grid is not None:
grids.append(grid)
if return_grids:
return grids
class StableDiffusionXLPipelineTextToImage(TextToImage):
def __init__(self, model_name, ckpt_dir, num_images, device="cuda", seed=42, max_sequence_length=77,
pad_encoders=[]):
super().__init__(model_name, ckpt_dir, num_images, device, seed)
self.max_sequence_length = max_sequence_length
self.pad_encoders = pad_encoders
self.pipe = None
self.refiner = None
self.load_model_components()
def load_model_components(self):
from diffusers import AutoPipelineForText2Image
from diffusers import DiffusionPipeline
self.generator = torch.manual_seed(self.seed)
if self.model_name == 'sdxl-turbo':
pipe = AutoPipelineForText2Image.from_pretrained(
# both models work with this code
"stabilityai/sdxl-turbo", # "stabilityai/sdxl-turbo", #"stabilityai/stable-diffusion-xl-base-1.0",
variant="fp16",
torch_dtype=torch.float16,
# generator=self.generator,
)
self.pipe = pipe.to(self.device)
elif self.model_name == 'sdxl':
pipe = DiffusionPipeline.from_pretrained(
"stabilityai/stable-diffusion-xl-base-1.0",
torch_dtype=torch.float16, variant="fp16",
use_safetensors=True
)
refiner = DiffusionPipeline.from_pretrained(
"stabilityai/stable-diffusion-xl-refiner-1.0",
text_encoder_2=pipe.text_encoder_2,
vae=pipe.vae,
torch_dtype=torch.float16,
use_safetensors=True,
variant="fp16",
)
self.pipe = pipe.to(self.device)
self.refiner = refiner.to(self.device)
def forward(self, prompt, num_images, output_path, save_grid=False, save_per_image=True, skip_layers=0,
ranges_to_keep=None,
specific_tokens=None, return_grids=False, specific_token_idx_to_keep_per_prompt_lists=None,
merge_ranges=None, primary_tokenizer_key="tokenizer_2"):
tokenizers = self.get_tokenizers()
skip_layers = self.validate_skip_layers(skip_layers, len(tokenizers))
clip_max = getattr(self.pipe.tokenizer, "model_max_length", 77)
tok2_max = getattr(self.pipe.tokenizer_2, "model_max_length", 77)
if specific_token_idx_to_keep_per_prompt_lists:
if primary_tokenizer_key == "tokenizer_2":
primary_tok, secondary_tok = self.pipe.tokenizer_2, self.pipe.tokenizer
primary_max, secondary_max = tok2_max, clip_max
else:
primary_tok, secondary_tok = self.pipe.tokenizer, self.pipe.tokenizer_2
primary_max, secondary_max = clip_max, tok2_max
ranges_primary = self.get_ranges_single_tokenizer(
prompt=prompt, tokenizer=primary_tok,
max_length=primary_max,
ranges_to_keep=ranges_to_keep, specific_tokens=None,
specific_token_idx_to_keep_per_prompt_lists=specific_token_idx_to_keep_per_prompt_lists)
primary_tokens = primary_tok(prompt, return_tensors="pt")["input_ids"][0]
ranges_to_try = {}
for key, primary_complement in ranges_primary.items():
primary_complement_set = set(primary_complement)
kept_indices = [i for i in range(len(primary_tokens)) if i not in primary_complement_set]
kept_text = primary_tok.decode(
[primary_tokens[i].item() for i in kept_indices]).strip()
secondary_ranges = self.get_ranges_single_tokenizer(
prompt=prompt, tokenizer=secondary_tok,
max_length=secondary_max,
ranges_to_keep=["tokens"] if kept_text else ["full"],
specific_tokens=[kept_text] if kept_text else None)
secondary_complement = next(iter(secondary_ranges.values()), [])
if primary_tokenizer_key == "tokenizer_2":
ranges_to_try[key] = [secondary_complement, primary_complement]
else:
ranges_to_try[key] = [primary_complement, secondary_complement]
else:
ranges_to_try = self.get_ranges_all_tokenizers(prompt=prompt, tokenizers=tokenizers,
max_lengths=[clip_max, tok2_max],
ranges_to_keep=ranges_to_keep,
specific_tokens=specific_tokens)
grids = []
for skip_tokens_name, skip_tokens in ranges_to_try.items():
lens_kwargs = {
'clip_skip': skip_layers,
'skip_tokens': skip_tokens,
'merge_ranges': merge_ranges,
}
if self.model_name == 'sdxl-turbo':
num_inference_steps = 1
images = self.pipe(
prompt=prompt,
# guidance_scale=0.,
# height=512,
# width=512,
# max_sequence_length=self.max_sequence_length,
num_inference_steps=num_inference_steps,
# generator=torch.Generator("cpu").manual_seed(self.seed),
# skip_tokens=skip_tokens,
lens_kwargs=lens_kwargs,
guidance_scale=0.0,
# timesteps=self.num_inference_steps,
# clip_skip=skip_layers,
num_images_per_prompt=num_images,
).images
elif self.model_name == 'sdxl':
# high_noise_frac = 0.8
num_inference_steps = 20
images = self.pipe(
prompt=prompt,
# generator=self.generator,
# lens_kwargs=lens_kwargs,
num_images_per_prompt=num_images,
num_inference_steps=num_inference_steps,
lens_kwargs=lens_kwargs,
# denoising_end=high_noise_frac,
# output_type="latent",
).images
# images = self.refiner(
# prompt=prompt,
# num_inference_steps=num_inference_steps,
# denoising_start=high_noise_frac,
# image=images,
# ).images
grid = self.save_images(images, output_path, skip_tokens_name, save_grid, save_per_image,
return_grids=return_grids, skip_layers=skip_layers)
if return_grids and grid is not None:
grids.append(grid)
if return_grids:
return grids
def get_tokenizers(self):
return {
'tokenizer': self.pipe.tokenizer,
'tokenizer_2': self.pipe.tokenizer_2,
}
class StableDiffusionTextToImage(TextToImage):
def __init__(self, model_name, ckpt_dir, num_images, device="cuda", seed=42, max_sequence_length=256):
super().__init__(model_name, ckpt_dir, num_images, device, seed)
self.max_sequence_length = max_sequence_length
def load_model_components(self):
from diffusers import StableDiffusionPipeline
generator = torch.manual_seed(self.seed)
self.generator = generator
pipe = StableDiffusionPipeline.from_pretrained("CompVis/stable-diffusion-v1-4", torch_dtype=torch.float16)
self.pipe = pipe.to(self.device)
def forward(self, prompt, num_images, output_path, save_grid=False, save_per_image=True, skip_layers=0, ranges_to_keep=None,
specific_tokens=None, return_grids=False, specific_token_idx_to_keep_per_prompt_lists=None):
# token_indices_aae = []
tokenizers = self.get_tokenizers()
# skip_layers = self.validate_skip_layers(skip_layers, len(tokenizers))
if specific_token_idx_to_keep_per_prompt_lists:
ranges_to_try = self.get_ranges_single_tokenizer(prompt=prompt, tokenizer=tokenizers['tokenizer'], max_length=self.max_sequence_length,
ranges_to_keep=ranges_to_keep, specific_tokens=None,
specific_token_idx_to_keep_per_prompt_lists=specific_token_idx_to_keep_per_prompt_lists)
for key, value in ranges_to_try.items():
ranges_to_try[key] = [value, value]
grids = []
for skip_tokens_name, skip_tokens in ranges_to_try.items():
lens_kwargs = {
'clip_skip': skip_layers,
'skip_tokens': skip_tokens,
} #
# print(self.pipe.get_indices(prompt))
pipe_output = self.pipe(prompt,
num_images_per_prompt=num_images,
generator=self.generator,
guidance_scale=7.5,
# max_iter_to_alter=25,
num_inference_steps=50,
# token_indices=[3, 8], # [2,3,6,7,8],
lens_kwargs=lens_kwargs
)
images = pipe_output.images
grid = self.save_images(images, output_path, skip_tokens_name, save_grid, save_per_image, return_grids, skip_layers)
if return_grids and grid is not None:
grids.append(grid)
if return_grids:
return grids
def get_tokenizers(self):
return {
'tokenizer': self.pipe.tokenizer,
}
class StableDiffusionAttendAndExciteTextToImage(TextToImage):
def __init__(self, model_name, ckpt_dir, num_images, device="cuda", seed=42, max_sequence_length=256):
super().__init__(model_name, ckpt_dir, num_images, device, seed)
self.max_sequence_length = max_sequence_length
def load_model_components(self):
from diffusers import StableDiffusionAttendAndExcitePipeline
generator = torch.manual_seed(self.seed)
self.generator = generator
pipe = StableDiffusionAttendAndExcitePipeline.from_pretrained("CompVis/stable-diffusion-v1-4", torch_dtype=torch.float16)
self.pipe = pipe.to(self.device)
def forward(self, prompt, num_images, output_path, save_grid=False, save_per_image=True, skip_layers=0, ranges_to_keep=None,
specific_tokens=None, return_grids=False, specific_token_idx_to_keep_per_prompt_lists=None):
token_indices_aae = []
tokenizers = self.get_tokenizers()
# skip_layers = self.validate_skip_layers(skip_layers, len(tokenizers))
if specific_token_idx_to_keep_per_prompt_lists:
ranges_to_try = self.get_ranges_single_tokenizer(prompt=prompt, tokenizer=tokenizers['tokenizer'], max_length=self.max_sequence_length,
ranges_to_keep=ranges_to_keep, specific_tokens=None,
specific_token_idx_to_keep_per_prompt_lists=specific_token_idx_to_keep_per_prompt_lists)
for key, value in ranges_to_try.items():
ranges_to_try[key] = [value, value]
grids = []
for skip_tokens_name, skip_tokens in ranges_to_try.items():
lens_kwargs = {
'clip_skip': skip_layers,
'skip_tokens': skip_tokens,
} #
print(self.pipe.get_indices(prompt))
pipe_output = self.pipe(prompt,
num_images_per_prompt=num_images,
generator=self.generator,
guidance_scale=7.5,
max_iter_to_alter=25,
num_inference_steps=50,
token_indices=[3, 8], # [2,3,6,7,8],
lens_kwargs=lens_kwargs)
images = pipe_output.images
grid = self.save_images(images, output_path, skip_tokens_name, save_grid, save_per_image, return_grids, skip_layers)
if return_grids and grid is not None:
grids.append(grid)
if return_grids:
return grids
def get_tokenizers(self):
return {
'tokenizer': self.pipe.tokenizer,
}
class SanaPipelineTextToImage(TextToImage):
def __init__(self, model_name, ckpt_dir, num_images, device="cuda", seed=42, max_sequence_length=256):
super().__init__(model_name, ckpt_dir, num_images, device, seed)
self.max_sequence_length = max_sequence_length
def load_model_components(self):
from diffusers import SanaPipeline
pipe = SanaPipeline.from_pretrained(
"Efficient-Large-Model/SANA1.5_1.6B_1024px_diffusers",
torch_dtype=torch.bfloat16,
)
pipe.to(self.device)
self.pipe = pipe
self.generator = torch.manual_seed(self.seed)
def forward(self, prompt, num_images, output_path, save_grid=False, save_per_image=True, skip_layers=0, ranges_to_keep=None,
specific_tokens=None, return_grids=False, specific_token_idx_to_keep_per_prompt_lists=None):
token_indices_aae = []
tokenizers = self.get_tokenizers()
# skip_layers = self.validate_skip_layers(skip_layers, len(tokenizers))
if specific_token_idx_to_keep_per_prompt_lists:
ranges_to_try = self.get_ranges_single_tokenizer(prompt=prompt, tokenizer=self.pipe.tokenizer, max_length=self.max_sequence_length,
ranges_to_keep=ranges_to_keep, specific_tokens=None,
specific_token_idx_to_keep_per_prompt_lists=specific_token_idx_to_keep_per_prompt_lists)
for key, value in ranges_to_try.items():
ranges_to_try[key] = [value, value]
else:
ranges_to_try = self.get_ranges_all_tokenizers(prompt=prompt, tokenizers=tokenizers, max_lengths=[self.max_sequence_length] * 2, ranges_to_keep=ranges_to_keep, specific_tokens=specific_tokens)
# pipeline's encode_prompt indexes skip_tokens[1]; duplicate for single-tokenizer model
for key in ranges_to_try:
if len(ranges_to_try[key]) == 1:
ranges_to_try[key] = ranges_to_try[key] * 2
grids = []
for skip_tokens_name, skip_tokens in ranges_to_try.items():
print("skip_tokens", skip_tokens)
lens_kwargs = {
'clip_skip': skip_layers,
'skip_tokens': skip_tokens,
}
# print(self.pipe.get_indices(prompt))
pipe_output = self.pipe(prompt,
num_images_per_prompt=num_images,
generator=self.generator,
guidance_scale=4.5,
# max_iter_to_alter=25,
num_inference_steps=20,
# token_indices=[3, 8], # [2,3,6,7,8],
lens_kwargs=lens_kwargs
)
images = pipe_output.images
grid = self.save_images(images, output_path, skip_tokens_name, save_grid, save_per_image, return_grids, skip_layers)
if return_grids and grid is not None:
grids.append(grid)
if return_grids:
return grids
def get_tokenizers(self):
return {
'tokenizer': self.pipe.tokenizer,
}
# def forward(self, prompt, num_images, output_path, save_grid=False, save_per_image=True, skip_layers=0, ranges_to_keep=None, specific_tokens=None, return_grids=False):
# tokenizers = {
# 'tokenizer': self.pipe.tokenizer,
# 'tokenizer_2': self.pipe.tokenizer_2,
# }
# skip_layers = self.validate_skip_layers(skip_layers, len(tokenizers))
# if specific_token_idx_to_keep_per_prompt_lists:
# ranges_to_try = self.get_ranges_single_tokenizer(prompt=prompt, tokenizer=self.pipe.tokenizer_2, max_length=self.max_sequence_length,
# ranges_to_keep=ranges_to_keep, specific_tokens=None,
# specific_token_idx_to_keep_per_prompt_lists=specific_token_idx_to_keep_per_prompt_lists)
# for key, value in ranges_to_try.items():
# ranges_to_try[key] = [value, value]
# else:
# ranges_to_try = self.get_ranges_all_tokenizers(prompt=prompt, tokenizers=tokenizers, max_lengths=[self.max_sequence_length] * 2, ranges_to_keep=ranges_to_keep, specific_tokens=specific_tokens)
# grids = []
# for skip_tokens_name, skip_tokens in ranges_to_try.items():
# lens_kwargs = {
# 'clip_skip': skip_layers,
# 'skip_tokens': skip_tokens,
# }
# pipe_output = self.pipe(prompt, lens_kwargs=lens_kwargs, num_images_per_prompt=num_images, generator=self.generator, skip_tokens=skip_tokens,
# num_inference_steps=4, clip_skip=skip_layers)
# images = pipe_output.images
# grid = self.save_images(images, output_path, skip_tokens_name, save_grid, save_per_image, return_grids, skip_layers)
# if return_grids and grid is not None:
# grids.append(grid)
# if return_grids:
# return grids
# def get_tokenizers(self):
# return {
# 'tokenizer': self.pipe.tokenizer,
# 'tokenizer_2': self.pipe.tokenizer_2,
# }
# import os
# import sys
# import argparse
# from tqdm.auto import tqdm
# from PIL import Image
# import torch
# from diffusers import AutoencoderKL, UNet2DConditionModel, UniPCMultistepScheduler
# from transformers import LlamaForCausalLM, LlamaTokenizer
# from huggingface_hub import hf_hub_download
# from argparse import ArgumentParser
# from box import Box
# import os
# import pandas as pd
# from PIL import Image
# from torchvision.transforms import functional as TF
# import torch
# from tqdm import tqdm
# from PIL import Image
# import torch
# from PIL import Image
# from scipy.spatial.distance import cosine
# class TextToImage:
# def __init__(self, model_name, ckpt_dir, num_images=1, device="cuda", seed=None):
# self.model_name = model_name
# self.ckpt_dir = ckpt_dir
# self.device = device
# self.num_images = num_images
# self.seed=seed
# self.load_model_components()
# def load_model_components(self):
# raise NotImplementedError("This method should be implemented by subclasses")
# def forward(self, prompt, num_images=1, per_token=False, per_range=False):
# raise NotImplementedError("This method should be implemented by subclasses")
# def get_complementory_range(self, range_to_keep, max_length):
# return list(set(range(max_length)) - set(range_to_keep))
# def get_ranges_single_tokenizer(self, prompt, tokenizer, max_length):
# ranges = {}
# tokens = tokenizer(prompt, return_tensors="pt")['input_ids'][0]
# token_length = len(tokens)
# decoded_tokens = [tokenizer.decode(token.item()) for token in tokens]
# pad_length = max_length - token_length
# ranges['full'] = []
# # ranges['tokens'] = list(range(0, token_length))
# # ranges['pads'] = list(range(token_length, max_length))
# # ranges['eot'] = list(range(0, len(tokens)-1)) + list(range(len(tokens), max_length))
# # ranges['None'] = list(range(0, max_length))
# return ranges
# def get_ranges_all_tokenizers(self, prompt, tokenizers, max_lengths):
# ranges = None
# for (tokenizer_key, tokenizer), max_len in zip(tokenizers.items(), max_lengths):
# print(f"Getting ranges for {tokenizer_key}")
# updated_ranges = self.get_ranges_single_tokenizer(prompt, tokenizer, max_len)
# if ranges is None:
# ranges = {key: [value] for key, value in updated_ranges.items()}
# else:
# for key, value in updated_ranges.items():
# if key in ranges:
# ranges[key].append(value)
# else:
# ranges[key] = [value]
# # Ensure each key has exactly len(tokenizers) items in its list
# bad_keys = []
# for key in ranges:
# if len(ranges[key]) < len(tokenizers):
# bad_keys.append(key)
# # while len(ranges[key]) < len(tokenizers):
# # ranges[key].append([range(0, 77)]) # Append a full mask range if not present
# for key in bad_keys:
# print(f"Removing key {key} from ranges")
# ranges.pop(key)
# return ranges
# def create_image_grid(self, images, output_path=None, number_of_images_per_row=2, do_save=False):
# print(f'Creating image grid for {len(images)} images')
# widths, heights = zip(*(i.size for i in images))
# total_width = sum(widths[:number_of_images_per_row]) # width of 10 images
# total_height = sum(heights[i] for i in range(0, len(heights), number_of_images_per_row)) # height of every 10th image
# new_img = Image.new('RGB', (total_width, total_height))
# x_offset = 0
# y_offset = 0
# for i, img in enumerate(images):
# new_img.paste(img, (x_offset, y_offset))
# x_offset += img.width
# if (i + 1) % number_of_images_per_row == 0: # move to next row after every 10 images
# x_offset = 0
# y_offset += img.height
# if do_save:
# if not os.path.exists(os.path.dirname(output_path)):
# os.makedirs(os.path.dirname(output_path))
# new_img.save(output_path)
# print(f'Image grid saved to {output_path}')
# else:
# return new_img
# class StableDiffusion3TextToImage(TextToImage):
# def __init__(self, model_name, ckpt_dir, num_images, device="cuda", seed=42, max_sequence_length=256):
# super().__init__(model_name, ckpt_dir, num_images, device, seed)
# self.max_sequence_length = max_sequence_length
# def load_model_components(self):
# # from diffusers_local.src.diffusers import StableDiffusion3Pipeline
# from diffusers import StableDiffusion3Pipeline
# generator = torch.manual_seed(self.seed)
# self.generator = generator
# pipe = StableDiffusion3Pipeline.from_pretrained("stabilityai/stable-diffusion-3-medium-diffusers", torch_dtype=torch.float16)
# self.pipe = pipe.to(self.device)
# def forward(self, prompt, num_images, output_path, pad_str=None, save_grid=False,
# save_per_image=True, zero_paddings=False, replace_with_pads=False, turn_attention_off=False, return_grids=False):
# tokenizer = self.pipe.tokenizer
# tokenizer_2 = self.pipe.tokenizer_2
# tokenizer_3 = self.pipe.tokenizer_3
# tokenizers = {
# 'tokenizer': tokenizer,
# 'tokenizer_2': tokenizer_2,
# 'tokenizer_3': tokenizer_3
# }
# number_of_tokens = self.pipe.tokenizer(prompt, return_tensors="pt")['input_ids'][0].shape[0]
# original_max_sequence_length = self.max_sequence_length
# max_sequence_length = self.max_sequence_length
# if max_sequence_length is None:
# max_sequence_length = 256
# elif max_sequence_length == 'prompt_len':
# max_sequence_length = number_of_tokens
# elif max_sequence_length == '2prompt_len':
# max_sequence_length = 2 * number_of_tokens
# ranges_to_try = self.get_ranges_all_tokenizers(prompt=prompt, tokenizers=tokenizers, max_lengths=[max_sequence_length] * 3)
# for skip_tokens_name, skip_tokens in ranges_to_try.items():
# prompt_words = prompt.split()
# prompt_snippet = "_".join(prompt_words[:15] if len(prompt_words) >= 5 else prompt_words)
# prompt_path = os.path.join(output_path, prompt_snippet)
# if not os.path.exists(prompt_path):
# os.makedirs(prompt_path)
# # check if exists before contine
# if save_grid and os.path.exists(os.path.join(prompt_path, f"{skip_tokens_name}.png")):
# print(f"Image grid for {skip_tokens_name} already exists, skipping")
# continue
# if save_per_image and os.path.exists(os.path.join(prompt_path, f"{skip_tokens_name}_0.png")):
# print(f"Image for {skip_tokens_name} already exists, skipping")
# continue
# pipe_output = self.pipe(prompt, num_images_per_prompt=num_images, generator=self.generator, skip_tokens=skip_tokens,
# num_inference_steps=50, pad_encoders=None, replace_with_pads=replace_with_pads,
# turn_attention_off=turn_attention_off, max_sequence_length=max_sequence_length)
# images = pipe_output.images
# if save_grid or return_grids:
# grid_output_path = os.path.join(prompt_path, f"{skip_tokens_name}.png")
# grid = self.create_image_grid(images=images, output_path=grid_output_path,
# number_of_images_per_row=(self.num_images + 1) // 2, do_save=not return_grids)
# if return_grids:
# return grid
# print(f"Saved the image grid to path: {grid_output_path}")
# if save_per_image:
# for i, image in enumerate(images):
# curr_image_path = os.path.join(prompt_path, f"{skip_tokens_name}_{i}.png")
# image.save(curr_image_path)
# print(f"Saved the image to path: {curr_image_path}")
# print(f"Generated image for {skip_tokens_name} to {output_path}")
# class FluxTextToImage(TextToImage):
# def __init__(self, model_name, ckpt_dir, num_images, device="cuda", max_sequence_length=512, seed=42, mask_diffusion=False):
# super().__init__(model_name, ckpt_dir, num_images, device)
# self.max_sequence_length = 512 if max_sequence_length == None else max_sequence_length
# self.seed = seed
# self.mask_diffusion = mask_diffusion
# def load_model_components(self):
# # Implement loading of Flux model components
# from diffusers import FluxPipeline
# if self.model_name == "flux-schnell":
# pipe = FluxPipeline.from_pretrained("black-forest-labs/FLUX.1-schnell", torch_dtype=torch.bfloat16)
# elif self.model_name == "flux-dev":
# pipe = FluxPipeline.from_pretrained("black-forest-labs/FLUX.1-dev", torch_dtype=torch.bfloat16)
# else:
# raise ValueError(f"Model name {self.model_name} not recognized")
# self.pipe = pipe.to(self.device)
# def forward(self, prompt, num_images, output_path, pad_str=None, save_grid=False,
# save_per_image=True, zero_paddings=False, replace_with_pads=False,
# turn_attention_off=False, num_inference_steps=4, clip_skip=None):
# tokenizer = self.pipe.tokenizer
# tokenizer_2 = self.pipe.tokenizer_2
# tokenizers = {
# 'tokenizer': tokenizer,
# 'tokenizer_2': tokenizer_2,
# }
# original_max_sequence_length = self.max_sequence_length
# prompt_ids_len = len(tokenizer_2(prompt)['input_ids'])
# if self.max_sequence_length == 'prompt_len':
# max_sequence_length = prompt_ids_len
# elif self.max_sequence_length == '2prompt_len':
# max_sequence_length = 2 * prompt_ids_len
# else:
# max_sequence_length = self.max_sequence_length
# ranges_to_try = self.get_ranges_all_tokenizers(prompt=prompt, tokenizers=tokenizers, max_lengths=[max_sequence_length] * 2)
# for skip_tokens_name, skip_tokens in ranges_to_try.items():
# prompt_words = prompt.split()
# prompt_snippet = "_".join(prompt_words[:15] if len(prompt_words) >= 5 else prompt_words)
# prompt_path = os.path.join(output_path, prompt_snippet)
# if not os.path.exists(prompt_path):
# os.makedirs(prompt_path)
# prompt_input_ids_len = len(tokenizer_2(prompt)['input_ids'])
# if self.mask_diffusion:
# prompt_copy = prompt
# prompt = [prompt_copy, '', prompt_copy, '']
# num_images = 1
# if self.model_name == "flux-dev":
# images = self.pipe(
# prompt=prompt,
# height=1024,
# width=1024,
# guidance_scale=3.5,
# num_inference_steps=50,
# max_sequence_length=max_sequence_length,
# generator=torch.Generator("cpu").manual_seed(self.seed),
# clip_skip=None,
# skip_tokens=skip_tokens,
# clip_skip=skip_layers,
# num_images_per_prompt=num_images,
# prompt_len=prompt_input_ids_len,
# mask_diffusion=self.mask_diffusion
# ).images
# elif self.model_name == "flux-schnell":
# images = self.pipe(
# prompt=prompt,
# guidance_scale=0.,
# height=512,
# width=512,
# max_sequence_length=max_sequence_length,
# generator=torch.Generator("cpu").manual_seed(self.seed),
# clip_skip=None,
# skip_tokens=skip_tokens,
# num_images_per_prompt=num_images,
# ).images
# else:
# raise ValueError(f"Model name {self.model_name} not recognized")
# os.makedirs(prompt_path, exist_ok=True)
# if save_grid:
# grid_output_path = os.path.join(prompt_path, f"{skip_tokens_name}.png")
# self.create_image_grid(images=images, output_path=grid_output_path,
# number_of_images_per_row=(self.num_images + 1) // 2, do_save=True)
# print(f"Saved the image grid to path: {grid_output_path}")
# if save_per_image:
# for i, image in enumerate(images):
# curr_image_path = os.path.join(prompt_path, f"{skip_tokens_name}_{i}.png")
# image.save(curr_image_path)
# print(f"Saved the image to path: {curr_image_path}")
# print(f"Generated image for {skip_tokens_name} to {output_path}")
# class StableDiffusion2TextToImage(TextToImage):
# def __init__(self, model_name, ckpt_dir, num_images, device="cuda", seed=42,):
# super().__init__(model_name, ckpt_dir, num_images, device, seed)
# def load_model_components(self):
# from diffusers import StableDiffusionPipeline
# generator = torch.manual_seed(self.seed)
# device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
# pipe = StableDiffusionPipeline.from_pretrained(
# 'stabilityai/stable-diffusion-2-1',
# # variant="fp16",
# # repo_type="huggingface",
# torch_dtype=torch.float32,
# # use_safetensors=True,
# generator=generator,
# )
# pipe.to(device)
# self.pipe = pipe
# self.generator = generator
# self.device = device
# def forward(self, prompt, num_images, output_path, pad_str=None, save_grid=False,
# save_per_image=True, zero_paddings=False, replace_with_pads=False, turn_attention_off=False):
# tokenizer = self.pipe.tokenizer
# ranges_to_try = self.get_ranges_single_tokenizer(prompt=prompt, tokenizer=tokenizer, max_length=77)
# for skip_tokens_name, skip_tokens in ranges_to_try.items():
# prompt_words = prompt.split()
# prompt_snippet = "_".join(prompt_words[:15] if len(prompt_words) >= 5 else prompt_words)
# prompt_path = os.path.join(output_path, prompt_snippet)
# if not os.path.exists(prompt_path):
# os.makedirs(prompt_path)
# # check if exists before contine
# if save_grid and os.path.exists(os.path.join(prompt_path, f"{skip_tokens_name}.png")):
# print(f"Image grid for {skip_tokens_name} already exists, skipping")
# continue
# if save_per_image and os.path.exists(os.path.join(prompt_path, f"{skip_tokens_name}_0.png")):
# print(f"Image for {skip_tokens_name} already exists, skipping")
# continue
# pipe_output = self.pipe(prompt, num_images_per_prompt=num_images, generator=self.generator, skip_tokens=skip_tokens,
# num_inference_steps=20) #, pad_encoders=None, replace_with_pads=replace_with_pads, turn_attention_off=turn_attention_off)
# images = pipe_output.images
# if save_grid:
# grid_output_path = os.path.join(prompt_path, f"{skip_tokens_name}.png")
# self.create_image_grid(images=images, output_path=grid_output_path,
# number_of_images_per_row=(self.num_images + 1) // 2, do_save=True)
# print(f"Saved the image grid to path: {grid_output_path}")
# if save_per_image:
# for i, image in enumerate(images):
# curr_image_path = os.path.join(prompt_path, f"{skip_tokens_name}_{i}.png")
# image.save(curr_image_path)
# print(f"Saved the image to path: {curr_image_path}")
# print(f"Generated image for {skip_tokens_name} to {output_path}")
# class StableDiffusionXLPipelineTextToImage(TextToImage):
# def __init__(self, model_name, ckpt_dir, num_images, device="cuda", seed=42, max_sequence_length=77):
# super().__init__(model_name, ckpt_dir, num_images, device, seed)
# self.max_sequence_length = max_sequence_length
# def load_model_components(self):
# from diffusers import StableDiffusionXLPipeline
# generator = torch.manual_seed(self.seed)
# device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
# pipe = StableDiffusionXLPipeline.from_pretrained(
# "stabilityai/stable-diffusion-xl-base-1.0",
# variant="fp16",
# # repo_type="huggingface",
# torch_dtype=torch.float16,
# # use_safetensors=True,
# generator=generator,
# )
# pipe.to(device)
# self.pipe = pipe
# self.generator = generator
# self.device = device
# def forward(self, prompt, num_images, output_path, pad_str=None, save_grid=False,
# save_per_image=True, zero_paddings=False, replace_with_pads=False, turn_attention_off=False):
# tokenizer = self.pipe.tokenizer
# tokenizer_2 = self.pipe.tokenizer_2
# tokenizers = {
# 'tokenizer': tokenizer,
# 'tokenizer_2': tokenizer_2,
# }
# ranges_to_try = self.get_ranges_all_tokenizers(prompt=prompt, tokenizers=tokenizers, max_lengths=[77] * 2)
# for skip_tokens_name, skip_tokens in ranges_to_try.items():
# prompt_words = prompt.split()
# prompt_snippet = "_".join(prompt_words[:15] if len(prompt_words) >= 5 else prompt_words)
# prompt_path = os.path.join(output_path, prompt_snippet)
# if not os.path.exists(prompt_path):
# os.makedirs(prompt_path)
# # check if exists before contine
# if save_grid and os.path.exists(os.path.join(prompt_path, f"{skip_tokens_name}.png")):
# print(f"Image grid for {skip_tokens_name} already exists, skipping")
# continue
# if save_per_image and os.path.exists(os.path.join(prompt_path, f"{skip_tokens_name}_0.png")):
# print(f"Image for {skip_tokens_name} already exists, skipping")
# continue
# max_sequence_length = self.max_sequence_length
# if max_sequence_length == None:
# max_sequence_length = 77
# elif max_sequence_length == 'prompt_len':
# max_sequence_length = len(tokenizer(prompt)['input_ids'])
# elif max_sequence_length == '2prompt_len':
# max_sequence_length = 2 * len(tokenizer(prompt)['input_ids'])
# else:
# print(f"max_sequence_length: {max_sequence_length}")
# with torch.no_grad():
# import time
# num_tokens = len(self.pipe.tokenizer_2(prompt)['input_ids'])
# max_sequence_length = num_tokens
# start_time = time.time()
# pipe_output = self.pipe(prompt, num_images_per_prompt=num_images, generator=self.generator, skip_tokens=skip_tokens,
# num_inference_steps=50, pad_encoders=None, max_sequence_length=max_sequence_length)
# end_time = time.time()
# print(f"Time taken for inference: {end_time - start_time}")
# images = pipe_output.images
# if save_grid:
# grid_output_path = os.path.join(prompt_path, f"{skip_tokens_name}.png")
# self.create_image_grid(images=images, output_path=grid_output_path,
# number_of_images_per_row=(self.num_images + 1) // 2, do_save=True)
# print(f"Saved the image grid to path: {grid_output_path}")
# if save_per_image:
# for i, image in enumerate(images):
# curr_image_path = os.path.join(prompt_path, f"{skip_tokens_name}_{i}.png")
# image.save(curr_image_path)
# print(f"Saved the image to path: {curr_image_path}")
# print(f"Generated image for {skip_tokens_name} to {output_path}")