(self, input_ids, attention_mask=None, use_cache=None)
| 106 | self.use_cache_calls: list[bool | None] = [] |
| 107 | |
| 108 | def forward(self, input_ids, attention_mask=None, use_cache=None): |
| 109 | self.use_cache_calls.append(use_cache) |
| 110 | hidden = self.embedding(input_ids.long()) |
| 111 | mask = attention_mask.to(device=hidden.device, dtype=hidden.dtype).unsqueeze(-1) |
| 112 | pooled = (hidden * mask).sum(dim=1) / mask.sum(dim=1).clamp_min(1.0) |
| 113 | return { |
| 114 | "logits": self.lm_head(hidden), |
| 115 | "rewards": self.reward_head(pooled).squeeze(-1), |
| 116 | "past_key_values": _tiny_past_key_values(hidden, use_cache), |
| 117 | } |
| 118 | |
| 119 | |
| 120 | def _tiny_past_key_values(hidden: torch.Tensor, use_cache: bool | None): |
no test coverage detected