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

Method __init__

tests/test_alignment_model_wrappers.py:21–25  ·  view source on GitHub ↗
(self, logits: torch.Tensor, *, output_kind: str = "tensor")

Source from the content-addressed store, hash-verified

19
20class FixedLogitsModel(torch.nn.Module):
21 def __init__(self, logits: torch.Tensor, *, output_kind: str = "tensor"):
22 super().__init__()
23 self.output_kind = output_kind
24 self.weight = torch.nn.Parameter(torch.tensor(1.0))
25 self.register_buffer("base_logits", logits.clone())
26
27 def forward(self, input_ids: torch.Tensor, *, logit_bias: float = 0.0):
28 del input_ids

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected