File size: 11,776 Bytes
da1fe36
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
"""One-pass typed decisions on a causal LFM backbone.

The model sees the state and the question once, and the answer is read off the
logits at the answer slot. Nothing is decoded, so a schema violation is
impossible and `tokens generated per decision` is zero.
"""

from __future__ import annotations

import math
from typing import Any, Sequence

import torch

from .api import SystemOneApi
from .lfm2_vl import Lfm2ForCausalLM, Lfm2VlForConditionalGeneration
from .prompt import (
    DEFAULT_MODEL,
    DEFAULT_STATE_STYLE,
    DEFAULT_SYSTEM,
    IM_START,
    Question,
    default_lead,
    prefix_text,
    readout,
    render,
    suffix_text,
)

# Every picture is bounded at this many pixels before the processor.
VISION_MAX_PIXELS = 1024 * 1024


# LFM2 models run on the hybrid stack (`hybrid.py`); any other model type through transformers as it is.
MODELS = {
    "lfm2_vl": Lfm2VlForConditionalGeneration,
    "lfm2": Lfm2ForCausalLM,
}


def load_backbone(model_id: str = DEFAULT_MODEL, dtype=torch.bfloat16):
    """The checkpoint and its tokenizer, with SDPA attention."""
    from transformers import AutoModelForImageTextToText, AutoTokenizer, PretrainedConfig

    tok = AutoTokenizer.from_pretrained(model_id, trust_remote_code=True)
    if tok.pad_token_id is None:
        tok.pad_token = tok.eos_token
    kind = PretrainedConfig.get_config_dict(model_id)[0].get("model_type")
    cls = MODELS.get(kind, AutoModelForImageTextToText)
    return cls.from_pretrained(model_id, dtype=dtype, attn_implementation="sdpa"), tok


def cap_pixels(image, max_pixels: int = VISION_MAX_PIXELS):
    """A picture downscaled to at most `max_pixels`, bicubic."""
    image = image.convert("RGB") if hasattr(image, "convert") else image
    w, h = image.size
    if w * h <= max_pixels:
        return image
    try:
        from PIL import Image
    except ImportError as e:  # optional for text
        raise ImportError("resizing a picture needs Pillow") from e

    scale = math.sqrt(max_pixels / (w * h))
    return image.resize((max(1, int(w * scale)), max(1, int(h * scale))), Image.Resampling.BICUBIC)


