(text_encoders, tokenizers, prompt, text_input_ids_list=None)
| 299 | |
| 300 | # Adapted from pipelines.StableDiffusionXLPipeline.encode_prompt |
| 301 | def encode_prompt(text_encoders, tokenizers, prompt, text_input_ids_list=None): |
| 302 | prompt_embeds_list = [] |
| 303 | |
| 304 | for i, text_encoder in enumerate(text_encoders): |
| 305 | if tokenizers is not None: |
| 306 | tokenizer = tokenizers[i] |
| 307 | text_input_ids = tokenize_prompt(tokenizer, prompt) |
| 308 | else: |
| 309 | assert text_input_ids_list is not None |
| 310 | text_input_ids = text_input_ids_list[i] |
| 311 | |
| 312 | prompt_embeds = text_encoder( |
| 313 | text_input_ids.to(text_encoder.device), |
| 314 | output_hidden_states=True, |
| 315 | ) |
| 316 | |
| 317 | # We are only ALWAYS interested in the pooled output of the final text encoder |
| 318 | pooled_prompt_embeds = prompt_embeds[0] |
| 319 | prompt_embeds = prompt_embeds.hidden_states[-2] |
| 320 | bs_embed, seq_len, _ = prompt_embeds.shape |
| 321 | prompt_embeds = prompt_embeds.view(bs_embed, seq_len, -1) |
| 322 | prompt_embeds_list.append(prompt_embeds) |
| 323 | |
| 324 | prompt_embeds = torch.concat(prompt_embeds_list, dim=-1) |
| 325 | pooled_prompt_embeds = pooled_prompt_embeds.view(bs_embed, -1) |
| 326 | return prompt_embeds, pooled_prompt_embeds |
| 327 | |
| 328 | |
| 329 | def tokenize_prompt(tokenizer, prompt): |
no test coverage detected