Skip to content

easysteer.hidden_states

Hidden-state and MoE router-logit capture from a running vLLM engine.

easysteer.hidden_states.capture

capture(llm: Any, prompts: Any, max_tokens: int = 1, layers: Optional[List[int]] = None, dtype: Optional[str] = None, select: Optional[Any] = None, per_prompt_selects: Optional[List[Optional[Any]]] = None, stream: str = 'hidden_states', **generate_kwargs) -> CaptureResult

Capture intermediate state for a batch of prompts.

Parameters:

Name Type Description Default
llm Any

vLLM LLM instance (any engine config: compiled or eager, prefix caching on or off).

required
prompts Any

prompt list (text or multimodal dicts).

required
max_tokens int

tokens to generate (1 = prompt-only forward).

1
layers Optional[List[int]]

layer-id subset (None = all hooked layers).

None
dtype Optional[str]

engine-side storage dtype (e.g. 'float16').

None
select Optional[Any]

global SelectSpec (or wire dict) row selection.

None
per_prompt_selects Optional[List[Optional[Any]]]

one SelectSpec (or wire dict) per prompt, overriding the global selection for that prompt; None entries keep the global selection. Requires positions='all' semantics (no reductions).

None
stream str

'hidden_states' or 'router_logits'.

'hidden_states'
**generate_kwargs Any

forwarded into SamplingParams.

{}

Returns:

Type Description
CaptureResult

CaptureResult with exact per-sample views.

Source code in easysteer/hidden_states/capture_result.py
def capture(
    llm: Any,
    prompts: Any,
    max_tokens: int = 1,
    layers: Optional[List[int]] = None,
    dtype: Optional[str] = None,
    select: Optional[Any] = None,
    per_prompt_selects: Optional[List[Optional[Any]]] = None,
    stream: str = "hidden_states",
    **generate_kwargs,
) -> CaptureResult:
    """Capture intermediate state for a batch of prompts.

    Args:
        llm: vLLM LLM instance (any engine config: compiled or eager,
            prefix caching on or off).
        prompts: prompt list (text or multimodal dicts).
        max_tokens: tokens to generate (1 = prompt-only forward).
        layers: layer-id subset (None = all hooked layers).
        dtype: engine-side storage dtype (e.g. 'float16').
        select: global SelectSpec (or wire dict) row selection.
        per_prompt_selects: one SelectSpec (or wire dict) per prompt,
            overriding the global selection for that prompt; None
            entries keep the global selection. Requires positions='all'
            semantics (no reductions).
        stream: 'hidden_states' or 'router_logits'.
        **generate_kwargs (Any): forwarded into SamplingParams.

    Returns:
        CaptureResult with exact per-sample views.
    """
    from vllm import SamplingParams
    from vllm.capture import deserialize_captured

    def to_wire(spec):
        if spec is None or isinstance(spec, dict):
            return spec
        return spec.to_wire()

    def rpc(method, *args, **kwargs):
        results = llm.llm_engine.collective_rpc(method, args=args, kwargs=kwargs)
        if len(results) != 1:
            raise RuntimeError(
                f"capture expects a single worker, got {len(results)} "
                "RPC results — tensor-parallel capture would return "
                "per-rank shards and is not supported"
            )
        return results

    enable_kwargs: Dict[str, Any] = {}
    if layers is not None:
        enable_kwargs["layers"] = list(layers)
    if dtype is not None:
        enable_kwargs["dtype"] = dtype
    if select is not None:
        enable_kwargs["select"] = to_wire(select)

    capture_select = None
    if per_prompt_selects is not None:
        if len(per_prompt_selects) != len(prompts):
            raise ValueError(
                f"per_prompt_selects ({len(per_prompt_selects)}) must "
                f"match prompts ({len(prompts)})"
            )
        capture_select = [
            None if s is None else {stream: to_wire(s)}
            for s in per_prompt_selects
        ]

    sampling_params = SamplingParams(
        max_tokens=max_tokens,
        temperature=generate_kwargs.pop("temperature", 0.0),
        **generate_kwargs,
    )

    rpc("start_capture", stream, **enable_kwargs)
    try:
        outputs = llm.generate(
            _salt_prompts(prompts),
            sampling_params=sampling_params,
            capture_select=capture_select,
            use_tqdm=False,
        )
        raw = rpc("fetch_captured", stream, clear=True)[0]
    finally:
        rpc("stop_capture", stream)
    tensors, meta = deserialize_captured(raw)
    return CaptureResult(tensors, meta, outputs)

easysteer.hidden_states.CaptureResult

Result of one capture call.

Attributes:

Name Type Description
layers

{true_layer_id: Tensor(total_rows, dim)} in fetch order.

outputs

