MCPcopy Create free account
hub / github.com/RL-Align/RL-Kernel / FakeReferenceModel

Class FakeReferenceModel

tests/test_stateless_executor.py:25–47  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

23
24
25class FakeReferenceModel(torch.nn.Module):
26 def __init__(self, logits: torch.Tensor):
27 super().__init__()
28 self.register_buffer("fixed_logits", logits)
29 self.config = SimpleNamespace(use_cache=True, _attn_implementation="eager")
30 self.generation_config = SimpleNamespace(use_cache=True)
31 self.use_cache_calls: list[bool | None] = []
32 self.config_use_cache_calls: list[bool | None] = []
33 self.generation_config_use_cache_calls: list[bool | None] = []
34 self.attn_implementation_calls: list[str | None] = []
35
36 def forward(self, input_ids, attention_mask=None, use_cache=None):
37 del attention_mask
38 self.use_cache_calls.append(use_cache)
39 self.config_use_cache_calls.append(getattr(self.config, "use_cache", None))
40 self.generation_config_use_cache_calls.append(
41 getattr(self.generation_config, "use_cache", None)
42 )
43 self.attn_implementation_calls.append(getattr(self.config, "_attn_implementation", None))
44 return SimpleNamespace(
45 logits=self.fixed_logits[: input_ids.shape[0], : input_ids.shape[1]],
46 past_key_values=None,
47 )
48
49
50class FakeRewardModel(torch.nn.Module):

Calls

no outgoing calls