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}")