File size: 1,496 Bytes
66e15a3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import torch
from transformers import AutoTokenizer, AutoModel

MODEL_DIR = "Llama-3.1-8B_sft/"

_tokenizer = None
_model = None

def _load_model():
    global _tokenizer, _model
    if _model is None:
        print(f"Loading LLaMA model from {MODEL_DIR}...")
        _tokenizer = AutoTokenizer.from_pretrained(MODEL_DIR)
        if _tokenizer.pad_token is None:
            _tokenizer.pad_token = _tokenizer.eos_token
        _model = AutoModel.from_pretrained(MODEL_DIR, dtype=torch.float16)
        _model.eval()
        if torch.cuda.is_available():
            _model = _model.cuda()
    return _tokenizer, _model


def embed_text_llama(texts, batch_size=32, max_length=512):
    """Embed a list of strings using LLaMA mean pooling. Returns CPU float32 tensor [N, 4096]."""
    tokenizer, model = _load_model()
    device = next(model.parameters()).device
    all_embeddings = []

    for i in range(0, len(texts), batch_size):
        batch = texts[i : i + batch_size]
        inputs = tokenizer(batch, return_tensors="pt", padding=True,
                           truncation=True, max_length=max_length).to(device)
        with torch.no_grad():
            outputs = model(**inputs)
        last_hidden = outputs.last_hidden_state  # [B, T, 4096]
        mask = inputs["attention_mask"].unsqueeze(-1).float()
        pooled = (last_hidden * mask).sum(dim=1) / mask.sum(dim=1).clamp(min=1e-9)
        all_embeddings.append(pooled.cpu().float())

    return torch.cat(all_embeddings, dim=0)