|
29 | 29 | from nemo_rl.environments.nemo_gym import ( |
30 | 30 | NemoGym, |
31 | 31 | NemoGymConfig, |
| 32 | + build_reward_component_columns, |
32 | 33 | extract_reward_components, |
33 | 34 | setup_nemo_gym_config, |
34 | 35 | validate_reward_components_match_scalar, |
@@ -62,6 +63,40 @@ def test_extract_reward_components(): |
62 | 63 | assert all(isinstance(v, float) for v in components.values()) |
63 | 64 |
|
64 | 65 |
|
| 66 | +def test_build_reward_component_columns(): |
| 67 | + """The bridge emission helper: reward/<name> keys, 0.0-fill, deterministic order. |
| 68 | +
|
| 69 | + Guards the producer->consumer contract end to end — the keys built here must be |
| 70 | + exactly what get_gdpo_reward_component_keys() selects (this is what would have caught |
| 71 | + the earlier reward1/reward2 vs reward/<name> mismatch). |
| 72 | + """ |
| 73 | + from nemo_rl.algorithms.utils import get_gdpo_reward_component_keys |
| 74 | + |
| 75 | + # Keys are reward/<name>; one entry per sample; values preserved. |
| 76 | + cols = build_reward_component_columns( |
| 77 | + [ |
| 78 | + {"correctness": 1.0, "format": 0.0}, |
| 79 | + {"correctness": 0.0, "format": 1.0}, |
| 80 | + ] |
| 81 | + ) |
| 82 | + assert set(cols) == {"reward/correctness", "reward/format"} |
| 83 | + assert torch.equal(cols["reward/correctness"], torch.tensor([1.0, 0.0])) |
| 84 | + assert torch.equal(cols["reward/format"], torch.tensor([0.0, 1.0])) |
| 85 | + |
| 86 | + # Union across the batch, deterministic (sorted) order, 0.0-fill for missing |
| 87 | + # components (and None samples). |
| 88 | + cols = build_reward_component_columns([{"b": 2.0}, {"a": 1.0, "b": 3.0}, None]) |
| 89 | + assert list(cols.keys()) == ["reward/a", "reward/b"] |
| 90 | + assert torch.equal(cols["reward/a"], torch.tensor([0.0, 1.0, 0.0])) |
| 91 | + assert torch.equal(cols["reward/b"], torch.tensor([2.0, 3.0, 0.0])) |
| 92 | + |
| 93 | + # The emitted keys are exactly what GDPO's consumer selects. |
| 94 | + assert get_gdpo_reward_component_keys(cols) == ["reward/a", "reward/b"] |
| 95 | + |
| 96 | + # No components anywhere -> no columns (single-reward path is untouched). |
| 97 | + assert build_reward_component_columns([None, None]) == {} |
| 98 | + |
| 99 | + |
65 | 100 | def test_validate_reward_components_match_scalar(): |
66 | 101 | """Multi-reward verifiers must set reward == sum(reward_components); mismatch raises.""" |
67 | 102 | # Contract satisfied: reward equals the component sum -> no error. |
|
0 commit comments