umer1995's picture
Fun CN: fp8 stream DiT + local bnb4 TE + xlarge (no bf16 host dump / no remote TE)
f0a4e91 verified
Raw
History Blame Contribute Delete
6.94 kB
# 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'"
)