Spaces:
Sleeping
Sleeping
| """Rainbow stella embedder na ZeroGPU — batchowy embedding dla pipeline'u. | |
| Ten sam model i to samo kodowanie co lokalny shim (embed_server.py): | |
| SentenceTransformer.encode(normalize_embeddings=True), BEZ prompta zapytania | |
| (pasaż surowy), max_seq_length=8192 — więc wektory dokumentów są SPÓJNE z | |
| wektorami zapytań liczonymi lokalnie (ten sam podprzestrzeń). | |
| """ | |
| import os | |
| import spaces # MUSI być przed torch / transformers (patch torch.cuda.*) | |
| import torch | |
| import gradio as gr | |
| from sentence_transformers import SentenceTransformer | |
| MODEL_ID = os.environ.get("EMBEDDER_MODEL", "sdadas/stella-pl-retrieval-8k") | |
| MAX_SEQ = int(os.environ.get("EMBED_MAX_SEQ", "8192")) | |
| # ładowanie w zakresie modułu, eager .to("cuda") — ZeroGPU pakuje wagi i streamuje | |
| # je do VRAM przy pierwszym wejściu w @spaces.GPU. config_kwargs wyłącza | |
| # memory-efficient attention (xformers) i unpad — działa na czystym torch/SDPA, | |
| # bez wheeli CUDA-extension. | |
| model = SentenceTransformer( | |
| MODEL_ID, device="cuda", trust_remote_code=True, | |
| config_kwargs={"use_memory_efficient_attention": False, "unpad_inputs": False}, | |
| ) | |
| model.max_seq_length = MAX_SEQ | |
| def embed(texts): | |
| """Lista tekstów (pasaże) -> {vectors:[[float]], dim, count}. Znormalizowane L2.""" | |
| if isinstance(texts, str): | |
| texts = [texts] | |
| texts = [("" if t is None else str(t)) for t in (texts or [])] | |
| if not texts: | |
| return {"vectors": [], "dim": 0, "count": 0} | |
| vecs = model.encode(texts, normalize_embeddings=True, batch_size=16, | |
| convert_to_numpy=True, show_progress_bar=False) | |
| return {"vectors": [[float(x) for x in v] for v in vecs], | |
| "dim": int(vecs.shape[1]), "count": int(vecs.shape[0])} | |
| demo = gr.Interface( | |
| fn=embed, | |
| inputs=gr.JSON(label="texts (list[str])"), | |
| outputs=gr.JSON(label="out"), | |
| title="Rainbow stella embedder (ZeroGPU)", | |
| description="POST listy pasaży -> wektory 1024-d (stella-pl-retrieval-8k, znormalizowane).", | |
| api_name="embed", | |
| ) | |
| demo.queue(max_size=128).launch() | |