from __future__ import annotations import unittest from types import SimpleNamespace import torch from torch import nn from metrics.phase2_critic_guided_math import _planned_processed_samples from networks.acdir import ACDiRPolicy from training.phase2_critic.rollouts import rollout_count_set_policy class _TinyTokenizer: eos_token_id = 18 pad_token_id = 0 all_special_ids = [0, 18] chat_template = None def convert_tokens_to_ids(self, token): return 19 if token == "<|mdm_mask|>" else -1 def __call__(self, prompts, return_tensors="pt", padding=True, **_): input_ids = torch.tensor([[1, 2] for _ in prompts], dtype=torch.long) return {"input_ids": input_ids, "attention_mask": torch.ones_like(input_ids)} def batch_decode(self, seqs, skip_special_tokens=True): return [" ".join(map(str, row.tolist())) for row in seqs] def decode(self, seq, skip_special_tokens=True): return " ".join(map(str, seq.tolist())) class _TinyActor(nn.Module): def __init__(self): super().__init__() self.config = SimpleNamespace(hidden_size=8, vocab_size=20) self.emb = nn.Embedding(20, 8) with torch.no_grad(): self.emb.weight.zero_() self.emb.weight[:, 0] = torch.arange(20, dtype=torch.float32) def get_input_embeddings(self): return self.emb def forward( self, input_ids, attention_mask=None, logits_indices=None, output_hidden_states=False, output_last_hidden_state=False, **_, ): hidden = self.emb(input_ids) positions = logits_indices.cpu() logits = torch.full((input_ids.shape[0], positions.shape[1], 20), -10.0) top = (positions % 17 + 1).long() logits.scatter_(2, top.unsqueeze(-1), 10.0) hidden_states = (hidden,) if output_last_hidden_state else None return SimpleNamespace(logits=logits, hidden_states=hidden_states) class _CountSetCritic(nn.Module): count_support = (0, 1, 2) def __init__(self): super().__init__() self.bias = nn.Parameter(torch.zeros(())) def forward_policy( self, hidden_states, token_embeddings, time_embed, candidate_mask, **_, ): token_scores = token_embeddings[..., 0].to(hidden_states.dtype) + self.bias count_logits = torch.tensor( [[-1.0, 2.0, 0.0]], dtype=hidden_states.dtype, device=hidden_states.device, ).expand(hidden_states.shape[0], -1).contiguous() return { "token_scores": token_scores, "remask_token_logits": token_scores, "count_logits": count_logits, "count_support": self.count_support, "retention_prior_logits": token_scores, "delta_value_logits": torch.zeros_like(token_scores), "state_value": hidden_states.new_zeros((hidden_states.shape[0],)), "encoded": hidden_states, } def _rollout(**overrides): kwargs = dict( actor=_TinyActor(), critic=_CountSetCritic(), tokenizer=_TinyTokenizer(), batch={"problems": ["a", "bb"]}, reward_fn=None, device=torch.device("cpu"), precision_dtype=torch.float32, time_embed_dim=8, steps=4, gen_length=4, block_length=2, no_sample=True, mask_id=19, eos_id=18, compute_rewards=False, return_responses=True, lookback_blocks=1, remask_min_age_current=0, max_total_remask_per_sample=2, reforward_after_remask=True, ) kwargs.update(overrides) return rollout_count_set_policy(**kwargs) class ReleaseRegressionTest(unittest.TestCase): def test_clean_progress_planning_handles_padded_ranks(self): self.assertEqual(_planned_processed_samples(5, 2, 2, 0), 4) self.assertEqual(_planned_processed_samples(5, 2, 2, 1), 5) self.assertEqual(_planned_processed_samples(6, 4, 1, 0), 4) self.assertEqual(_planned_processed_samples(6, 4, 1, 1), 6) def test_empty_context_rows_are_made_attention_safe(self): class Recorder(nn.Module): def __init__(self): super().__init__() self.padding_mask = None def forward(self, encoded, src_key_padding_mask=None): self.padding_mask = src_key_padding_mask.detach().clone() return encoded policy = ACDiRPolicy( hidden_size=4, time_embed_dim=4, mlp_hidden=8, policy_dim=8, encoder_layers=1, encoder_heads=2, dropout=0.0, ) recorder = Recorder() policy.encoder = recorder candidate = torch.tensor([[False, False, False], [True, False, True]]) context = torch.tensor([[False, False, False], [False, True, False]]) policy.forward_policy( torch.randn(2, 3, 4), torch.randn(2, 3, 4), torch.randn(2, 4), candidate, context_mask=context, ) safe_context = ~recorder.padding_mask self.assertEqual(safe_context.tolist(), [[True, False, False], [True, True, True]]) def test_deterministic_rollout_matches_preoptimization_golden(self): torch.manual_seed(123) out = _rollout() self.assertEqual(out["tokens"].tolist(), [[1, 2, 3, 4, 5, 6]] * 2) self.assertEqual(out["responses"], ["3 4 5 6"] * 2) self.assertEqual(out["remask_counts"].tolist(), [2, 2]) self.assertEqual(out["unique_remask_counts"].tolist(), [2, 2]) self.assertEqual(out["unmask_decisions"].tolist(), [4, 4]) self.assertEqual(out["policy_decisions"].tolist(), [4, 4]) self.assertEqual(out["count0_decisions"].tolist(), [2, 2]) self.assertEqual(out["count1_decisions"].tolist(), [2, 2]) self.assertEqual(out["count2_decisions"].tolist(), [0, 0]) def test_fork_resume_trace_has_fresh_time_left(self): common = dict( no_sample=False, temperature=0.4, return_tokens=False, return_responses=False, acdir_deferred_unmask=False, sample_remask=False, ) capture = _rollout(mti_capture_transition=True, **common) self.assertIsNotNone(capture["mti_transition_state"]) repaired = _rollout( mti_transition_state=capture["mti_transition_state"], cfpg_branch="remask", return_mti_trace=True, **common, ) self.assertGreaterEqual(len(repaired["mti_trace"]["repair_steps"]), 1) for step in repaired["mti_trace"]["repair_steps"]: self.assertEqual(tuple(step["time_left"].shape), (2,)) def test_ablation_trace_has_fresh_time_left(self): torch.manual_seed(0) out = _rollout( critic=None, ablation_remask_policy="low_confidence", ablation_remask_probability=1.0, return_mti_trace=True, ) self.assertGreaterEqual(len(out["mti_trace"]["repair_steps"]), 1) for step in out["mti_trace"]["repair_steps"]: self.assertEqual(tuple(step["time_left"].shape), (2,)) if __name__ == "__main__": unittest.main()