File size: 5,573 Bytes
14bf8c2
2d8be88
 
14bf8c2
 
2d8be88
14bf8c2
 
 
2d8be88
 
 
 
 
 
 
 
 
 
 
14bf8c2
2d8be88
14bf8c2
 
 
 
2d8be88
 
 
14bf8c2
2d8be88
 
14bf8c2
2d8be88
14bf8c2
 
2d8be88
14bf8c2
 
 
 
2d8be88
 
 
 
 
 
14bf8c2
 
 
 
 
 
 
 
 
2d8be88
 
 
14bf8c2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2d8be88
 
 
 
14bf8c2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
"""Self-contained Hugging Face runtime for Pivot."""
import torch
from torch import nn
from transformers import AutoConfig, PreTrainedModel

from .configuration_pivot import PivotConfig
from .modeling_lfm2_bidirectional import Lfm2BidirectionalModel
from .pivot_model import HuggingFaceEncoder, MLPScorer, SetBrierEncoder, StubEncoder
from .pivot_infer import decide_native, decide_typed, predict


class PivotModel(PreTrainedModel):
    config_class = PivotConfig
    base_model_prefix = "network"
    _no_split_modules = ["SetBrierEncoder"]

    def __init__(self, config):
        super().__init__(config)
        cfg = config.dsbt_config
        backbone = cfg["backbone"]

        if backbone.get("encoder") == "stub":
            encoder = StubEncoder(
                int(backbone.get("stub_vocab_size", 256)),
                int(backbone.get("stub_hidden_size", 32)),
            )
        else:
            body_config = dict(config.encoder_config)
            model_type = body_config.pop("model_type")
            body = Lfm2BidirectionalModel(AutoConfig.for_model(model_type, **body_config))
            encoder = HuggingFaceEncoder.__new__(HuggingFaceEncoder)
            nn.Module.__init__(encoder)
            encoder.model = body
            encoder.hidden_size = int(body.config.hidden_size)
            encoder.pretrained_audit = {"packaged_runtime": True}

        hidden = encoder.hidden_size
        scorer_cfg = cfg["scorer"]
        if scorer_cfg.get("type", "mlp") != "mlp":
            raise ValueError("This Pivot package supports the reviewed MLP scorer only")
        scorer = MLPScorer(hidden, int(scorer_cfg.get("hidden_size", hidden)))
        self.network = SetBrierEncoder(encoder, scorer)
        self.post_init()

    def forward(self, ctx_ids, ctx_mask, opt_ids, opt_mask, opt_attn):
        return self.network(ctx_ids, ctx_mask, opt_ids, opt_mask, opt_attn)

    def _limits(self):
        cfg = self.config.dsbt_config
        data = cfg["data"]
        serving = cfg.get("serving") or {}
        return (
            int(serving.get("max_context_tokens", data["max_context_tokens"])),
            int(serving.get("max_option_tokens", data["max_option_tokens"])),
        )

    @torch.no_grad()
    def decide(self, tokenizer, state, questions):
        self.eval()
        max_context_tokens, max_option_tokens = self._limits()
        return decide_typed(
            self.network,
            tokenizer,
            state,
            questions,
            max_context_tokens=max_context_tokens,
            max_option_tokens=max_option_tokens,
            device=next(self.parameters()).device,
            model_id="Pivot",
        )

    @torch.no_grad()
    def decide_native(self, tokenizer, context, candidates):
        self.eval()
        max_context_tokens, max_option_tokens = self._limits()
        return decide_native(
            self.network,
            tokenizer,
            context,
            candidates,
            max_context_tokens=max_context_tokens,
            max_option_tokens=max_option_tokens,
            device=next(self.parameters()).device,
        )

    @torch.no_grad()
    def choose(self, tokenizer, context, options):
        self.eval()
        max_context_tokens, max_option_tokens = self._limits()
        return predict(
            self.network,
            tokenizer,
            context,
            list(options),
            max_context_tokens=max_context_tokens,
            max_option_tokens=max_option_tokens,
            device=next(self.parameters()).device,
        )

    @torch.no_grad()
    def _encode_texts(self, tokenizer, texts, *, max_length, padding):
        batch = tokenizer(
            list(texts),
            max_length=int(max_length),
            padding=padding,
            truncation=True,
            return_tensors="pt",
        )
        device = next(self.parameters()).device
        return self.network.encode(
            batch["input_ids"].to(device),
            batch["attention_mask"].to(device),
        )

    @torch.no_grad()
    def encode_context(self, tokenizer, context):
        self.eval()
        max_context_tokens, _ = self._limits()
        return self._encode_texts(
            tokenizer,
            [context],
            max_length=max_context_tokens,
            padding=True,
        )[0]

    @torch.no_grad()
    def encode_candidates(self, tokenizer, candidate_texts):
        self.eval()
        _, max_option_tokens = self._limits()
        return self._encode_texts(
            tokenizer,
            list(candidate_texts),
            max_length=max_option_tokens,
            padding="max_length",
        )

    @torch.no_grad()
    def choose_cached(self, context_embedding, candidate_embeddings, candidate_texts):
        self.eval()
        texts = list(candidate_texts)
        if candidate_embeddings.ndim == 2:
            k = int(candidate_embeddings.shape[0])
        elif candidate_embeddings.ndim == 3 and candidate_embeddings.shape[0] == 1:
            k = int(candidate_embeddings.shape[1])
        else:
            raise ValueError("candidate_embeddings must be [K,d] or [1,K,d]")
        if len(texts) != k or k < 2:
            raise ValueError("candidate embedding/text mismatch")
        out = self.network.score_preencoded(context_embedding, candidate_embeddings)
        probs = out.probs[0].detach().cpu().float()
        index = int(out.pred_index[0].item())
        return {
            "choice": texts[index],
            "index": index,
            "probs": probs.tolist(),
        }