| """Internal drafting experiments: cross-examination and counterfactual rehearsal.""" |
|
|
| from __future__ import annotations |
|
|
| import torch |
|
|
| from nexum_core.internal_agents import merge_selected_agent_workspaces |
|
|
|
|
| def _workspace( |
| knowledge: list[list[float]], |
| capability: list[list[float]], |
| ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: |
| selected_t = torch.tensor( |
| [ |
| [ |
| [1.0, 0.0, 0.0], |
| [0.9, 0.1, 0.0], |
| [-1.0, 0.0, 0.0], |
| [0.0, 1.0, 0.0], |
| ] |
| ] |
| ) |
| route_ids_t = torch.arange(4).reshape(1, 4) |
| route_weight_t = torch.tensor([[0.45, 0.35, 0.10, 0.10]]) |
| knowledge_t = torch.tensor([knowledge], dtype=torch.float32) |
| capability_t = torch.tensor([capability], dtype=torch.float32) |
| return selected_t, route_ids_t, route_weight_t, knowledge_t, capability_t |
|
|
|
|
| def test_counterfactual_rehearsal_moves_merge_toward_supported_dissenter() -> None: |
| selected_t, route_ids_t, route_weight_t, knowledge_t, capability_t = _workspace( |
| [[0.1], [0.1], [3.0], [0.2]], |
| [[0.1], [0.1], [3.0], [0.2]], |
| ) |
|
|
| first = merge_selected_agent_workspaces( |
| selected_t, |
| route_ids_t, |
| route_weight_t, |
| knowledge_t=knowledge_t, |
| capability_t=capability_t, |
| correction_pressure_t=torch.zeros(()), |
| release_confidence_t=torch.ones(()), |
| ) |
| corrected = merge_selected_agent_workspaces( |
| selected_t, |
| route_ids_t, |
| route_weight_t, |
| knowledge_t=knowledge_t, |
| capability_t=capability_t, |
| correction_pressure_t=torch.ones(()), |
| release_confidence_t=torch.ones(()), |
| ) |
|
|
| assert torch.all(first.acceptance_t > 0) |
| assert torch.all(corrected.acceptance_t > 0) |
| torch.testing.assert_close( |
| corrected.acceptance_t.sum(dim=-1), torch.ones(1) |
| ) |
| assert corrected.acceptance_t[0, 2] > first.acceptance_t[0, 2] |
| assert corrected.merged_t[0, 0] < first.merged_t[0, 0] |
| assert not torch.allclose(corrected.merged_t, first.merged_t) |
| assert torch.count_nonzero(corrected.rehearsal_delta_t) > 0 |
|
|
|
|
| def test_disagreement_preserves_diversity_until_challenged() -> None: |
| route_ids_t = torch.arange(3).reshape(1, 3) |
| route_weight_t = torch.tensor([[0.5, 0.3, 0.2]]) |
| knowledge_t = torch.ones(1, 3, 1) |
| capability_t = torch.ones(1, 3, 1) |
| aligned_t = torch.tensor([[[1.0, 0.0], [1.0, 0.0], [1.0, 0.0]]]) |
| diverse_t = torch.tensor([[[1.0, 0.0], [0.0, 1.0], [-1.0, 0.0]]]) |
|
|
| aligned = merge_selected_agent_workspaces( |
| aligned_t, |
| route_ids_t, |
| route_weight_t, |
| knowledge_t=knowledge_t, |
| capability_t=capability_t, |
| correction_pressure_t=torch.zeros(()), |
| release_confidence_t=torch.ones(()), |
| ) |
| diverse = merge_selected_agent_workspaces( |
| diverse_t, |
| route_ids_t, |
| route_weight_t, |
| knowledge_t=knowledge_t, |
| capability_t=capability_t, |
| correction_pressure_t=torch.zeros(()), |
| release_confidence_t=torch.ones(()), |
| ) |
|
|
| assert diverse.disagreement_t[0] > aligned.disagreement_t[0] |
| aligned_delta = aligned.rehearsal_delta_t.float().square().mean().sqrt() |
| diverse_delta = diverse.rehearsal_delta_t.float().square().mean().sqrt() |
| assert diverse_delta > aligned_delta |
|
|
|
|
| def test_cross_examination_scores_proposals_against_peer_evidence_union() -> None: |
| selected_t, route_ids_t, route_weight_t, _, _ = _workspace( |
| [[0.1], [0.1], [0.1], [0.1]], |
| [[0.1], [0.1], [0.1], [0.1]], |
| ) |
| weak_peers_knowledge_t = torch.tensor( |
| [[[0.05], [3.0], [3.0], [3.0]]], dtype=torch.float32 |
| ) |
| all_weak_knowledge_t = torch.full((1, 4, 1), 0.05) |
|
|
| strong_peers = merge_selected_agent_workspaces( |
| selected_t, |
| route_ids_t, |
| route_weight_t, |
| knowledge_t=weak_peers_knowledge_t, |
| capability_t=weak_peers_knowledge_t.clone(), |
| correction_pressure_t=torch.ones(()), |
| release_confidence_t=torch.ones(()), |
| ) |
| all_weak = merge_selected_agent_workspaces( |
| selected_t, |
| route_ids_t, |
| route_weight_t, |
| knowledge_t=all_weak_knowledge_t, |
| capability_t=all_weak_knowledge_t.clone(), |
| correction_pressure_t=torch.ones(()), |
| release_confidence_t=torch.ones(()), |
| ) |
|
|
| assert not torch.allclose( |
| strong_peers.acceptance_t[0, 0], all_weak.acceptance_t[0, 0] |
| ) |
| for packet in (strong_peers, all_weak): |
| assert torch.all(packet.acceptance_t > 0) |
| torch.testing.assert_close( |
| packet.acceptance_t.sum(dim=-1), torch.ones(1) |
| ) |
|
|
|
|
| def test_default_signals_keep_plain_route_weighted_merge() -> None: |
| selected_t, route_ids_t, route_weight_t, knowledge_t, capability_t = _workspace( |
| [[0.1], [0.2], [3.0], [0.2]], |
| [[0.1], [0.2], [3.0], [0.2]], |
| ) |
|
|
| packet = merge_selected_agent_workspaces( |
| selected_t, |
| route_ids_t, |
| route_weight_t, |
| knowledge_t=knowledge_t, |
| capability_t=capability_t, |
| ) |
|
|
| expected_t = ( |
| selected_t * route_weight_t.unsqueeze(-1).to(selected_t) |
| ).sum(dim=1) |
| torch.testing.assert_close(packet.merged_t, expected_t, rtol=1e-5, atol=1e-6) |
| torch.testing.assert_close(packet.acceptance_t.sum(dim=-1), torch.ones(1)) |
| assert torch.count_nonzero(packet.rehearsal_delta_t) == 0 |
|
|