Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
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
17 changes: 11 additions & 6 deletions jevany/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -373,6 +373,10 @@ def hidden_to_device(module, args, output):
self.verbalizers = (list(verbalizers) if verbalizers is not None
else decision_verbalizers(tokenizer)) if decision_mode == "lm_token" else []
self.verbalizer_ids = verbalizer_token_ids(tokenizer, self.verbalizers) if self.verbalizers else []
index_device = self.lm_head.weight.device if self.lm_head is not None else device
self.register_buffer("_verbalizer_index",
torch.tensor(self.verbalizer_ids, dtype=torch.long, device=index_device),
persistent=False)
self.temperature = 1.0
self.device = device
self.device_map = device_map
Expand Down Expand Up @@ -479,12 +483,13 @@ def hidden_batch(self, encs):
def _question_readout(self, h, decide, options):
query = h[decide]
if self.decision_mode == "lm_token":
logits = F.linear(query.to(self.lm_head.weight.device, self.lm_head.weight.dtype), self.lm_head.weight).float()
if self.training:
return logits
candidates = torch.tensor(self.verbalizer_ids[:len(options)], device=logits.device)
logits = logits.index_select(0, candidates)
return logits if self.temperature == 1.0 else logits / self.temperature
weight = self.lm_head.weight
if not self.training:
# Inference normalizes over candidates only. Selecting rows
# before projection avoids scoring the rest of the vocabulary.
weight = weight.index_select(0, self._verbalizer_index[:len(options)])
logits = F.linear(query.to(weight.device, weight.dtype), weight).float()
return logits if self.training or self.temperature == 1.0 else logits / self.temperature
return self.head(query, h[torch.tensor(options, device=self.device)])

def _readout(self, h, enc):
Expand Down
3 changes: 3 additions & 0 deletions tests/test_backbones.py
Original file line number Diff line number Diff line change
Expand Up @@ -95,6 +95,9 @@ def test_text_causal_lm_direct_token_preserves_vocabulary_head(tmp_path):
tokenizer = load_tokenizer(base)
model = DecisionModel(base, tokenizer, "cpu", lora=2, decision_mode="lm_token",
verbalizers=["yes", "no"])
assert model._verbalizer_index.device == model.lm_head.weight.device
assert model._verbalizer_index.tolist() == model.verbalizer_ids
assert "_verbalizer_index" not in model.state_dict()
encoded = model.encode(tokenizer, RECORD)
model.train()
logits = model(encoded)
Expand Down
25 changes: 24 additions & 1 deletion tests/test_decision_modes.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@ def direct_model(vocabulary_size=8):
model.device = "cpu"
model.verbalizers = ["A", "B", "C"]
model.verbalizer_ids = [1, 4, 6]
model.register_buffer("_verbalizer_index", torch.tensor(model.verbalizer_ids), persistent=False)
model.temperature = 1.0
model.lm_head = torch.nn.Linear(3, vocabulary_size, bias=False)
model.option_isolation = False
Expand All @@ -28,9 +29,16 @@ def direct_model(vocabulary_size=8):

def test_lm_token_training_uses_full_vocabulary_logits():
model = direct_model()
model.temperature = 1.7
model.train()
logits = model._question_readout(torch.tensor([[1.0, 2.0, 3.0]]), 0, [0, 1])
hidden = torch.tensor([[1.0, 2.0, 3.0]], requires_grad=True)
logits = model._question_readout(hidden, 0, [0, 1])
assert logits.shape == (8,)
expected = torch.nn.functional.linear(hidden[0], model.lm_head.weight).float()
torch.testing.assert_close(logits, expected, rtol=0, atol=0)
actual_gradient, = torch.autograd.grad(logits.square().sum(), hidden, retain_graph=True)
expected_gradient, = torch.autograd.grad(expected.square().sum(), hidden)
torch.testing.assert_close(actual_gradient, expected_gradient, rtol=0, atol=0)


def test_lm_token_eval_selects_and_calibrates_candidates():
Expand All @@ -43,6 +51,21 @@ def test_lm_token_eval_selects_and_calibrates_candidates():
assert torch.equal(logits, torch.tensor([1.5, 6.0]))


@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16])
@pytest.mark.parametrize("options", [1, 2, 3])
def test_candidate_projection_matches_full_vocabulary_readout(dtype, options):
torch.manual_seed(17)
model = direct_model(vocabulary_size=1024)
model.to(dtype)
model.temperature = 1.7
model.eval()
hidden = torch.randn(4, 3, dtype=dtype)
full = torch.nn.functional.linear(hidden[2], model.lm_head.weight).float()
expected = full[model.verbalizer_ids[:options]] / model.temperature
actual = model._question_readout(hidden, 2, list(range(options)))
torch.testing.assert_close(actual, expected, rtol=1e-6, atol=1e-6)


def test_lm_token_ce_has_full_vocabulary_denominator():
logits = torch.tensor([0.0, 2.0, -1.0, 0.5, 1.0, -2.0, 3.0])
question = {"options": ["a", "b", "c"], "label": 2,
Expand Down
Loading