From f84eebdae6c5726b4bc1381ff912e95083971d53 Mon Sep 17 00:00:00 2001 From: HabibaTahir Date: Wed, 19 Aug 2026 10:12:44 +0500 Subject: [PATCH] Fix inference crash: undefined true_rule_ids in transformer.py forward(), beam_search.py tuple-unpack + src_positions kwarg mismatch --- inference/beam_search.py | 19 +++++++++++++++---- model/transformer.py | 4 ++-- 2 files changed, 17 insertions(+), 6 deletions(-) diff --git a/inference/beam_search.py b/inference/beam_search.py index 16a0c11..8c221dd 100644 --- a/inference/beam_search.py +++ b/inference/beam_search.py @@ -19,9 +19,20 @@ def beam_search( beam_size: int = 5, max_len: int = 32, node_pool: Optional[NodeValidityPool] = None, + src_positions: Optional[torch.Tensor] = None, + parent_child_pairs: Optional[torch.Tensor] = None, ) -> Dict[str, Any]: - """Simplified beam search for SimpleCalculusModel -- one model call - per step (src_seq, tgt_in_seq), no rule_embeddings, no tree kwargs.""" + """Beam search for the tree-based CalculusSolverModel (model/transformer.py). + + NOTE: CalculusSolverModel.forward(src_seq, tgt_in_seq, true_rule_ids=None) + computes src_positions/parent_child_pairs internally (as zero tensors) and + does not take them as inputs. src_positions/parent_child_pairs are accepted + here only so callers built for the older tree-kwarg interface (e.g. + inference/solve.py) don't break -- they are unused. + + forward() returns (decoder_logits, rule_logits, verifier_logits); only + decoder_logits is used for next-token scoring here. + """ device = src_tokens.device vocab = vocab_map["token_to_id"] id_to_token = vocab_map["id_to_token"] @@ -53,8 +64,8 @@ def beam_search( ) tgt = torch.tensor([current_tokens], device=device) - logits = model(src_tokens, tgt) - next_logits = logits[0, -1, :] + decoder_logits, _rule_logits, _verifier_logits = model(src_tokens, tgt) + next_logits = decoder_logits[0, -1, :] mask = node_pool.mask(validity_tokens, all_candidate_tokens) invalid_mask = torch.tensor([not v for v in mask], device=device) diff --git a/model/transformer.py b/model/transformer.py index 8497a45..2749e2a 100644 --- a/model/transformer.py +++ b/model/transformer.py @@ -58,7 +58,7 @@ def __init__( templates=templates ) - def forward(self, src_seq, tgt_in_seq): + def forward(self, src_seq, tgt_in_seq, true_rule_ids=None): device = src_seq.device batch_size, seq_len = src_seq.size() @@ -95,4 +95,4 @@ def forward(self, src_seq, tgt_in_seq): # 5. Trace steps (verifier) verifier_logits = self.step_tracer(rule_ids, decoder_hidden_states) - return decoder_logits, rule_logits, verifier_logits + return decoder_logits, rule_logits, verifier_logits \ No newline at end of file