class SystemOne(SystemOneApi):
    """State in, calibrated distribution out, one forward pass."""

    def __init__(
        self,
        model_id: str = DEFAULT_MODEL,
        device: str | None = None,
        calibration=None,
        lead: str | None = None,
        state_style: str = DEFAULT_STATE_STYLE,
        system: str = DEFAULT_SYSTEM,
        option_style: str = "desc",
        compile: bool = False,
        token_budget: int = 65536,
        model=None,
        tokenizer=None,
    ):
        """`model` and `tokenizer`, when given, are a backbone already loaded (`D1Model` passes itself);
        it stays on its device unless `device` says otherwise."""
        if model is None:
            model, tokenizer = load_backbone(model_id)
        else:
            model_id, device = model.config._name_or_path, device or next(model.parameters()).device
        self.model, self.tokenizer = model, tokenizer
        self.model_id = model_id
        if lead is None:
            lead = default_lead(self.model.config.model_type)
        self.device = torch.device(device or ("cuda" if torch.cuda.is_available() else "cpu"))
        self.model.to(self.device).eval()
        bos = getattr(self.tokenizer, "bos_token", None)
        self.bos = bos if isinstance(bos, str) else ""
        self.calibration = calibration
        self.lead = lead
        self.state_style = state_style
        self.system = system
        self.option_style = option_style
        self.token_budget = token_budget
        self.processor = None
        # CUDA graphs for single questions on NVIDIA, where eager time is mostly kernel launches.
        if compile and torch.version.hip:
            raise ValueError("compile=True needs CUDA: on ROCm the CUDA graphs fault after a few dozen calls")
        self._one_pass = (torch.compile(self.model.forward, mode="reduce-overhead")
                          if compile else self.model)

    # ---------------------------------------------------------------- prompt

    def render(self, state: Any, q: Question) -> str:
        return render(
            self.tokenizer, state, q, self.bos, self.lead, self.state_style,
            self.system, self.option_style,
        )

    # --------------------------------------------------------------- forward

    def _logz_ids(self, rows: list[list[int]]) -> list[torch.Tensor]:
        """Log-softmax at the answer slot, one row per token list, in one pass.

        The rows are one tree (`hybrid.py`): their common start is its trunk and
        is read once; the rest of each row is a branch, packed with no padding.
        Mathematically each row alone; in bf16 the kernels differ by batch shape.
        """
        if len(rows) == 1:  # nothing to share: a plain chain is faster than a tree of one
            row = self._one_pass(input_ids=torch.tensor(rows, device=self.device), logits_to_keep=1).logits[0, -1]
            return [row.float() - torch.logsumexp(row.float(), dim=-1)]
        shared = 0  # every row keeps at least its last token
        while shared < min(map(len, rows)) - 1 and len({r[shared] for r in rows}) == 1:
            shared += 1
        return self._tree_logz(rows[0][:shared], [r[shared:] for r in rows])

    def _tree_logz(self, trunk: list[int], rows: list[list[int]], **vision) -> list[torch.Tensor]:
        """`trunk` then each of `rows`, read at each row's end."""
        packed = torch.tensor([t for r in rows for t in r], device=self.device)
        lengths = torch.tensor([len(r) for r in rows], device=self.device)
        trunk = torch.tensor([trunk], dtype=torch.long, device=self.device)
        logits = self.model.answer(trunk, packed, lengths, **vision).float()
        return list(logits - torch.logsumexp(logits, dim=-1, keepdim=True))

    def plan_batches(self, texts: Sequence[str], token_budget: int | None = None) -> list[list[int]]:
        """Consecutive batches of at most `token_budget` tokens: rows are packed,
        so a batch costs its real tokens."""
        return self._plan([len(self.tokenizer.encode(t, add_special_tokens=False)) for t in texts], token_budget)

    def _plan(self, lengths: Sequence[int], token_budget: int | None = None) -> list[list[int]]:
        if not lengths:
            return []
        budget = token_budget or self.token_budget
        out: list[list[int]] = [[]]
        used = 0
        for i, n in enumerate(lengths):
            if out[-1] and used + n > budget:
                out.append([])
                used = 0
            out[-1].append(i)
            used += n
        return out

    # --------------------------------------------------------------- readout

    def _readout(self, q: Question, logz: torch.Tensor) -> list[float]:
        return readout(self.tokenizer, q, logz, self.calibration)

    # ------------------------------------------------------------------- api

    @torch.inference_mode()
    def run(self, requests: Sequence[tuple[Any, list[Question], Sequence]]) -> list[tuple[list[list[float]], int]]:
        """Each `(state, questions, images)` request's probabilities and the tokens it read. Requests of one
        question and no pictures are packed together, one tree per token budget; any other request is its
        own pass, its state (and pictures) the trunk and its questions the branches."""
        out: list = [None] * len(requests)
        single = [i for i, (_, qs, images) in enumerate(requests) if len(qs) == 1 and not images]
        rows = [self.tokenizer.encode(self.render(requests[i][0], requests[i][1][0]), add_special_tokens=False)
                for i in single]
        for chunk in self._plan([len(r) for r in rows]):
            for j, z in zip(chunk, self._logz_ids([rows[j] for j in chunk])):
                out[single[j]] = ([self._readout(requests[single[j]][1][0], z)], len(rows[j]))
        for i, (state, qs, images) in enumerate(requests):
            if out[i] is None:
                out[i] = self._request(state, qs, images)
        return out

    def _request(self, state: Any, qs: list[Question], images: Sequence) -> tuple[list[list[float]], int]:
        pics = [cap_pixels(im) for im in images]
        prefix = prefix_text(self.tokenizer, state, self.bos, self.state_style, self.system,
                             self._image_markup(len(pics)) if pics else "")
        suffixes = [suffix_text(self.tokenizer, q, self.lead, self.option_style) for q in qs]
        vision: dict = {}
        if not pics:
            trunk = self.tokenizer.encode(prefix, add_special_tokens=False)
        elif len(qs) == 1:  # the whole prompt in one plain pass
            inputs = self._image_inputs(prefix + suffixes[0], pics)
            row = self._one_pass(**inputs, logits_to_keep=1).logits[0, -1].float()
            return [self._readout(qs[0], row - torch.logsumexp(row, dim=-1))], int(inputs["input_ids"].shape[1])
        else:
            vision = self._image_inputs(prefix, pics)
            trunk = vision.pop("input_ids")[0].tolist()
            vision.pop("attention_mask", None)
        branches = [self.tokenizer.encode(s, add_special_tokens=False) for s in suffixes]
        probs: list[list[float]] = []
        for chunk in self._plan([len(b) for b in branches]):
            zs = self._tree_logz(trunk, [branches[j] for j in chunk], **vision)
            probs += [self._readout(qs[j], z) for j, z in zip(chunk, zs)]
        return probs, len(trunk) + sum(map(len, branches))

    def tokens(self, state: Any, questions: Sequence[Question]) -> int:
        """The longest prompt one of `questions` makes over a text state: its state's tokens and its own."""
        prefix = prefix_text(self.tokenizer, state, self.bos, self.state_style, self.system)
        return len(self.tokenizer.encode(prefix, add_special_tokens=False)) + max(
            len(self.tokenizer.encode(suffix_text(self.tokenizer, q, self.lead, self.option_style),
                                      add_special_tokens=False)) for q in questions)

    # ---------------------------------------------------------------- vision

    def _image_markup(self, n: int) -> str:
        """What the chat template writes for `n` images at the head of a user turn (`<image>` each on LFM2-VL)."""
        if self.processor is None:
            self.processor = self._load_processor()
        msgs = [{"role": "user", "content": [*([{"type": "image"}] * n), {"type": "text", "text": "\x00"}]}]
        text = self.processor.apply_chat_template(msgs, add_generation_prompt=False, tokenize=False)
        head = f"{IM_START}user\n"
        return text[text.index(head) + len(head):text.index("\x00")]

    def _load_processor(self):
        from transformers import AutoProcessor

        return AutoProcessor.from_pretrained(self.model.config._name_or_path, trust_remote_code=True)

    def _image_inputs(self, text: str, images: Sequence) -> dict:
        """Token ids and pixel inputs for one prompt holding `images`; the prompt carries its own BOS."""
        inputs = self.processor(text=[text], images=[list(images)], return_tensors="pt", add_special_tokens=False)
        if "pixel_attention_mask" in inputs:
            # LFM2-VL's processor pads every image to 1024 patches and the tower
            # masks the padding out; cutting it gives the same answer for up to
            # half the work.
            n = int(inputs["pixel_attention_mask"].sum(1).max())
            inputs["pixel_values"] = inputs["pixel_values"][:, :n]
            inputs["pixel_attention_mask"] = inputs["pixel_attention_mask"][:, :n]
        return {k: (v.to(self.device) if hasattr(v, "to") else v) for k, v in inputs.items()}