File size: 17,018 Bytes
8752091
 
 
 
843a654
a468182
8752091
 
 
2ca20a1
8752091
6b7efd8
 
 
 
 
 
 
 
 
8752091
 
6577ccb
8752091
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
843a654
8752091
 
 
 
 
 
 
 
 
843a654
 
 
8752091
 
 
 
843a654
 
8752091
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2ca20a1
8752091
 
 
 
 
 
 
 
 
 
6b7efd8
 
 
 
 
 
 
8752091
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2ca20a1
8752091
 
6b7efd8
8752091
6577ccb
6b7efd8
 
 
 
 
 
 
 
 
8752091
 
 
 
 
 
6577ccb
 
 
 
 
 
8752091
 
 
 
6577ccb
 
 
 
 
8752091
 
 
 
 
 
 
6b7efd8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
8752091
 
 
 
 
 
 
 
 
 
 
 
 
 
6b7efd8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
8752091
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6b7efd8
 
 
 
 
 
 
 
 
 
 
8752091
 
 
2ca20a1
8752091
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
"""Transformers-compatible attention selection for FastPLMs models."""

from __future__ import annotations

import torch

from collections.abc import Mapping
from functools import partial
from typing import Any
from transformers import AttentionInterface, AttentionMaskInterface, PretrainedConfig

from ._auto import (
    AUTO_ATTENTION,
    AttentionResolution,
    attention_execution_context,
    deferred_resolution,
    needs_execution_context,
    provisional_implementation,
    resolve_auto_attention,
)
from ._core import (
    AttentionBackend,
    canonical_checkpoint_attention_backend,
    get_attn_implementation,
    kernels_flash_attention_func,
    resolve_attention_backend,
    set_config_attn_implementation,
)
from ._kernel_lock import require_kernels_package


def _kernels_attention_forward(
    module: torch.nn.Module,
    query: torch.Tensor,
    key: torch.Tensor,
    value: torch.Tensor,
    attention_mask: torch.Tensor | None,
    *,
    implementation: str,
    **kwargs: Any,
) -> tuple[torch.Tensor, None]:
    """Run one canonical FlashAttention backend through Hugging Face kernels.

    Transformers attention functions receive Q, K, and V with shape
    (b, h, l, d) and return an output with shape (b, l, h, d). The shared
    FastPLMs kernel adapter uses the latter layout internally.
    """

    # query, key, value: (b, h, l, d); attention_mask: (b, l) or None
    dropout = float(kwargs.get("dropout", 0.0) or 0.0)
    if module.training and dropout:
        raise RuntimeError(
            "Hugging Face kernels FlashAttention is inference-only when attention dropout "
            "is nonzero. Use SDPA for this training configuration."
        )
    causal = bool(kwargs.get("is_causal", getattr(module, "is_causal", False)))
    softmax_scale = kwargs.get("scaling")
    output = kernels_flash_attention_func(
        query_states=query.transpose(1, 2).contiguous(),  # (b, l, h, d)
        key_states=key.transpose(1, 2).contiguous(),  # (b, l, h, d)
        value_states=value.transpose(1, 2).contiguous(),  # (b, l, h, d)
        attention_mask_2d=attention_mask,
        causal=causal,
        softmax_scale=softmax_scale,
        implementation=implementation,
    )  # (b, l, h, d)
    return output, None  # (b, l, h, d), None


# Keep FastPLMs' kernels-only adapters local to this registry instance.
# ``GeneralInterface.register`` updates Transformers' class-wide mapping, so
# using it here would replace the canonical FlashAttention handlers for every
# model in the process, including models unrelated to FastPLMs.
FASTPLMS_ATTENTION_FUNCTIONS = AttentionInterface()
FASTPLMS_ATTENTION_MASKS = AttentionMaskInterface()
FASTPLMS_ATTENTION_FUNCTIONS["flash_attention_2"] = partial(
    _kernels_attention_forward,
    implementation="flash_attention_2",
)
FASTPLMS_ATTENTION_FUNCTIONS["flash_attention_3"] = partial(
    _kernels_attention_forward,
    implementation="flash_attention_3",
)
for _flash_name in ("flash_attention_2", "flash_attention_3"):
    FASTPLMS_ATTENTION_MASKS[_flash_name] = FASTPLMS_ATTENTION_MASKS[_flash_name]


