| import os |
| import argparse |
| from tqdm.auto import tqdm |
| from PIL import Image |
| import torch |
| import re |
| |
| |
| |
| from box import Box |
| import pandas as pd |
| |
| from scipy.spatial.distance import cosine |
|
|
|
|
| def sanitize_filename(name): |
| """Remove or replace characters that are invalid in filenames.""" |
| |
| name = re.sub(r'[<>:"/\\|?*]', '_', name) |
| |
| name = name.strip('. ') |
| |
| 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 [] |
| 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] |
| 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": |
| |
| |
| for specific_token_idx_to_keep_per_prompt in specific_token_idx_to_keep_per_prompt_lists: |
| |
| 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: |
| |
| positions_str = "-".join(str(i) for i in valid_indices) |
| token_str = "_".join(t for t in tokens_to_keep if t) |
| |
| token_str = sanitize_filename(token_str) if token_str else 'tokens' |
| range_name = f"st_pos{positions_str}_{token_str}" |
| |
| complement = [i for i in range(len(tokens)) if i not in valid_indices] |
| |
| 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): |
| |
| 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) |
| |
| 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")): |
| |
| skip_tokens_name+='_2' |
| |
| 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): |
| |
| 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: |
| |
| 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}") |
| |
| 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 |
| """ |
| |
| 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] |
| |
| |
| entity_tokens = tokenizer(entity, return_tensors="pt")['input_ids'][0] |
| entity_token_ids = entity_tokens.tolist()[1:] |
| entity_token_texts = [tokenizer.decode([t]) for t in entity_token_ids] |
| |
| |
| matches = [] |
| |
| |
| for i in range(len(full_token_ids) - len(entity_token_ids) + 1): |
| |
| 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] |
| |
| |
| 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 |
| |
| |
| 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": |
| |
| primary_tok = self.pipe.tokenizer_2 |
| secondary_tok = self.pipe.tokenizer |
| primary_max = self.max_sequence_length |
| secondary_max = clip_max |
| else: |
| |
| 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()), []) |
| |
| 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( |
| |
| "stabilityai/sdxl-turbo", |
| variant="fp16", |
| torch_dtype=torch.float16, |
| |
| ) |
| 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, |
| |
| |
| |
| |
| num_inference_steps=num_inference_steps, |
| |
| |
| lens_kwargs=lens_kwargs, |
| guidance_scale=0.0, |
| |
| |
| num_images_per_prompt=num_images, |
| ).images |
| elif self.model_name == 'sdxl': |
| |
| num_inference_steps = 20 |
| images = self.pipe( |
| prompt=prompt, |
| |
| |
| num_images_per_prompt=num_images, |
| num_inference_steps=num_inference_steps, |
| lens_kwargs=lens_kwargs, |
| |
| |
| ).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): |
| |
| tokenizers = self.get_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, |
| } |
| |
| pipe_output = self.pipe(prompt, |
| num_images_per_prompt=num_images, |
| generator=self.generator, |
| guidance_scale=7.5, |
| |
| 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, |
| } |
|
|
| 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() |
| |
| 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], |
| 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() |
| |
|
|
| 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) |
| |
| 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, |
| } |
| |
| pipe_output = self.pipe(prompt, |
| num_images_per_prompt=num_images, |
| generator=self.generator, |
| guidance_scale=4.5, |
| |
| num_inference_steps=20, |
| |
| 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, |
| } |
|
|
|
|
|
|
|
|
|
|
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
|
|
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
|
|
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| |
| |
| |
| |
|
|
| |
| |
|
|
| |
| |
| |
| |
| |
| |
|
|
| |
| |
| |
| |
| |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| |
| |
|
|
| |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
|
|
| |
| |
| |
| |
|
|
| |
| |
| |
| |
| |
| |
| |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| |
| |
| |
| |
| |
|
|
| |
| |
|
|
| |
| |
| |
| |
| |
|
|
| |
| |
| |
| |
| |
| |
| |
|
|
| |
| |
| |
| |
| |
| |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| |
|
|
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| |
| |
| |
|
|
| |
| |
| |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| |
| |
| |
| |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| |
| |
| |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
|
|