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'"
)
|