File size: 6,387 Bytes
33f617f
 
 
 
 
 
7812b29
33f617f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
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