"""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