| 59 | ) |
| 60 | |
| 61 | def get_text_embeddings(self, texts, device, text_max_length=256): |
| 62 | inputs = self.text_tokenizer( |
| 63 | texts, |
| 64 | return_tensors="pt", |
| 65 | padding=True, |
| 66 | truncation=True, |
| 67 | max_length=text_max_length, |
| 68 | ) |
| 69 | inputs = {key: value.to(device) for key, value in inputs.items()} |
| 70 | with torch.no_grad(): |
| 71 | outputs = self.text_encoder_model(**inputs) |
| 72 | last_hidden_states = outputs.last_hidden_state |
| 73 | attention_mask = inputs["attention_mask"] |
| 74 | return last_hidden_states, attention_mask |
| 75 | |
| 76 | def diffusion_process( |
| 77 | self, |