File size: 2,079 Bytes
f35855e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
"""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


@spaces.GPU(duration=120)
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()