SAI-Embedding_100M / EmbeddingModel.py
thongbuind's picture
Đóng gói AutoModel/AutoTokenizer (trust_remote_code), xoá bố cục cũ
7812b29 verified
Raw History Blame Contribute Delete
6.39 kB
import json
from contextlib import nullcontext
from pathlib import Path
import torch
import torch.nn as nn
import torch.nn.functional as F
from .TransformerModel import TransformerModel
project_root = Path(__file__).resolve().parent.parent.parent
config_dir = project_root / "config"
def build_backbone(model_size: str) -> TransformerModel:
with open(config_dir / "base.json", "r") as f:
config = json.load(f)
with open(config_dir / f"{model_size}.json", "r") as f:
config.update(json.load(f))
return TransformerModel(
config["vocab_size"], config["d_model"], config["num_heads"], config["num_kv_heads"],
config["num_layers"], config["ff_dim"], config["max_seq_len"], config["dropout"],
)
def _extract_state_dict(obj):
"""Nhận state_dict thuần (pretrained_*.pt), checkpoint resume ({model_state_dict})
hoặc checkpoint embedding ({embedding_config, state_dict})."""
config = {}
if isinstance(obj, dict) and "embedding_config" in obj:
config, state_dict = obj["embedding_config"], obj["state_dict"]
elif isinstance(obj, dict) and "model_state_dict" in obj:
state_dict = obj["model_state_dict"]
else:
state_dict = obj
# Bỏ prefix do torch.compile thêm vào (nếu có).
state_dict = {k.replace("_orig_mod.", ""): v for k, v in state_dict.items()}
return config, state_dict
class EmbeddingModel(nn.Module):
"""SAI backbone -> pooling -> vector câu.
- causal=False: attention hai chiều (cần adapt bằng MNTP trước khi dùng).
- pooling="mean": trung bình các token không phải PAD; "last": token cuối
(chỉ hợp lý khi causal=True, dùng để so sánh baseline).
- Không có projection head: vector 768 chiều, cắt được theo Matryoshka (mrl_dims).
"""
def __init__(self, backbone: TransformerModel, model_size: str, causal: bool = False,
pooling: str = "mean", query_prefix: str = "", passage_prefix: str = "",
mrl_dims=None):
super().__init__()
assert pooling in ("mean", "last"), pooling
self.backbone = backbone
self.model_size = model_size
self.causal = causal
self.pooling = pooling
self.query_prefix = query_prefix
self.passage_prefix = passage_prefix
self.mrl_dims = list(mrl_dims or [backbone.d_model])
@property
def embedding_config(self):
return {
"model_size": self.model_size, "causal": self.causal, "pooling": self.pooling,
"query_prefix": self.query_prefix, "passage_prefix": self.passage_prefix,
"mrl_dims": self.mrl_dims,
}
@classmethod
def from_checkpoint(cls, path, map_location="cpu", defaults=None, **overrides):
"""Thứ tự ưu tiên cấu hình: overrides > cấu hình lưu trong checkpoint > defaults.
Checkpoint LM thuần (pretrained_100M.pt, mntp_100M.pt) không lưu cấu hình,
nên truyền defaults lấy từ config/embedding.json."""
obj = torch.load(path, map_location=map_location, weights_only=False)
saved_config, state_dict = _extract_state_dict(obj)
config = {"model_size": "100M", "causal": False, "pooling": "mean",
"query_prefix": "", "passage_prefix": "", "mrl_dims": None}
config.update({k: v for k, v in (defaults or {}).items() if k in config})
config.update(saved_config)
config.update({k: v for k, v in overrides.items() if v is not None})
backbone = build_backbone(config["model_size"])
backbone.load_state_dict(state_dict)
return cls(backbone, **config)
def save(self, path):
path = Path(path)
path.parent.mkdir(parents=True, exist_ok=True)
torch.save({"embedding_config": self.embedding_config,
"state_dict": self.backbone.state_dict()}, path)
def pool(self, hidden: torch.Tensor, attention_mask: torch.Tensor) -> torch.Tensor:
mask = attention_mask.to(hidden.dtype)
if self.pooling == "mean":
summed = (hidden.float() * mask.float().unsqueeze(-1)).sum(dim=1)
return summed / mask.float().sum(dim=1, keepdim=True).clamp_min(1.0)
last = attention_mask.float().sum(dim=1).long() - 1
return hidden[torch.arange(hidden.size(0), device=hidden.device), last].float()
def forward(self, input_ids, attention_mask, has_padding: bool = True, normalize: bool = True):
hidden = self.backbone.forward_features(
input_ids, attention_mask, has_padding=has_padding, causal=self.causal,
)
pooled = self.pool(hidden, attention_mask)
return F.normalize(pooled, dim=-1) if normalize else pooled
@torch.no_grad()
def encode(self, texts, tokenizer, is_query: bool = True, max_len: int = 512,
batch_size: int = 128, dim: int = None, prefix: str = None, show_progress: bool = False):
"""Encode list text -> tensor (N, dim) float32 đã L2-normalize, trên CPU.
Text được sắp theo độ dài để giảm padding rồi trả về đúng thứ tự ban đầu.
"""
if prefix is None:
prefix = self.query_prefix if is_query else self.passage_prefix
device = next(self.parameters()).device
amp = torch.autocast("cuda", dtype=torch.bfloat16) if device.type == "cuda" else nullcontext()
was_training = self.training
self.eval()
tokenized = tokenizer.encode(list(texts), prefix, max_len)
order = sorted(range(len(tokenized)), key=lambda i: -len(tokenized[i]))
dim = dim or self.backbone.d_model
out = torch.empty(len(tokenized), dim, dtype=torch.float32)
for step, start in enumerate(range(0, len(order), batch_size)):
idx = order[start:start + batch_size]
input_ids, attention_mask, has_padding = tokenizer.pad([tokenized[i] for i in idx])
with amp:
emb = self(input_ids.to(device), attention_mask.to(device), has_padding, normalize=False)
out[idx] = F.normalize(emb[:, :dim].float(), dim=-1).cpu()
if show_progress and step % 50 == 0:
print(f" encode {start + len(idx):,}/{len(order):,}", end="\r")
if show_progress:
print()
self.train(was_training)
return out