File size: 11,608 Bytes
80aea5b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
"""HF loading and an explicitly experimental masked-model adapter.

Only imported for an actual evaluation, never for planning.
"""

import torch
from lm_eval.models.huggingface import HFLM
from transformers import AutoTokenizer
from tqdm import tqdm


class UL2HFLM(HFLM):
    """S-denoiser continuation scoring; only answer text contributes to NLL.

    Source: S + context + sentinel_0 + EOS.
    Decoder input: BOS + sentinel_0 + answer[:-1]. Controls are conditioned
    on, never scored; the softmax still includes the complete vocabulary.
    """

    def _encode_pair(self, context, continuation):
        # Match causal HFLM's text boundary, without tokenizer-added controls.
        spaces = len(context) - len(context.rstrip())
        if spaces:
            continuation = context[-spaces:] + continuation
            context = context[:-spaces]
        whole = self.tok_encode(context + continuation, add_special_tokens=False)
        prefix = self.tok_encode(context, add_special_tokens=False)
        return prefix, whole[len(prefix):]

    @torch.inference_mode()
    def _loglikelihood_tokens(self, requests, disable_tqdm=False, override_bs=None):
        contract = self.model.config.ul2
        span = contract["sentinel_ids"][0]
        bos = self.model.config.decoder_start_token_id
        batch_size = override_bs or self.batch_size
        results = [None] * len(requests)
        ordered = sorted(enumerate(requests), key=lambda x: -(len(x[1][1]) + len(x[1][2])))
        for start in tqdm(range(0, len(ordered), batch_size), disable=disable_tqdm, desc="UL2 likelihood"):
            batch = ordered[start:start + batch_size]
            sources, decoders = [], []
            for _, (_, context, answer) in batch:
                if not answer or len(answer) >= self.max_length:
                    raise ValueError("UL2 scoring requires a nonempty answer shorter than the context limit")
                budget = min(self.max_length + 1 - len(answer), self.model.config.max_position_embeddings - 3)
                context = context[-budget:]
                sources.append([contract["mode_ids"]["S"], *context, span, contract["eos_id"]])
                decoders.append([bos, span, *answer[:-1]])

            def pad(rows):
                ids = torch.full((len(rows), max(map(len, rows))), contract["pad_id"], dtype=torch.long, device=self.device)
                mask = torch.zeros_like(ids)
                for i, row in enumerate(rows):
                    ids[i, :len(row)] = torch.tensor(row, device=self.device)
                    mask[i, :len(row)] = 1
                return ids, mask

            source, source_mask = pad(sources)
            decoder, decoder_mask = pad(decoders)
            hidden = self.model.model(input_ids=source, attention_mask=source_mask,
                                      decoder_input_ids=decoder, decoder_attention_mask=decoder_mask,
                                      use_cache=False).last_hidden_state
            # Project only scored positions and bound temporary vocabulary logits.
            for row, (index, (key, _, answer)) in enumerate(batch):
                score, greedy = 0.0, True
                for offset in range(0, len(answer), 256):
                    target = torch.tensor(answer[offset:offset + 256], device=self.device)
                    logits = self.model.lm_head(hidden[row, 1 + offset:1 + offset + len(target)]).float()
                    score += torch.log_softmax(logits, -1).gather(1, target[:, None]).sum().item()
                    greedy = greedy and bool((logits.argmax(-1) == target).all())
                results[index] = (score, greedy)
                if key is not None:
                    self.cache_hook.add_partial("loglikelihood", key, results[index])
        return results

    def loglikelihood_rolling(self, requests, disable_tqdm=False):
        raise NotImplementedError("UL2 continuation likelihood is not rolling autoregressive perplexity")

    def generate_until(self, requests, disable_tqdm=False):
        raise NotImplementedError("UL2 benchmark adapter currently supports likelihood tasks only")


class PrefixHFLM(HFLM):
    """Score causal answers conditioned on a fully bidirectional text prefix."""

    @torch.inference_mode()
    def _loglikelihood_tokens(self, requests, disable_tqdm=False, override_bs=None):
        results = [None] * len(requests)
        ordered = sorted(enumerate(requests), key=lambda x: -(len(x[1][1]) + len(x[1][2])))
        size = override_bs or self.batch_size
        for start in tqdm(range(0, len(ordered), size), disable=disable_tqdm, desc="Prefix likelihood"):
            batch = ordered[start:start + size]
            rows, prefixes = [], []
            for _, (_, context, answer) in batch:
                if not answer or len(answer) > self.max_length:
                    raise ValueError("Require nonempty answer no longer than context limit")
                context = context[-(self.max_length + 1 - len(answer)):]
                rows.append(context + answer[:-1])
                prefixes.append(len(context))
            ids = torch.full((len(rows), max(map(len, rows))), self.tokenizer.pad_token_id or 0,
                             device=self.device, dtype=torch.long)
            mask = torch.zeros_like(ids)
            for i, row in enumerate(rows):
                ids[i, :len(row)] = torch.tensor(row, device=self.device)
                mask[i, :len(row)] = 1
            hidden = self.model.model(ids, attention_mask=mask,
                                      prefix_lengths=torch.tensor(prefixes, device=self.device),
                                      use_cache=False).last_hidden_state
            for i, (index, (key, _, answer)) in enumerate(batch):
                score, greedy = 0., True
                for offset in range(0, len(answer), 256):
                    target = torch.tensor(answer[offset:offset+256], device=self.device)
                    logits = self.model.lm_head(hidden[i, prefixes[i]-1+offset:prefixes[i]-1+offset+len(target)]).float()
                    score += torch.log_softmax(logits, -1).gather(1, target[:, None]).sum().item()
                    greedy = greedy and bool((logits.argmax(-1) == target).all())
                results[index] = (score, greedy)
                if key is not None:
                    self.cache_hook.add_partial("loglikelihood", key, results[index])
        return results

    def loglikelihood_rolling(self, requests, disable_tqdm=False):
        raise NotImplementedError("Use an explicitly defined prefix-continuation protocol")