class FastPLMsAttentionMixin:
    """Synchronize Transformers attention selection with custom model layers.

    Model families retain their checkpoint parameter names. Only runtime
    attributes are updated when ``set_attn_implementation`` is called.
    """

    _supports_sdpa = True
    _supports_flex_attn = True
    # Transformers checks the singular flag during model construction. A
    # family opts in only when its manifest entry advertises at least one of
    # the two FastPLMs kernels-only FlashAttention implementations.
    _supports_flash_attn = False
    _supports_flash_attn_2 = False
    _supports_flash_attn_3 = False
    _fastplms_attention_implementations = (
        "eager",
        "sdpa",
        "flex_attention",
    )
    # Preference order for ``attn_implementation="auto"``, mirrored from the
    # family's ``attention_auto_order`` in models.toml. Empty rejects the request.
    _fastplms_attention_auto_order: tuple[str, ...] = ()
    # Supplied by the Transformers ``PreTrainedModel`` that follows this mixin in every
    # MRO. Each family's configuration class adds the ``attn_backend`` field this mixin
    # reads, and they share no base that declares it, so the boundary is dynamic.
    config: Any

    def _validate_attention_name(self, implementation: str) -> None:
        if implementation not in self._fastplms_attention_implementations:
            raise ValueError(
                f"{type(self).__name__} does not support {implementation!r}; expected one of "
                f"{self._fastplms_attention_implementations}."
            )

    def _check_and_adjust_attn_implementation(
        self,
        attn_implementation: str | None,
        is_init_check: bool = False,
        allow_all_kernels: bool = False,
    ) -> str:
        """Resolve attention without invoking Transformers' source-Flash probe.

        The standard ``flash_attention_2`` and ``flash_attention_3`` names are
        retained for the Transformers API, but FastPLMs resolves them only
        through the exact Hugging Face ``kernels`` artifacts pinned by
        ``models.toml``. Repository-qualified or otherwise external kernels
        are never accepted through this model hook.
        """

        if allow_all_kernels:
            raise ValueError("FastPLMs does not load external attention kernels.")
        if attn_implementation is None:
            return super()._check_and_adjust_attn_implementation(
                None,
                is_init_check=is_init_check,
                allow_all_kernels=False,
            )

        self._validate_attention_name(attn_implementation)
        if attn_implementation in {"flash_attention_2", "flash_attention_3"}:
            if not self._supports_flash_attn:
                raise ValueError(
                    f"{type(self).__name__} does not advertise kernels-only FlashAttention."
                )
            # Validate the lightweight Python dependency here, but defer binary
            # download and import until Q, K, and V have passed the CUDA gate.
            require_kernels_package()
            return attn_implementation

        return super()._check_and_adjust_attn_implementation(
            attn_implementation,
            is_init_check=is_init_check,
            allow_all_kernels=False,
        )

    def __init__(self, config: PretrainedConfig, *args: Any, **kwargs: Any) -> None:
        sentinel = object()
        internal = getattr(config, "_attn_implementation_internal", sentinel)
        stored = getattr(config, "_attn_implementation", None) if internal is sentinel else internal
        legacy = getattr(config, "attn_backend", None)
        requested = stored if stored is not None else legacy
        auto_requested = requested == AUTO_ATTENTION
        serialized_backend: str | None = None
        if auto_requested:
            # The configuration never holds ``auto``. Family layers are built on the
            # provisional implementation, and a saved copy keeps the backend that a
            # named load of the same checkpoint would have stored.
            requested = provisional_implementation(self._require_attention_auto_order())
            serialized_backend = legacy if legacy not in (None, AUTO_ATTENTION) else requested
            stored = None
        if requested is not None:
            if not isinstance(requested, str):
                raise TypeError(
                    "The configured attention implementation must be a string or None; "
                    f"received {type(requested).__name__}."
                )
            # A serialized configuration can name a backend with the historical
            # spelling used by the official source it was converted from. That
            # names the same implementation, so translate it here rather than
            # rejecting a checkpoint that asked for an implementation FastPLMs has.
            canonical = canonical_checkpoint_attention_backend(requested)
            self._validate_attention_name(canonical)
            # ``PreTrainedModel.__init__`` resolves a missing Transformers
            # implementation to the family default.  Legacy FastPLMs configs
            # persist their explicit choice in ``attn_backend``, so forward it
            # into the canonical Transformers field before the base class can
            # replace it with SDPA.  A stored canonical value already agrees and
            # is left untouched, including an explicit
            # ``attn_implementation=...`` load override.
            if canonical != stored:
                set_config_attn_implementation(config, canonical)
        super().__init__(config, *args, **kwargs)
        # Transformers resolves an unspecified implementation during the base
        # model initialization. Synchronize that choice before family layers
        # are constructed.
        resolved = get_attn_implementation(config)
        self._validate_attention_name(resolved)
        set_config_attn_implementation(config, resolved)
        if auto_requested:
            self.__dict__["_fastplms_serialized_attn_backend"] = serialized_backend
            self._begin_auto_attention()

    def _require_attention_auto_order(self) -> tuple[str, ...]:
        order = self._fastplms_attention_auto_order
        if not order:
            raise ValueError(
                f"{type(self).__name__} does not support attn_implementation='auto'; "
                f"request one of {self._fastplms_attention_implementations}."
            )
        return order

    @property
    def attention_resolution(self) -> AttentionResolution | None:
        """The record of an ``auto`` request, or None when a backend was named."""
        return self.__dict__.get("_fastplms_attention_resolution")

    def _begin_auto_attention(self) -> None:
        """Resolve now when no candidate needs a device, else at the first forward."""
        order = self._require_attention_auto_order()
        self._cancel_pending_auto_attention()
        if not needs_execution_context(order):
            resolution = resolve_auto_attention(order, None)
            self._apply_attn_implementation(resolution.resolved)
            self.__dict__["_fastplms_attention_resolution"] = resolution
            return
        resolution = deferred_resolution(order)
        self._apply_attn_implementation(resolution.resolved)
        self.__dict__["_fastplms_attention_resolution"] = resolution
        # The first forward runs inside the caller's autocast context, which is
        # what decides FlashAttention eligibility for FP32 parameters.
        self.__dict__["_fastplms_auto_attention_hook"] = (
            self._as_module().register_forward_pre_hook(_resolve_auto_attention_before_forward)
        )

    def _as_module(self) -> torch.nn.Module:
        if not isinstance(self, torch.nn.Module):
            raise TypeError(
                f"{type(self).__name__} must be a torch.nn.Module to defer attention selection."
            )
        return self

    def _cancel_pending_auto_attention(self) -> None:
        hook = self.__dict__.pop("_fastplms_auto_attention_hook", None)
        if hook is not None:
            hook.remove()

    def resolve_attn_implementation(
        self,
        device: torch.device | str | None = None,
        dtype: torch.dtype | None = None,
    ) -> AttentionResolution:
        """Settle a pending ``auto`` request for the device and dtype of the next forward.

        The first forward does this by itself. Call it earlier, for example before
        ``torch.compile`` or before fingerprinting an embedding run, and pass
        ``dtype`` when the forward will run under an autocast context that is not
        active yet. A settled request returns its record unchanged.
        """
        resolution = self.attention_resolution
        if resolution is None:
            raise RuntimeError(
                f"{type(self).__name__} was not configured with attn_implementation='auto'."
            )
        if not resolution.deferred:
            return resolution
        self._cancel_pending_auto_attention()
        resolution = resolve_auto_attention(
            self._require_attention_auto_order(),
            attention_execution_context(self._as_module(), device=device, dtype=dtype),
        )
        self._apply_attn_implementation(resolution.resolved)
        self.__dict__["_fastplms_attention_resolution"] = resolution
        return resolution

    def save_pretrained(self, *args: Any, **kwargs: Any) -> Any:
        """Save without the machine-specific outcome of an ``auto`` request."""
        # ``save_pretrained`` comes from the ``PreTrainedModel`` later in the MRO.
        if self.attention_resolution is None:
            return super().save_pretrained(*args, **kwargs)  # type: ignore[misc]
        selected_backend = self.config.attn_backend
        self.config.attn_backend = self.__dict__["_fastplms_serialized_attn_backend"]
        try:
            return super().save_pretrained(*args, **kwargs)  # type: ignore[misc]
        finally:
            self.config.attn_backend = selected_backend

    def set_attn_implementation(
        self,
        attn_implementation: str | Mapping[str, str],
        allow_all_kernels: bool = False,
    ) -> None:
        """Select an advertised backend and update every instantiated layer."""
        if isinstance(attn_implementation, Mapping):
            if set(attn_implementation) == {""}:
                attn_implementation = attn_implementation[""]
            else:
                raise ValueError(
                    "FastPLMs models have one attention backbone; pass a string or {'': name}."
                )
        if attn_implementation == AUTO_ATTENTION:
            if allow_all_kernels:
                raise ValueError("FastPLMs does not load external attention kernels.")
            self.__dict__.setdefault(
                "_fastplms_serialized_attn_backend", getattr(self.config, "attn_backend", None)
            )
            self._begin_auto_attention()
            return
        # A named request replaces any earlier automatic selection.
        self._cancel_pending_auto_attention()
        self.__dict__.pop("_fastplms_attention_resolution", None)
        self.__dict__.pop("_fastplms_serialized_attn_backend", None)
        self._apply_attn_implementation(attn_implementation, allow_all_kernels)

    def _apply_attn_implementation(
        self, attn_implementation: str, allow_all_kernels: bool = False
    ) -> None:
        resolved_name = self._check_and_adjust_attn_implementation(
            attn_implementation,
            is_init_check=False,
            allow_all_kernels=allow_all_kernels,
        )
        set_config_attn_implementation(self.config, resolved_name)
        resolved = resolve_attention_backend(resolved_name)
        for module in self.modules():
            if module is self:
                continue
            for attribute in ("attn_backend", "attention_backend", "_attn_backend"):
                if attribute not in module.__dict__:
                    continue
                current = module.__dict__[attribute]
                module.__dict__[attribute] = (
                    resolved if isinstance(current, AttentionBackend) else resolved_name
                )