the vLLM RequestOutput list, prompt order.

Source code in easysteer/hidden_states/capture_result.py
class CaptureResult:
    """Result of one capture call.

    Attributes:
        layers: {true_layer_id: Tensor(total_rows, dim)} in fetch order.
        outputs: the vLLM RequestOutput list, prompt order.
    """

    def __init__(
        self,
        layers: Dict[int, torch.Tensor],
        meta: Optional[Dict[int, Any]],
        outputs: Any,
    ):
        self.layers = layers
        self.outputs = outputs
        self._meta = meta
        self._sample_rows: Optional[List[List[int]]] = None
        if meta is not None:
            for lid, m in meta.items():
                if len(m) != layers[lid].shape[0]:
                    raise RuntimeError(
                        f"layer {lid}: {len(m)} row labels for "
                        f"{layers[lid].shape[0]} rows — engine/client "
                        "label desync"
                    )
            self._sample_rows = self._index_samples()

    @property
    def layer_ids(self) -> List[int]:
        return sorted(self.layers)

    @property
    def labelled(self) -> bool:
        return self._sample_rows is not None

    def rows(self, layer: int) -> torch.Tensor:
        return self.layers[layer]

    def meta(self, layer: int):
        """Row labels (req_ids/positions/token_ids) for a layer."""
        if self._meta is None:
            raise RuntimeError("this capture has no row labels")
        return self._meta[layer]

    def _index_samples(self) -> List[List[int]]:
        from vllm.capture import match_capture_request_id

        first = self._meta[self.layer_ids[0]]
        by_label: Dict[str, List[int]] = {}
        for row, rid in enumerate(first.req_ids):
            by_label.setdefault(rid, []).append(row)
        sample_rows: List[List[int]] = []
        claimed = set()
        for output in self.outputs:
            matches = [
                label
                for label in by_label
                if label not in claimed
                and match_capture_request_id(label, output.request_id)
            ]
            if len(matches) > 1:
                raise RuntimeError(
                    f"request {output.request_id!r} matches several row "
                    f"label groups {matches!r}; duplicate client request "
                    "ids cannot be attributed"
                )
            if matches:
                claimed.add(matches[0])
                rows = by_label[matches[0]]
                rows.sort(key=lambda r: int(first.positions[r]))
                sample_rows.append(rows)
            else:
                sample_rows.append([])
        stale = set(by_label) - claimed
        if stale:
            raise RuntimeError(
                f"captured rows belong to requests outside this call: "
                f"{sorted(stale)[:5]} — the capture store was stale"
            )
        return sample_rows

    def __len__(self) -> int:
        return len(self.outputs)

    def sample(self, i: int) -> Dict[int, torch.Tensor]:
        """One sample's rows for every layer: {layer_id: (rows, dim)}."""
        if self._sample_rows is None:
            raise RuntimeError(
                "this capture has no row labels; per-sample views are "
                "unavailable"
            )
        idx = torch.tensor(self._sample_rows[i], dtype=torch.long)
        return {lid: t[idx] for lid, t in self.layers.items()}

    def sample_positions(self, i: int) -> List[int]:
        """Absolute sequence positions of sample i's rows (row order)."""
        first = self.meta(self.layer_ids[0])
        return [int(first.positions[r]) for r in self._sample_rows[i]]

    def sample_token_ids(self, i: int) -> List[int]:
        """Input token ids of sample i's rows (row order)."""
        first = self.meta(self.layer_ids[0])
        return [int(first.token_ids[r]) for r in self._sample_rows[i]]

    def to_nested(self) -> List[List[torch.Tensor]]:
        """Legacy extractor shape: `[sample][layer_pos]` (layers sorted by id)."""
        return [
            [self.sample(i)[lid] for lid in self.layer_ids]
            for i in range(len(self))
        ]

meta

meta(layer: int)

Row labels (req_ids/positions/token_ids) for a layer.

Source code in easysteer/hidden_states/capture_result.py
def meta(self, layer: int):
    """Row labels (req_ids/positions/token_ids) for a layer."""
    if self._meta is None:
        raise RuntimeError("this capture has no row labels")
    return self._meta[layer]

sample

sample(i: int) -> Dict[int, torch.Tensor]

One sample's rows for every layer: {layer_id: (rows, dim)}.

Source code in easysteer/hidden_states/capture_result.py
def sample(self, i: int) -> Dict[int, torch.Tensor]:
    """One sample's rows for every layer: {layer_id: (rows, dim)}."""
    if self._sample_rows is None:
        raise RuntimeError(
            "this capture has no row labels; per-sample views are "
            "unavailable"
        )
    idx = torch.tensor(self._sample_rows[i], dtype=torch.long)
    return {lid: t[idx] for lid, t in self.layers.items()}

