@@ -216,25 +216,27 @@ class TestComputeRowReward(unittest.TestCase):
216216 def _two_fns (self ):
217217 """Return two reward functions whose per-response scores can be summed."""
218218
219- # Each fn must accept prompts, completions, answer as keyword args and
220- # return a list of per-completion scores. The helper calls fn once per
221- # response with single-element lists, so the returned list has length 1.
222- def fn1 (prompts , completions , answer ): # pylint: disable=unused-argument
219+ # Each fn must accept prompts, completions, answer, question as keyword
220+ # args and return a list of per-completion scores. The helper calls fn
221+ # once per response with single-element lists, so the returned list has
222+ # length 1.
223+ def fn1 (prompts , completions , answer , question ): # pylint: disable=unused-argument
223224 return [1.0 for _ in completions ]
224225
225- def fn2 (prompts , completions , answer ): # pylint: disable=unused-argument
226+ def fn2 (prompts , completions , answer , question ): # pylint: disable=unused-argument
226227 return [float (len (c )) for c in completions ]
227228
228229 return [fn1 , fn2 ]
229230
230231 @pytest .mark .cpu_only
231232 def test_single_response_single_fn (self ):
232- def fn (prompts , completions , answer ): # pylint: disable=unused-argument
233+ def fn (prompts , completions , answer , question ): # pylint: disable=unused-argument
233234 return [2.5 for _ in completions ]
234235
235236 score_sum , count = evaluate_rl ._compute_row_reward (
236237 reward_fns = [fn ],
237238 prompt = "p" ,
239+ question = "q" ,
238240 responses = ["abc" ],
239241 answer = "gold" ,
240242 row_idx = 0 ,
@@ -247,6 +249,7 @@ def test_sums_across_reward_fns_for_single_response(self):
247249 score_sum , count = evaluate_rl ._compute_row_reward (
248250 reward_fns = self ._two_fns (),
249251 prompt = "p" ,
252+ question = "q" ,
250253 responses = ["abcd" ],
251254 answer = "gold" ,
252255 row_idx = 0 ,
@@ -261,6 +264,7 @@ def test_sums_across_passes_for_multiple_responses(self):
261264 score_sum , count = evaluate_rl ._compute_row_reward (
262265 reward_fns = self ._two_fns (),
263266 prompt = "p" ,
267+ question = "q" ,
264268 responses = ["a" , "bcd" , "ef" ],
265269 answer = "gold" ,
266270 row_idx = 0 ,
@@ -276,6 +280,7 @@ def test_empty_responses_returns_zero_and_zero_count(self):
276280 score_sum , count = evaluate_rl ._compute_row_reward (
277281 reward_fns = self ._two_fns (),
278282 prompt = "p" ,
283+ question = "q" ,
279284 responses = [],
280285 answer = "gold" ,
281286 row_idx = 0 ,
@@ -293,13 +298,65 @@ def _boom(**kwargs): # pylint: disable=unused-argument
293298 score_sum , count = evaluate_rl ._compute_row_reward (
294299 reward_fns = [_boom ],
295300 prompt = "p" ,
301+ question = "q" ,
296302 responses = ["abc" ],
297303 answer = "gold" ,
298304 row_idx = 0 ,
299305 )
300306 self .assertEqual (score_sum , 0.0 )
301307 self .assertEqual (count , 0 ) # zero count so the caller's mean isn't biased
302308
309+ @pytest .mark .cpu_only
310+ def test_question_is_forwarded_to_reward_fn (self ):
311+ """Regression: helper must pass `question` through to reward fns.
312+
313+ The built-in `check_numbers` reward reads `kwargs["question"]`; if the
314+ helper omits it, every eval row produces `KeyError('question')` and
315+ `mean_reward` collapses to 0.0.
316+ """
317+ received = {}
318+
319+ def fn (prompts , completions , answer , question ): # pylint: disable=unused-argument
320+ received ["question" ] = question
321+ return [1.0 for _ in completions ]
322+
323+ evaluate_rl ._compute_row_reward (
324+ reward_fns = [fn ],
325+ prompt = "p" ,
326+ question = "What is 2+2?" ,
327+ responses = ["abc" ],
328+ answer = "gold" ,
329+ row_idx = 0 ,
330+ )
331+ self .assertEqual (received ["question" ], "What is 2+2?" )
332+
333+ @pytest .mark .cpu_only
334+ def test_integrates_with_real_check_numbers_reward (self ):
335+ """End-to-end: real `check_numbers` from utils_rl must not raise on the
336+ eval-time kwargs the helper passes (regression for the original
337+ `KeyError('question')` failure mode in production)."""
338+ from maxtext .trainers .post_train .rl import utils_rl # pylint: disable=import-outside-toplevel
339+
340+ config = _make_config (eval_mode = "pass" )
341+
342+ # `check_numbers` takes tmvp_config positionally via partial; mirror what
343+ # train_rl.py's make_reward_fn does.
344+ def wrapped_check_numbers (** kwargs ):
345+ return utils_rl .check_numbers (tmvp_config = config , ** kwargs )
346+
347+ # A correct-answer response should yield a non-zero score (proves the
348+ # kwargs all reached the inside of check_numbers).
349+ score_sum , count = evaluate_rl ._compute_row_reward (
350+ reward_fns = [wrapped_check_numbers ],
351+ prompt = "solve: 2+2" ,
352+ question = "What is 2+2?" ,
353+ responses = ["<reasoning>2+2=4</reasoning><answer>4</answer>" ],
354+ answer = '["4"]' , # json-encoded list of acceptable answers
355+ row_idx = 0 ,
356+ )
357+ self .assertEqual (count , 1 )
358+ self .assertGreater (score_sum , 0.0 )
359+
303360
304361if __name__ == "__main__" :
305362 unittest .main ()
0 commit comments