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

Class ReferenceModelWrapper

rl_engine/alignment/model_wrappers.py:113–160  ·  view source on GitHub ↗

Standard adapter for the frozen reference model used by KL penalties.

Source from the content-addressed store, hash-verified

111
112
113class 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 )

Calls

no outgoing calls