Setup the embedding model for dense vectors
()
| 452 | return client |
| 453 | |
| 454 | def setup_embedding_model(): |
| 455 | """Setup the embedding model for dense vectors""" |
| 456 | try: |
| 457 | embedding_model = TextEmbedding( |
| 458 | model_name="thenlper/gte-base", |
| 459 | providers=["CUDAExecutionProvider", "CPUExecutionProvider"], |
| 460 | model_kwargs={"torch_dtype": "float16"}, |
| 461 | max_length=512, |
| 462 | ) |
| 463 | test_texts = ["test"] |
| 464 | _ = list(embedding_model.embed(test_texts)) |
| 465 | print("Using GPU-accelerated embeddings") |
| 466 | return embedding_model |
| 467 | except Exception as e: |
| 468 | print(f"GPU not available, falling back to CPU: {e}") |
| 469 | return TextEmbedding(model_name="thenlper/gte-base", max_length=512) |
| 470 | |
| 471 | def get_beir_dataset(dataset: str) -> Tuple[Dict, Dict, Dict]: |
| 472 | """Download and load BEIR dataset""" |