Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
188 changes: 188 additions & 0 deletions tests/integration/model_bridge/test_inspect_attention_mask.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,188 @@
"""Offline mask parity through the real Inspect driver and HF provider."""

import pytest
import torch
from tokenizers import Tokenizer
from tokenizers.models import WordLevel
from transformers import (
GPT2Config,
GPT2LMHeadModel,
LlamaConfig,
LlamaForCausalLM,
OPTConfig,
OPTForCausalLM,
PreTrainedTokenizerFast,
)

pytest.importorskip("inspect_ai")
pytestmark = pytest.mark.inspect


@pytest.fixture(scope="module", params=["gpt2", "llama", "opt"])
def inspect_models(tmp_path_factory, request):
from transformer_lens.model_bridge.remote_bridge import RemoteBridge

path = tmp_path_factory.mktemp("inspect_mask_model")
torch.manual_seed(2)
if request.param == "gpt2":
model = GPT2LMHeadModel(
GPT2Config(
vocab_size=32,
n_embd=16,
n_layer=1,
n_head=2,
n_positions=32,
bos_token_id=1,
eos_token_id=2,
pad_token_id=0,
attn_implementation="eager",
)
)
elif request.param == "llama":
model = LlamaForCausalLM(
LlamaConfig(
vocab_size=32,
hidden_size=16,
intermediate_size=32,
num_hidden_layers=1,
num_attention_heads=2,
num_key_value_heads=2,
max_position_embeddings=32,
bos_token_id=1,
eos_token_id=2,
pad_token_id=0,
attn_implementation="eager",
)
)
else:
model = OPTForCausalLM(
OPTConfig(
vocab_size=32,
hidden_size=16,
ffn_dim=32,
num_hidden_layers=1,
num_attention_heads=2,
word_embed_proj_dim=16,
max_position_embeddings=32,
bos_token_id=1,
eos_token_id=2,
pad_token_id=0,
attn_implementation="eager",
)
)
model.eval()
model.save_pretrained(path)
tokenizer = PreTrainedTokenizerFast(
tokenizer_object=Tokenizer(WordLevel({str(i): i for i in range(32)}, unk_token="3")),
pad_token="0",
bos_token="1",
eos_token="2",
unk_token="3",
)
tokenizer.save_pretrained(path)
bridge = RemoteBridge.boot_inspect(str(path), device="cpu")
yield bridge, model
bridge.close()


@pytest.mark.parametrize("left_padding", [True, False])
def test_masked_logits_cache_and_loss(inspect_models, left_padding):
bridge, model = inspect_models
tokens = torch.tensor([[0, 0, 7, 8, 9] if left_padding else [7, 8, 9, 0, 0]])
mask = torch.tensor([[0, 0, 1, 1, 1] if left_padding else [1, 1, 1, 0, 0]])
kwargs = {"attention_mask": mask}
if left_padding and model.config.model_type != "opt":
kwargs["position_ids"] = (mask.cumsum(-1) - 1).clamp_min(0)
with torch.no_grad():
expected = model(tokens, **kwargs).logits
actual, cache = bridge.run_with_cache(
tokens, attention_mask=mask, names_filter=["blocks.0.hook_out"]
)
torch.testing.assert_close(actual, expected, rtol=1e-5, atol=1e-6)
alone, alone_cache = bridge.run_with_cache(
torch.tensor([[7, 8, 9]]), names_filter=["blocks.0.hook_out"]
)
valid = mask[0].bool()
torch.testing.assert_close(actual[:, valid], alone, rtol=1e-5, atol=1e-6)
torch.testing.assert_close(
cache["blocks.0.hook_out"][:, valid], alone_cache["blocks.0.hook_out"], rtol=1e-5, atol=1e-6
)
loss = bridge.forward(tokens, attention_mask=mask, return_type="loss")
alone_loss = bridge.forward(torch.tensor([[7, 8, 9]]), return_type="loss")
torch.testing.assert_close(loss, alone_loss)
labeled_loss = bridge.forward(tokens, labels=tokens, attention_mask=mask, return_type="loss")
torch.testing.assert_close(labeled_loss, alone_loss)


def test_all_ones_mask_preserves_unmasked_forward(inspect_models):
bridge, _ = inspect_models
tokens = torch.tensor([[7, 8, 9]])
torch.testing.assert_close(
bridge.forward(tokens, attention_mask=torch.ones_like(tokens)),
bridge.forward(tokens),
rtol=0,
atol=0,
)


@pytest.mark.parametrize("mask", [[[1, 1]], [[1, 0.5, 1]], [[0, 0, 0]], [[[1, 1, 1]]]])
def test_invalid_masks_fail_at_public_and_provider_boundaries(inspect_models, mask):
from inspect_ai.model import GenerateConfig

