File size: 6,936 Bytes
f0a4e91
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
# Modified from https://github.com/microsoft/Lens
"""GPT-OSS text encoder for Lens.



We subclass ``transformers.GptOssForCausalLM`` so we can:



1. Return hidden states *only* at a configured layer subset (default

   ``[5, 11, 17, 23]``), avoiding the memory cost of HF's stock

   ``output_hidden_states=True`` which materializes every layer.

2. Early-exit after the last selected layer, since we don't need the

   downstream LM head at all when extracting features.



Standard ``generate(...)`` is inherited unchanged and is used by the optional

prompt reasoner.

"""
from __future__ import annotations

from typing import List, Optional, Sequence

import torch

try:
    from transformers.masking_utils import (create_causal_mask,
                                            create_sliding_window_causal_mask)
    from transformers.models.gpt_oss.modeling_gpt_oss import GptOssForCausalLM
    _HAS_GPT_OSS = True
except ImportError:
    _HAS_GPT_OSS = False
    GptOssForCausalLM = None


if _HAS_GPT_OSS:

    class LensGptOssEncoder(GptOssForCausalLM):
        """``GptOssForCausalLM`` subclass that exposes selected hidden states."""

        def set_selected_layers(self, layer_indices: Sequence[int]) -> None:
            layers = [int(i) for i in layer_indices]
            if not layers:
                raise ValueError("layer_indices must be non-empty")
            if len(set(layers)) != len(layers):
                raise ValueError(f"layer_indices must be unique; got {layers}")
            if min(layers) < 0 or max(layers) >= len(self.model.layers):
                raise ValueError(
                    f"layer_indices out of range; got {layers}, "
                    f"model has {len(self.model.layers)} layers"
                )
            self._lens_selected_layers = layers
            self._lens_max_layer = max(layers)

        @torch.no_grad()
        def forward(  # type: ignore[override]

            self,

            input_ids: Optional[torch.LongTensor] = None,

            attention_mask: Optional[torch.Tensor] = None,

            *args,

            **kwargs,

        ):
            """Lens-specific forward.



            When ``input_ids`` and ``attention_mask`` are provided AND

            ``set_selected_layers(...)`` has been called, this returns the list of

            hidden states at the configured selected layers (the Lens feature

            extraction path).



            Otherwise, falls back to ``GptOssForCausalLM.forward`` so that

            ``generate(...)`` (used by the prompt reasoner) still works unchanged.

            """
            is_lens_feature_call = (
                input_ids is not None
                and attention_mask is not None
                and hasattr(self, "_lens_selected_layers")
                and not args
                and not kwargs
            )

            target_device = self.model.embed_tokens.weight.device
            if input_ids is not None and input_ids.device != target_device:
                input_ids = input_ids.to(target_device)
            if attention_mask is not None and attention_mask.device != target_device:
                attention_mask = attention_mask.to(target_device)

            if not is_lens_feature_call:
                return super().forward(input_ids, attention_mask, *args, **kwargs)

            model = self.model
            inputs_embeds = model.embed_tokens(input_ids)
            position_ids = torch.arange(
                inputs_embeds.shape[1], device=inputs_embeds.device
            ).unsqueeze(0).expand_as(input_ids)

            mask_kwargs = {
                "config": model.config,
                "inputs_embeds": inputs_embeds,
                "attention_mask": attention_mask,
                "past_key_values": None,
                "position_ids": position_ids,
            }
            causal_mask_mapping = {
                "full_attention": create_causal_mask(**mask_kwargs),
                "sliding_attention": create_sliding_window_causal_mask(**mask_kwargs),
            }

            hidden_states = inputs_embeds
            position_embeddings = model.rotary_emb(hidden_states, position_ids)

            captured: List[torch.Tensor] = [None] * len(self._lens_selected_layers)
            index_lookup = {idx: pos for pos, idx in enumerate(self._lens_selected_layers)}

            for i, decoder_layer in enumerate(model.layers):
                hidden_states = decoder_layer(
                    hidden_states,
                    attention_mask=causal_mask_mapping[model.config.layer_types[i]],
                    position_embeddings=position_embeddings,
                    position_ids=position_ids,
                    past_key_values=None,
                    use_cache=False,
                )
                if i in index_lookup:
                    captured[index_lookup[i]] = hidden_states
                if i == self._lens_max_layer:
                    break

            for pos, layer_idx in enumerate(self._lens_selected_layers):
                if captured[pos] is None:
                    raise RuntimeError(
                        f"Failed to capture hidden state for layer {layer_idx}"
                    )
            return captured

        def encode_layers(

            self,

            input_ids: torch.LongTensor,

            attention_mask: torch.Tensor,

        ) -> List[torch.Tensor]:
            """Backwards-compatible alias for the Lens feature path.



            Kept so existing call sites (``LensPipeline._get_text_embeddings``,

            external users) keep working. New code should call the encoder

            directly: ``encoder(input_ids, attention_mask)``.

            """
            if not hasattr(self, "_lens_selected_layers"):
                raise RuntimeError("Call set_selected_layers(...) before encode_layers().")
            return self(input_ids=input_ids, attention_mask=attention_mask)

else:

    class LensGptOssEncoder:  # type: ignore[no-redef]
        """Placeholder when transformers does not have GptOssForCausalLM.



        Lens requires ``transformers >= 5.8.0`` for the GPT-OSS model class.

        Please upgrade: ``pip install 'transformers>=5.8.0'``

        """

        def __init__(self, *args, **kwargs):
            raise ImportError(
                "LensGptOssEncoder requires GptOssForCausalLM from "
                "transformers >= 5.8.0. Please upgrade: "
                "pip install 'transformers>=5.8.0'"
            )

        @classmethod
        def from_pretrained(cls, *args, **kwargs):
            raise ImportError(
                "LensGptOssEncoder requires GptOssForCausalLM from "
                "transformers >= 5.8.0. Please upgrade: "
                "pip install 'transformers>=5.8.0'"
            )