# Selection reads the manifest, can load a kernel, and rewrites layer attributes.
# It runs eagerly so that a compiled model never traces it.
@torch.compiler.disable  # type: ignore[untyped-decorator]
def _resolve_auto_attention_before_forward(
    module: torch.nn.Module, _arguments: tuple[Any, ...]
) -> None:
    if not isinstance(module, FastPLMsAttentionMixin):
        raise TypeError("The automatic attention hook belongs on a FastPLMs model.")
    module.resolve_attn_implementation()


def validate_transformers_attention_interfaces() -> None:
    """Verify that Transformers exposes functions and masks for every backend.

    The validated Transformers registers these canonical names. The FastPLMs function
    overrides remain instance-local and do not replace process-global handlers.
    """
    function_registry = FASTPLMS_ATTENTION_FUNCTIONS
    mask_registry = FASTPLMS_ATTENTION_MASKS
    missing_functions = [
        name
        for name in (
            "sdpa",
            "flex_attention",
            "flash_attention_2",
            "flash_attention_3",
        )
        if name not in function_registry
    ]
    missing_masks = [
        name
        for name in (
            "eager",
            "sdpa",
            "flex_attention",
            "flash_attention_2",
            "flash_attention_3",
        )
        if name not in mask_registry
    ]
    if missing_functions or missing_masks:
        raise RuntimeError(
            "Transformers attention registry is incomplete: "
            f"functions={missing_functions}, masks={missing_masks}."
        )