bridge, _ = inspect_models
with pytest.raises(ValueError, match="attention_mask"):
bridge.forward(torch.tensor([[7, 8, 9]]), attention_mask=torch.tensor(mask))
with pytest.raises(ValueError, match="attention_mask"):
bridge._driver._model.api._generate_capture(
[], {"input_ids": [7, 8, 9], "attention_mask": mask}, GenerateConfig()
)


@pytest.mark.parametrize("dtype", [torch.bool, torch.float32, torch.bfloat16])
def test_binary_mask_dtypes(inspect_models, dtype):
bridge, _ = inspect_models
tokens = torch.tensor([[0, 7, 8]])
mask = torch.tensor([[0, 1, 1]])
torch.testing.assert_close(
bridge.forward(tokens, attention_mask=mask.to(dtype)),
bridge.forward(tokens, attention_mask=mask),
rtol=0,
atol=0,
)


def test_masked_capture_and_intervention(inspect_models):
bridge, _ = inspect_models
tokens = torch.tensor([[0, 0, 7, 8, 9]])
mask = torch.tensor([[0, 0, 1, 1, 1]])
names = ["blocks.0.hook_out", "blocks.0.attn.hook_pattern"]
intervention = {"blocks.0.attn.hook_out": {"op": "scale", "factor": 0.5}}
padded, cache = bridge.run_with_cache(
tokens, attention_mask=mask, names_filter=names, intervene=intervention
)
alone, alone_cache = bridge.run_with_cache(
torch.tensor([[7, 8, 9]]), names_filter=names, intervene=intervention
)
torch.testing.assert_close(padded[:, 2:], alone, rtol=1e-5, atol=1e-6)
torch.testing.assert_close(cache[names[0]][:, 2:], alone_cache[names[0]], rtol=1e-5, atol=1e-6)
pattern = cache[names[1]]
torch.testing.assert_close(pattern[:, :, 2:, 2:], alone_cache[names[1]], rtol=1e-5, atol=1e-6)
assert torch.count_nonzero(pattern[:, :, 2:, :2]) == 0


def test_provider_completion_uses_last_attended_token(inspect_models):
from inspect_ai.model import GenerateConfig

bridge, model = inspect_models
api = bridge._driver._model.api
with torch.no_grad():
logits = model(torch.tensor([[7, 8, 9]])).logits[0, -1]
next_id = int(logits.argmax())
output = api._generate_capture(
[],
{"input_ids": [7, 8, 9, 0, 0], "attention_mask": [1, 1, 1, 0, 0]},
GenerateConfig(logprobs=True, top_logprobs=3),
)
assert output.completion == api._tokenizer.decode([next_id])
entry = output.choices[0].logprobs.content[0]
assert entry.logprob == pytest.approx(float(logits.log_softmax(-1)[next_id]), abs=1e-6)
22 changes: 22 additions & 0 deletions tests/unit/model_bridge/test_inspect_driver.py
Original file line number Diff line number Diff line change
Expand Up @@ -80,6 +80,28 @@ def _driver(model=None) -> InspectDriver:
return InspectDriver(model=model or _fake_model(), adapter=_adapter(), tokenizer=None)


@pytest.mark.parametrize("provider", ["tl_bridge_vllm", "vllm-lens"])
def test_unsupported_provider_rejects_attention_mask(provider):
driver = InspectDriver(_fake_model(), _adapter(), None, profiles.for_provider(provider))
try:
with pytest.raises(NotImplementedError, match="attention_mask"):
driver.forward(np.array([[7, 8]]), attention_mask=np.array([[1, 1]]))
finally:
driver.close()


def test_vllm_provider_direct_capture_rejects_attention_mask():
from inspect_ai.model import GenerateConfig

from transformer_lens.model_bridge.sources.inspect.vllm_provider import (
TransformerLensVLLMModelAPI,
)

api = object.__new__(TransformerLensVLLMModelAPI)
with pytest.raises(NotImplementedError, match="attention_mask"):
api._generate_capture([], {"input_ids": [7, 8], "attention_mask": [1, 1]}, GenerateConfig())


class TestProtocolConformance:
def test_is_driver_and_validates(self):
driver = _driver()
Expand Down
9 changes: 9 additions & 0 deletions transformer_lens/model_bridge/sources/inspect/driver.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,7 @@
from transformer_lens.model_bridge.sources._driver_base import DriverBase

from . import hooks, wire
from .masks import normalize_attention_mask
from .profiles import TLBridgeProfile

