Standard adapter for the frozen reference model used by KL penalties.
| 111 | |
| 112 | |
| 113 | class ReferenceModelWrapper(PolicyModelWrapper): |
| 114 | """Standard adapter for the frozen reference model used by KL penalties.""" |
| 115 | |
| 116 | def __init__( |
| 117 | self, |
| 118 | model: torch.nn.Module, |
| 119 | *, |
| 120 | freeze: bool = True, |
| 121 | eval_mode: bool = True, |
| 122 | ): |
| 123 | super().__init__(model) |
| 124 | if freeze: |
| 125 | self.freeze() |
| 126 | if eval_mode: |
| 127 | self.eval() |
| 128 | |
| 129 | def freeze(self) -> "ReferenceModelWrapper": |
| 130 | for parameter in self.model.parameters(): |
| 131 | parameter.requires_grad_(False) |
| 132 | return self |
| 133 | |
| 134 | def forward_logits(self, input_ids: torch.Tensor, **model_kwargs: Any) -> torch.Tensor: |
| 135 | with torch.no_grad(): |
| 136 | return super().forward_logits(input_ids, **model_kwargs) |
| 137 | |
| 138 | def selected_logprobs( |
| 139 | self, |
| 140 | input_ids: torch.Tensor, |
| 141 | token_ids: torch.Tensor, |
| 142 | *, |
| 143 | mask: Optional[torch.Tensor] = None, |
| 144 | logits_start: Optional[int] = None, |
| 145 | logits_end: Optional[int] = None, |
| 146 | temperature: float = 1.0, |
| 147 | output_dtype: torch.dtype = torch.float32, |
| 148 | **model_kwargs: Any, |
| 149 | ) -> torch.Tensor: |
| 150 | with torch.no_grad(): |
| 151 | return super().selected_logprobs( |
| 152 | input_ids, |
| 153 | token_ids, |
| 154 | mask=mask, |
| 155 | logits_start=logits_start, |
| 156 | logits_end=logits_end, |
| 157 | temperature=temperature, |
| 158 | output_dtype=output_dtype, |
| 159 | **model_kwargs, |
| 160 | ) |
no outgoing calls