Align interleaved outcome targets and add open diagnostics

This commit is contained in:
2026-05-08 17:25:27 +09:00
parent b5550d3840
commit b33a55c76d
5 changed files with 277 additions and 1 deletions
@@ -459,6 +459,7 @@ def test_deep_cfr_interleaved_scheduler_matches_recursive_single_traversal() ->
"opponent_policy": "network",
"max_depth": 3,
"max_nodes_per_traversal": 64,
"outcome_unsampled_regret": "zero",
},
}
)
@@ -522,6 +523,7 @@ def test_deep_cfr_interleaved_scheduler_matches_recursive_single_traversal() ->
0,
1,
**common,
outcome_unsampled_regret=config.traversal.outcome_unsampled_regret,
opponent_policy=config.traversal.opponent_policy,
interleave_width=4,
interleave_max_batch=8,
@@ -536,6 +538,10 @@ def test_deep_cfr_interleaved_scheduler_matches_recursive_single_traversal() ->
[sample.target.sum() for sample in recursive_advantage],
atol=1.0e-5,
)
for interleaved_sample, recursive_sample in zip(
interleaved_advantage, recursive_advantage, strict=True
):
assert np.allclose(interleaved_sample.target, recursive_sample.target, atol=1.0e-5)
assert np.allclose(
[sample.target.sum() for sample in interleaved_strategy],
[sample.target.sum() for sample in recursive_strategy],
@@ -619,6 +625,7 @@ def test_deep_cfr_interleaved_scheduler_matches_average_strategy_opponent() -> N
0,
1,
**common,
outcome_unsampled_regret=config.traversal.outcome_unsampled_regret,
opponent_policy=config.traversal.opponent_policy,
interleave_width=4,
interleave_max_batch=8,