class DiffusionHFLM(HFLM):
    """Continuation PLL, with other answer tokens visible.

    The bool is full-token reconstruction accuracy, NOT AR exact match.
    Generation uses one masked next-token slot at a time, not the model's
    unconditional parallel denoiser. Both are experimental protocols.
    """

    @torch.inference_mode()
    def _loglikelihood_tokens(self, requests, disable_tqdm=False, override_bs=None):
        results = []
        for cache_key, context, continuation in tqdm(requests, disable=disable_tqdm, desc="Diffusion PLL"):
            if not continuation:
                results.append((0.0, True))
                continue
            if len(continuation) >= self.max_length:
                raise ValueError("Diffusion PLL needs room for context and the entire continuation")
            context = context[-(self.max_length - len(continuation)):]
            tokens = context + continuation
            base = torch.tensor([tokens], dtype=torch.long, device=self.device)
            timestep = torch.tensor([1 / len(tokens)], device=self.device)
            score, greedy = 0.0, True
            # Independent single-mask replicas preserve the PLL definition.
            # Project only masked positions instead of every vocabulary logit.
            for start in range(len(context), len(tokens), 16):
                positions = torch.arange(start, min(start + 16, len(tokens)), device=self.device)
                rows = torch.arange(len(positions), device=self.device)
                masked = base.expand(len(positions), -1).clone()
                masked[rows, positions] = self.model.config.mask_token_id
                hidden = self.model.model(input_ids=masked, timesteps=timestep.expand(len(positions)))
                logits = self.model.lm_head(hidden[rows, positions]).float()
                targets = base[0, positions]
                score += torch.log_softmax(logits, dim=-1).gather(1, targets[:, None]).sum().item()
                greedy = greedy and bool((logits.argmax(-1) == targets).all())
            result = (score, greedy)
            results.append(result)
            if cache_key is not None:
                self.cache_hook.add_partial("loglikelihood", cache_key, result)
        return results

    def loglikelihood_rolling(self, requests, disable_tqdm=False):
        raise NotImplementedError("Diffusion PLL must not be reported as autoregressive perplexity")

    @torch.inference_mode()
    def _model_generate(self, context, max_length, stop, **generation_kwargs):
        if context.shape[0] != 1:
            raise ValueError("Diffusion generation requires batch size 1")
        if generation_kwargs.get("do_sample", False):
            raise ValueError("Diffusion adapter implements greedy generation only")
        sequence = context
        prompt_length = context.shape[1]
        while sequence.shape[1] < min(max_length, self.max_length):
            slot = torch.full((1, 1), self.model.config.mask_token_id,
                              dtype=torch.long, device=self.device)
            masked = torch.cat((sequence, slot), dim=1)
            timestep = torch.tensor([1 / masked.shape[1]], device=self.device)
            logits = self.model(input_ids=masked, timesteps=timestep).logits[:, -1]
            predicted = logits.argmax(dim=-1, keepdim=True)
            sequence = torch.cat((sequence, predicted), dim=1)
            text = self.tok_decode(sequence[0, prompt_length:].tolist())
            if predicted.item() == self.eot_token_id or any(s and s in text for s in stop):
                break
        return sequence


def load_model(plan, args):
    common = dict(batch_size=args.batch_size, device=args.device, dtype=args.dtype,
                  max_length=args.max_length, trust_remote_code=args.trust_remote_code,
                  revision=args.revision)
    tokenizer = AutoTokenizer.from_pretrained(
        plan["tokenizer"], trust_remote_code=args.trust_remote_code,
        revision=args.tokenizer_revision or args.revision,
    )
    if plan["local_model"] is None:
        return HFLM(pretrained=plan["checkpoint"], tokenizer=tokenizer, **common)

    from tiny_llm.models import MODEL_REGISTRY

    entry = MODEL_REGISTRY[plan["local_model"]]
    dtype = args.dtype if args.dtype == "auto" else getattr(torch, args.dtype)
    # Checkpoints have no auto_map or tokenizer. Never load train_state.pt.
    config = entry["config_class"].from_pretrained(plan["checkpoint"])
    model = entry["model_class"].from_pretrained(plan["checkpoint"], config=config, dtype=dtype)
    model.to(args.device).eval()
    backend = "seq2seq" if plan["kind"] == "seq2seq" else "causal"
    adapter = DiffusionHFLM if plan["kind"] == "diffusion" else HFLM
    if plan['local_model'] == 'prefixlm':
        adapter = PrefixHFLM
    if plan["kind"] == "seq2seq" and getattr(config, "ul2", None):
        adapter = UL2HFLM
    return adapter(pretrained=model, tokenizer=tokenizer, backend=backend, **common)