acdir-llada-math500 / tests /test_release_regression.py
NYCU-MLLab's picture
Upload folder using huggingface_hub
4a28d4d verified
Raw
History Blame Contribute Delete
7.39 kB
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()