# Cap a single provider call so a hung remote/provider forward unblocks the sync caller
Expand Down Expand Up @@ -68,6 +69,7 @@ def forward(
intervene: Mapping[str, Intervention] | None = None,
max_new_tokens: int = 1,
return_logits: bool = True,
attention_mask: TensorLike | None = None,
**kwargs: Any,
) -> ForwardResult:
if self._model is None:
Expand All @@ -79,6 +81,11 @@ def forward(
"InspectDriver supports max_new_tokens=1 only (single-forward capture)."
)
ids = self._normalize_input_ids(input_ids)
mask = None
if attention_mask is not None:
if not getattr(self._profile, "supports_attention_mask", False):
raise NotImplementedError("This Inspect provider does not support attention_mask.")
mask = normalize_attention_mask(attention_mask, len(ids))
# capture is authoritative: the bridge passes exactly the hooks with handlers,
# so () means "capture nothing" (logits only), not "capture everything".
names = list(capture)
Expand All @@ -89,6 +96,8 @@ def forward(
prompt, extra_args = self._profile.build_request(
ids, wire_keys, interventions, return_logits, self.tokenizer
)
if mask is not None:
extra_args["attention_mask"] = mask

output = self._run_coro(self._generate(prompt, extra_args))

Expand Down
21 changes: 21 additions & 0 deletions transformer_lens/model_bridge/sources/inspect/masks.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,21 @@
"""Torch-free validation of the single-sequence Inspect padding-mask wire format."""

from typing import Any

import numpy as np


def normalize_attention_mask(mask: Any, n_tokens: int) -> list[int]:
"""Accept a binary flat wire mask or a single-row public padding mask."""
if hasattr(mask, "detach"):
mask = mask.detach().cpu().tolist()
array = np.asarray(mask)
if array.ndim == 2 and array.shape[0] == 1:
array = array[0]
if array.ndim != 1 or array.shape[0] != n_tokens:
raise ValueError("Inspect attention_mask must match the single input sequence length.")
if not np.all((array == 0) | (array == 1)):
raise ValueError("Inspect attention_mask must be a binary 0/1 padding mask.")
if not np.any(array):
raise ValueError("Inspect attention_mask must retain at least one token.")
return [int(value) for value in array.tolist()]
11 changes: 9 additions & 2 deletions transformer_lens/model_bridge/sources/inspect/profiles.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,9 +32,15 @@ class TLBridgeProfile:
# only populate the gen position; lets RemoteBridge.forward reject loss/both there.
provides_sequence_logits = True

def __init__(self, supported_kinds: Any = None, provides_sequence_logits: bool = True) -> None:
def __init__(
self,
supported_kinds: Any = None,
provides_sequence_logits: bool = True,
supports_attention_mask: bool = True,
) -> None:
self._kinds = supported_kinds
self.provides_sequence_logits = provides_sequence_logits
self.supports_attention_mask = supports_attention_mask

def supported_hooks(self, n_layers: int) -> frozenset[str]:
return hooks.supported_hook_points(n_layers, self._kinds)
Expand Down Expand Up @@ -79,6 +85,7 @@ class VLLMLensProfile:
"""

provides_sequence_logits = False
supports_attention_mask = False

def supported_hooks(self, n_layers: int) -> frozenset[str]:
return frozenset(f"blocks.{i}.hook_out" for i in range(n_layers))
Expand Down Expand Up @@ -142,7 +149,7 @@ def for_provider(provider: str) -> Any:
if provider.startswith("vllm-lens"):
return VLLMLensProfile()
if provider in ("tl_bridge", "tl_bridge_vllm"):
return TLBridgeProfile()
return TLBridgeProfile(supports_attention_mask=provider == "tl_bridge")
# An unknown provider would otherwise get full-capability codec and NaN downstream.
raise ValueError(
f"No Inspect codec for provider {provider!r}. Known providers: 'tl_bridge', "
Expand Down
11 changes: 10 additions & 1 deletion transformer_lens/model_bridge/sources/inspect/source.py
Original file line number Diff line number Diff line change
Expand Up @@ -61,6 +61,11 @@ def boot_inspect(
default) and eager attention. Full-sequence logits ride on ``return_logits=True`` (the
default); pass ``return_logits=False`` to skip the (seq × d_vocab) payload for pure
activation capture (``run_with_cache`` keeps them since it returns logits).

``tl_bridge`` accepts a binary, single-row ``attention_mask`` matching the input
length. Masked forwards use padding-aware positions where the model accepts 2-D
position IDs; model-owned position derivation is left intact. The vLLM Inspect
providers reject masks rather than silently ignore them.
"""
from inspect_ai.model import get_model
from transformers import AutoConfig, AutoTokenizer
Expand Down Expand Up @@ -118,7 +123,11 @@ def boot_inspect(
kinds = api.supported_kinds() if hasattr(api, "supported_kinds") else None
note = api.capability_note() if hasattr(api, "capability_note") else ""
psl = bool(getattr(api, "provides_sequence_logits", True))
profile = profiles.TLBridgeProfile(supported_kinds=kinds, provides_sequence_logits=psl)
profile = profiles.TLBridgeProfile(
supported_kinds=kinds,
provides_sequence_logits=psl,
supports_attention_mask=provider == "tl_bridge",
Comment thread
jlarson4 marked this conversation as resolved.
Outdated
)
if note:
warnings.warn(note, UserWarning, stacklevel=2)
else:
Expand Down
Loading
Loading