sample_positions

sample_positions(i: int) -> List[int]

Absolute sequence positions of sample i's rows (row order).

Source code in easysteer/hidden_states/capture_result.py
def sample_positions(self, i: int) -> List[int]:
    """Absolute sequence positions of sample i's rows (row order)."""
    first = self.meta(self.layer_ids[0])
    return [int(first.positions[r]) for r in self._sample_rows[i]]

sample_token_ids

sample_token_ids(i: int) -> List[int]

Input token ids of sample i's rows (row order).

Source code in easysteer/hidden_states/capture_result.py
def sample_token_ids(self, i: int) -> List[int]:
    """Input token ids of sample i's rows (row order)."""
    first = self.meta(self.layer_ids[0])
    return [int(first.token_ids[r]) for r in self._sample_rows[i]]

to_nested

to_nested() -> List[List[torch.Tensor]]

Legacy extractor shape: [sample][layer_pos] (layers sorted by id).

Source code in easysteer/hidden_states/capture_result.py
def to_nested(self) -> List[List[torch.Tensor]]:
    """Legacy extractor shape: `[sample][layer_pos]` (layers sorted by id)."""
    return [
        [self.sample(i)[lid] for lid in self.layer_ids]
        for i in range(len(self))
    ]

Compatibility wrappers

Nested-list wrappers over capture(); prefer capture() for new code.

easysteer.hidden_states.get_all_hidden_states_generate

get_all_hidden_states_generate(llm: Any, prompts: Union[List[str], List[Dict[str, Any]]], max_tokens: int = 1, split_by_samples: bool = True, token_ids: Optional[List[int]] = None, positions: Optional[List[int]] = None, layers: Optional[List[int]] = None, dtype: Optional[str] = None, select: Optional[Union[dict, Any]] = None, **generate_kwargs) -> Union[Tuple[List[List[torch.Tensor]], Any], Tuple[List[torch.Tensor], Any]]

Capture every layer's hidden states while running generate.

Works for any generate-capable model, including multimodal models (Qwen-VL, LLaVA, ...) that do not support the embed task. With the default max_tokens=1 only the prompt forward is captured, matching what an embed task would produce.

Parameters:

Name Type Description Default
llm Any

vLLM LLM instance (any engine config: compiled or eager, prefix caching on or off).

required
prompts Union[List[str], List[Dict[str, Any]]]

text prompts, or multimodal dicts with prompt and multi_modal_data keys.

required
max_tokens int

tokens to generate (1 = prompt-only forward).

1
split_by_samples bool

if True return [sample][layer] tensors; if False return per-layer tensors concatenated over samples.

True
token_ids Optional[List[int]]

only capture rows whose input token id is in this list (source-side filter; unions with positions).

None
positions Optional[List[int]]

only capture these absolute positions (negatives resolve from the prompt end; unions with token_ids).

None
layers Optional[List[int]]

layer-id subset (None = all hooked layers).

None
dtype Optional[str]

engine-side storage dtype (e.g. 'float16').

None
select Optional[Union[dict, Any]]

SelectSpec (or wire dict) — the full where-clause selection language; cannot combine with the shortcuts.

None
**generate_kwargs Any

forwarded into SamplingParams.

{}

Returns:

Type Description
Union[Tuple[List[List[Tensor]], Any], Tuple[List[Tensor], Any]]

(hidden_states, outputs) where hidden_states is

Union[Tuple[List[List[Tensor]], Any], Tuple[List[Tensor], Any]]

[sample][layer] (split) or [layer] (concatenated),

Union[Tuple[List[List[Tensor]], Any], Tuple[List[Tensor], Any]]

layers ordered by layer id.

