| 23 | |
| 24 | |
| 25 | class 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 | |
| 50 | class FakeRewardModel(torch.nn.Module): |
no outgoing calls