Encode text using CLIP model.
(text: str, clip_model, device: str)
| 51 | |
| 52 | |
| 53 | def get_text_embedding(text: str, clip_model, device: str) -> torch.Tensor: |
| 54 | """Encode text using CLIP model.""" |
| 55 | try: |
| 56 | with torch.no_grad(): |
| 57 | text_tokens = clip.tokenize([text]).to(device) |
| 58 | text_embedding = clip_model.encode_text(text_tokens) |
| 59 | # text_embedding = text_embedding / text_embedding.norm(dim=-1, |
| 60 | # keepdim=True) |
| 61 | return text_embedding.float() |
| 62 | except Exception as e: |
| 63 | logger.warning(f"Failed to encode text '{text}': {e}") |
| 64 | return torch.zeros(1, 512, device=device, dtype=torch.float32) |
| 65 | |
| 66 | |
| 67 | def interactive_input_thread(loop_state: LoopState): |