Source code in easysteer/hidden_states/capture_generate.py
def get_all_hidden_states_generate(
    llm: Any,
    prompts: Union[List[str], List[Dict[str, Any]]],
    max_tokens: int = 1,
    split_by_samples: bool = True,
    token_ids: Optional[List[int]] = None,
    positions: Optional[List[int]] = None,
    layers: Optional[List[int]] = None,
    dtype: Optional[str] = None,
    select: Optional[Union[dict, Any]] = None,
    **generate_kwargs,
) -> Union[
    Tuple[List[List[torch.Tensor]], Any], Tuple[List[torch.Tensor], Any]
]:
    """Capture every layer's hidden states while running generate.

    Works for any generate-capable model, including multimodal models
    (Qwen-VL, LLaVA, ...) that do not support the embed task. With the
    default ``max_tokens=1`` only the prompt forward is captured,
    matching what an embed task would produce.

    Args:
        llm: vLLM LLM instance (any engine config: compiled or eager,
            prefix caching on or off).
        prompts: text prompts, or multimodal dicts with ``prompt`` and
            ``multi_modal_data`` keys.
        max_tokens: tokens to generate (1 = prompt-only forward).
        split_by_samples: if True return ``[sample][layer]`` tensors;
            if False return per-layer tensors concatenated over samples.
        token_ids: only capture rows whose input token id is in this
            list (source-side filter; unions with ``positions``).
        positions: only capture these absolute positions (negatives
            resolve from the prompt end; unions with ``token_ids``).
        layers: layer-id subset (None = all hooked layers).
        dtype: engine-side storage dtype (e.g. ``'float16'``).
        select: SelectSpec (or wire dict) — the full where-clause
            selection language; cannot combine with the shortcuts.
        **generate_kwargs (Any): forwarded into SamplingParams.

    Returns:
        ``(hidden_states, outputs)`` where hidden_states is
        ``[sample][layer]`` (split) or ``[layer]`` (concatenated),
        layers ordered by layer id.
    """
    result = capture(
        llm,
        prompts,
        max_tokens=max_tokens,
        layers=layers,
        dtype=dtype,
        select=_sugar_select(token_ids, positions, select),
        **generate_kwargs,
    )
    if split_by_samples:
        return result.to_nested(), result.outputs
    return [result.layers[lid] for lid in result.layer_ids], result.outputs

easysteer.hidden_states.get_moe_router_logits_generate

get_moe_router_logits_generate(llm: Any, prompts: Union[List[str], List[Dict[str, Any]]], max_tokens: int = 1, split_by_samples: bool = False, **generate_kwargs) -> Union[Tuple[Dict[int, torch.Tensor], Any], Tuple[List[Dict[int, torch.Tensor]], Any]]

Capture MoE router logits while running generate.

Works for any generate-capable MoE model, including multimodal ones (e.g. Qwen3-VL). With the default max_tokens=1 only the prompt forward is captured. When router-logits steering is active, the captured logits are the post-steering ones.

Parameters:

Name Type Description Default
llm Any

vLLM LLM instance (any engine config: compiled or eager, prefix caching on or off).

required
prompts Union[List[str], List[Dict[str, Any]]]

text prompts, or multimodal dicts with prompt and multi_modal_data keys.

required
max_tokens int

tokens to generate (1 = prompt-only forward).

1
split_by_samples bool

if True return one {layer_id: tensor} dict per sample; if False return a single dict with all samples' rows concatenated per layer.

False
**generate_kwargs Any

forwarded into SamplingParams.

{}

Returns:

Type Description
Union[Tuple[Dict[int, Tensor], Any], Tuple[List[Dict[int, Tensor]], Any]]

(router_logits, outputs) where router_logits is

Union[Tuple[Dict[int, Tensor], Any], Tuple[List[Dict[int, Tensor]], Any]]

{layer_id: (rows, n_experts)} (concatenated) or

Union[Tuple[Dict[int, Tensor], Any], Tuple[List[Dict[int, Tensor]], Any]]

[sample_idx]{layer_id: (rows, n_experts)} (split).

Source code in easysteer/hidden_states/moe_capture_generate.py
def get_moe_router_logits_generate(
    llm: Any,
    prompts: Union[List[str], List[Dict[str, Any]]],
    max_tokens: int = 1,
    split_by_samples: bool = False,
    **generate_kwargs,
) -> Union[
    Tuple[Dict[int, torch.Tensor], Any],
    Tuple[List[Dict[int, torch.Tensor]], Any],
]:
    """Capture MoE router logits while running generate.

    Works for any generate-capable MoE model, including multimodal
    ones (e.g. Qwen3-VL). With the default ``max_tokens=1`` only the
    prompt forward is captured. When router-logits steering is active,
    the captured logits are the post-steering ones.

    Args:
        llm: vLLM LLM instance (any engine config: compiled or eager,
            prefix caching on or off).
        prompts: text prompts, or multimodal dicts with ``prompt`` and
            ``multi_modal_data`` keys.
        max_tokens: tokens to generate (1 = prompt-only forward).
        split_by_samples: if True return one ``{layer_id: tensor}``
            dict per sample; if False return a single dict with all
            samples' rows concatenated per layer.
        **generate_kwargs (Any): forwarded into SamplingParams.

    Returns:
        ``(router_logits, outputs)`` where router_logits is
        ``{layer_id: (rows, n_experts)}`` (concatenated) or
        ``[sample_idx]{layer_id: (rows, n_experts)}`` (split).
    """
    result = capture(
        llm,
        prompts,
        max_tokens=max_tokens,
        stream="router_logits",
        **generate_kwargs,
    )
    if split_by_samples:
        return [result.sample(i) for i in range(len(result))], result.outputs
    return dict(result.layers), result.outputs