(
&mut self,
prompt: &str,
uncond_prompt: &str,
use_guide_scale: bool,
first: bool,
)
| 578 | } |
| 579 | |
| 580 | async fn text_embeddings( |
| 581 | &mut self, |
| 582 | prompt: &str, |
| 583 | uncond_prompt: &str, |
| 584 | use_guide_scale: bool, |
| 585 | first: bool, |
| 586 | ) -> Result<Tensor> { |
| 587 | let tokenizer; |
| 588 | let text_model; |
| 589 | let pad_id; |
| 590 | let max_token_embeddings; |
| 591 | |
| 592 | if first { |
| 593 | tokenizer = &self.tokenizer; |
| 594 | text_model = &mut self.text_model; |
| 595 | pad_id = self.pad_id; |
| 596 | max_token_embeddings = self.sd_config.clip.max_position_embeddings; |
| 597 | } else { |
| 598 | tokenizer = self.tokenizer_2.as_ref().unwrap(); |
| 599 | text_model = self.text_model_2.as_mut().unwrap(); |
| 600 | pad_id = self.pad_id_2.unwrap(); |
| 601 | max_token_embeddings = self |
| 602 | .sd_config |
| 603 | .clip2 |
| 604 | .as_ref() |
| 605 | .unwrap() |
| 606 | .max_position_embeddings; |
| 607 | } |
| 608 | |
| 609 | info!("Running with prompt \"{prompt}\"."); |
| 610 | |
| 611 | let mut tokens = tokenizer |
| 612 | .encode(prompt, true) |
| 613 | .map_err(E::msg)? |
| 614 | .get_ids() |
| 615 | .to_vec(); |
| 616 | |
| 617 | if tokens.len() > max_token_embeddings { |
| 618 | anyhow::bail!( |
| 619 | "the prompt is too long, {} > max-tokens ({})", |
| 620 | tokens.len(), |
| 621 | max_token_embeddings |
| 622 | ) |
| 623 | } |
| 624 | |
| 625 | while tokens.len() < max_token_embeddings { |
| 626 | tokens.push(pad_id) |
| 627 | } |
| 628 | |
| 629 | let tokens = Tensor::new(tokens.as_slice(), &self.context.device)?.unsqueeze(0)?; |
| 630 | |
| 631 | let text_embeddings = text_model |
| 632 | .forward_mut(&tokens, 0, 0, &mut self.context) |
| 633 | .await?; |
| 634 | |
| 635 | let text_embeddings = if use_guide_scale { |
| 636 | let mut uncond_tokens = tokenizer |
| 637 | .encode(uncond_prompt, true) |
no test coverage detected