From cbd01e20ba30b6ce6ae4b2fed31f82a402eba3e6 Mon Sep 17 00:00:00 2001 From: Ange Lou Date: Thu, 1 Oct 2026 06:59:51 +0000 Subject: [PATCH] Reduce direct-token latency with candidate-only projection Cache candidate token indices on the LM-head device and select the requested vocabulary rows before evaluation projection. Retain full-vocabulary training, temperature calibration and the existing device/dtype handling. Add FP32/BF16 candidate-readout parity, training-gradient and derived-buffer checks. The independent branch passes 431 CPU tests (37 optional skips), and builds both a source distribution and wheel. The PR description contains the 4B/27B accuracy and latency comparison with explicit deltas. --- jevany/model.py | 17 +++++++++++------ tests/test_backbones.py | 3 +++ tests/test_decision_modes.py | 25 ++++++++++++++++++++++++- 3 files changed, 38 insertions(+), 7 deletions(-) diff --git a/jevany/model.py b/jevany/model.py index a2bcde9..52d4610 100644 --- a/jevany/model.py +++ b/jevany/model.py @@ -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 @@ -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): diff --git a/tests/test_backbones.py b/tests/test_backbones.py index 5e25b28..7fac13f 100644 --- a/tests/test_backbones.py +++ b/tests/test_backbones.py @@ -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) diff --git a/tests/test_decision_modes.py b/tests/test_decision_modes.py index 52f5d0e..aa93a86 100644 --- a/tests/test_decision_modes.py +++ b/tests/test_decision_modes.py @@ -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 @@ -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(): @@ -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,