Team Ai
Apppublic

michaelcreatesstuff/llm-grounded-diffusion

sourceHugging Faceupdated 3y agoView on Hugging Face
2likes
generation.py229 linesDownload Raw Back to root
1version = "v3.0"2 3import torch4import numpy as np5import models6import utils7from models import pipelines, sam8from utils import parse, latents9from shared import model_dict, sam_model_dict, DEFAULT_SO_NEGATIVE_PROMPT, DEFAULT_OVERALL_NEGATIVE_PROMPT10import gc11from io import BytesIO12import base6413import PIL.Image14 15verbose = False16 17vae, tokenizer, text_encoder, unet, dtype = model_dict.vae, model_dict.tokenizer, model_dict.text_encoder, model_dict.unet, model_dict.dtype18 19model_dict.update(sam_model_dict)20 21 22# Hyperparams23height = 512  # default height of Stable Diffusion24width = 512  # default width of Stable Diffusion25H, W = height // 8, width // 8 # size of the latent26guidance_scale = 7.5  # Scale for classifier-free guidance27 28# batch size that is not 1 is not supported29overall_batch_size = 130 31# discourage masks with confidence below32discourage_mask_below_confidence = 0.8533 34# discourage masks with iou (with coarse binarized attention mask) below35discourage_mask_below_coarse_iou = 0.2536 37run_ind = None38 39 40def generate_single_object_with_box_batch(prompts, bboxes, phrases, words, input_latents_list, input_embeddings, 41                                    sam_refine_kwargs, num_inference_steps, gligen_scheduled_sampling_beta=0.3, 42                                    verbose=False, scheduler_key=None, visualize=True, batch_size=None):43    # batch_size=None: does not limit the batch size (pass all input together)44    45    # prompts and words are not used since we don't have cross-attention control in this function46    47    input_latents = torch.cat(input_latents_list, dim=0)48    49    # We need to "unsqueeze" to tell that we have only one box and phrase in each batch item50    bboxes, phrases = [[item] for item in bboxes], [[item] for item in phrases]51    52    input_len = len(bboxes)53    assert len(bboxes) == len(phrases), f"{len(bboxes)} != {len(phrases)}"54    55    if batch_size is None:56        batch_size = input_len57    58    run_times = int(np.ceil(input_len / batch_size))59    mask_selected_list, single_object_pil_images_box_ann, latents_all = [], [], []60    for batch_idx in range(run_times):61        input_latents_batch, bboxes_batch, phrases_batch = input_latents[batch_idx * batch_size:(batch_idx + 1) * batch_size], \62            bboxes[batch_idx * batch_size:(batch_idx + 1) * batch_size], phrases[batch_idx * batch_size:(batch_idx + 1) * batch_size]63        input_embeddings_batch = input_embeddings[0], input_embeddings[1][batch_idx * batch_size:(batch_idx + 1) * batch_size]64        65        _, single_object_images_batch, single_object_pil_images_box_ann_batch, latents_all_batch = pipelines.generate_gligen(66            model_dict, input_latents_batch, input_embeddings_batch, num_inference_steps, bboxes_batch, phrases_batch, gligen_scheduled_sampling_beta=gligen_scheduled_sampling_beta, 67            guidance_scale=guidance_scale, return_saved_cross_attn=False,68            return_box_vis=True, save_all_latents=True, batched_condition=True, scheduler_key=scheduler_key69        )70        71        gc.collect()72        torch.cuda.empty_cache()73        74        # `sam_refine_boxes` also calls `empty_cache` so we don't need to explicitly empty the cache again.75        mask_selected, _ = sam.sam_refine_boxes(sam_input_images=single_object_images_batch, boxes=bboxes_batch, model_dict=model_dict, verbose=verbose, **sam_refine_kwargs)76        77        mask_selected_list.append(np.array(mask_selected)[:, 0])78        single_object_pil_images_box_ann.append(single_object_pil_images_box_ann_batch)79        latents_all.append(latents_all_batch)80    81    single_object_pil_images_box_ann, latents_all = sum(single_object_pil_images_box_ann, []), torch.cat(latents_all, dim=1)82    83    # mask_selected_list: List(batch)[List(image)[List(box)[Array of shape (64, 64)]]]84    85    mask_selected = np.concatenate(mask_selected_list, axis=0)86    mask_selected = mask_selected.reshape((-1, *mask_selected.shape[-2:]))87    88    assert mask_selected.shape[0] == input_latents.shape[0], f"{mask_selected.shape[0]} != {input_latents.shape[0]}"89    90    print(mask_selected.shape)91    92    mask_selected_tensor = torch.tensor(mask_selected)93    94    latents_all = latents_all.transpose(0,1)[:,:,None,...]95    96    gc.collect()97    torch.cuda.empty_cache()98    99    return latents_all, mask_selected_tensor, single_object_pil_images_box_ann100 101def get_masked_latents_all_list(so_prompt_phrase_word_box_list, input_latents_list, so_input_embeddings, verbose=False, **kwargs):102    latents_all_list, mask_tensor_list = [], []103   104    if not so_prompt_phrase_word_box_list:105        return latents_all_list, mask_tensor_list106    107    prompts, bboxes, phrases, words = [], [], [], []108 109    for prompt, phrase, word, box in so_prompt_phrase_word_box_list:110        prompts.append(prompt)111        bboxes.append(box)112        phrases.append(phrase)113        words.append(word)114    115    latents_all_list, mask_tensor_list, so_img_list = generate_single_object_with_box_batch(prompts, bboxes, phrases, words, input_latents_list, input_embeddings=so_input_embeddings, verbose=verbose, **kwargs)116 117    return latents_all_list, mask_tensor_list, so_img_list118 119 120# Note: need to keep the supervision, especially the box corrdinates, corresponds to each other in single object and overall.121 122def run(123    spec, bg_seed = 1, overall_prompt_override="", fg_seed_start = 20, frozen_step_ratio=0.4, gligen_scheduled_sampling_beta = 0.3, num_inference_steps = 20,124    so_center_box = False, fg_blending_ratio = 0.1, scheduler_key='dpm_scheduler', so_negative_prompt = DEFAULT_SO_NEGATIVE_PROMPT, overall_negative_prompt = DEFAULT_OVERALL_NEGATIVE_PROMPT, so_horizontal_center_only = True, 125    align_with_overall_bboxes = False, horizontal_shift_only = True, use_autocast = False, so_batch_size = None126):127    """    128    so_center_box: using centered box in single object generation129    so_horizontal_center_only: move to the center horizontally only130    131    align_with_overall_bboxes: Align the center of the mask, latents, and cross-attention with the center of the box in overall bboxes132    horizontal_shift_only: only shift horizontally for the alignment of mask, latents, and cross-attention133    """134    135    print("generation:", spec, bg_seed, fg_seed_start, frozen_step_ratio, gligen_scheduled_sampling_beta)136    137    frozen_step_ratio = min(max(frozen_step_ratio, 0.), 1.)138    frozen_steps = int(num_inference_steps * frozen_step_ratio)139 140    if True:141        so_prompt_phrase_word_box_list, overall_prompt, overall_phrases_words_bboxes = parse.convert_spec(spec, height, width, verbose=verbose)142 143    if overall_prompt_override and overall_prompt_override.strip():144        overall_prompt = overall_prompt_override.strip()145 146    overall_phrases, overall_words, overall_bboxes = [item[0] for item in overall_phrases_words_bboxes], [item[1] for item in overall_phrases_words_bboxes], [item[2] for item in overall_phrases_words_bboxes]147 148    # The so box is centered but the overall boxes are not (since we need to place to the right place).149    if so_center_box:150        so_prompt_phrase_word_box_list = [(prompt, phrase, word, utils.get_centered_box(bbox, horizontal_center_only=so_horizontal_center_only)) for prompt, phrase, word, bbox in so_prompt_phrase_word_box_list]151        if verbose:152            print(f"centered so_prompt_phrase_word_box_list: {so_prompt_phrase_word_box_list}")153    so_boxes = [item[-1] for item in so_prompt_phrase_word_box_list]154 155    sam_refine_kwargs = dict(156        discourage_mask_below_confidence=discourage_mask_below_confidence, discourage_mask_below_coarse_iou=discourage_mask_below_coarse_iou,157        height=height, width=width, H=H, W=W158    )159    160    # Note that so and overall use different negative prompts161 162    with torch.autocast("cuda", enabled=use_autocast):163        so_prompts = [item[0] for item in so_prompt_phrase_word_box_list]164        if so_prompts:165            so_input_embeddings = models.encode_prompts(prompts=so_prompts, tokenizer=tokenizer, text_encoder=text_encoder, negative_prompt=so_negative_prompt, one_uncond_input_only=True)166        else:167            so_input_embeddings = []168 169        overall_input_embeddings = models.encode_prompts(prompts=[overall_prompt], tokenizer=tokenizer, negative_prompt=overall_negative_prompt, text_encoder=text_encoder)170        171        input_latents_list, latents_bg = latents.get_input_latents_list(172            model_dict, bg_seed=bg_seed, fg_seed_start=fg_seed_start, 173            so_boxes=so_boxes, fg_blending_ratio=fg_blending_ratio, height=height, width=width, verbose=False174        )175        latents_all_list, mask_tensor_list, so_img_list = get_masked_latents_all_list(176            so_prompt_phrase_word_box_list, input_latents_list, 177            gligen_scheduled_sampling_beta=gligen_scheduled_sampling_beta,178            sam_refine_kwargs=sam_refine_kwargs, so_input_embeddings=so_input_embeddings, num_inference_steps=num_inference_steps, scheduler_key=scheduler_key, verbose=verbose, batch_size=so_batch_size179        )180 181        182 183        composed_latents, foreground_indices, offset_list = latents.compose_latents_with_alignment(184            model_dict, latents_all_list, mask_tensor_list, num_inference_steps, 185            overall_batch_size, height, width, latents_bg=latents_bg, 186            align_with_overall_bboxes=align_with_overall_bboxes, overall_bboxes=overall_bboxes,187            horizontal_shift_only=horizontal_shift_only188        )189        190        overall_bboxes_flattened, overall_phrases_flattened = [], []191        for overall_bboxes_item, overall_phrase in zip(overall_bboxes, overall_phrases):192            for overall_bbox in overall_bboxes_item:193                overall_bboxes_flattened.append(overall_bbox)194                overall_phrases_flattened.append(overall_phrase)195 196        # Generate with composed latents197 198        # Foreground should be frozen199        frozen_mask = foreground_indices != 0200        201        regen_latents, images = pipelines.generate_gligen(202            model_dict, composed_latents, overall_input_embeddings, num_inference_steps, 203            overall_bboxes_flattened, overall_phrases_flattened, guidance_scale=guidance_scale,204            gligen_scheduled_sampling_beta=gligen_scheduled_sampling_beta,205            frozen_steps=frozen_steps, frozen_mask=frozen_mask, scheduler_key=scheduler_key206        )207 208        print(f"Generation with spatial guidance from input latents and first {frozen_steps} steps frozen (directly from the composed latents input)")209        print("Generation from composed latents (with semantic guidance)")210 211        # display(Image.fromarray(images[0]), "img", run_ind)212        213    gc.collect()214    torch.cuda.empty_cache()215 216    # Convert to PIL Image217    image = PIL.Image.fromarray(images[0])218    219    # Save as PNG in memory220    buffer = BytesIO()221    image.save(buffer, format='PNG')222    223    # Encode PNG to base64224    png_bytes = buffer.getvalue()225    base64_string = base64.b64encode(png_bytes).decode('utf-8')\226        227    return images[0], so_img_list, base64_string228 229