mtg-draft-viz / src /bert /embedd_text.py
TimoBertram's picture
Upload src/bert/embedd_text.py with huggingface_hub
66e15a3 verified
Raw
History Blame Contribute Delete
1.5 kB
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)