Forward pass that returns **fused log‑probs** as logits. We keep separate caches for each sub‑model.
(
self,
input_ids: torch.LongTensor,
attention_mask: Optional[torch.LongTensor] = None,
past_key_values: Optional[Tuple] = None,
knn_past_key_values: Optional[Tuple] = None,
use_cache: bool = True,
**kwargs,
)
| 55 | # 1. forward() |
| 56 | # ------------------------------------------------------------------ # |
| 57 | def forward( |
| 58 | self, |
| 59 | input_ids: torch.LongTensor, |
| 60 | attention_mask: Optional[torch.LongTensor] = None, |
| 61 | past_key_values: Optional[Tuple] = None, |
| 62 | knn_past_key_values: Optional[Tuple] = None, |
| 63 | use_cache: bool = True, |
| 64 | **kwargs, |
| 65 | ): |
| 66 | """ |
| 67 | Forward pass that returns **fused log‑probs** as logits. |
| 68 | We keep separate caches for each sub‑model. |
| 69 | """ |
| 70 | base_outputs = self.base_lm( |
| 71 | input_ids=input_ids, |
| 72 | attention_mask=attention_mask, |
| 73 | past_key_values=past_key_values, |
| 74 | use_cache=use_cache, |
| 75 | **kwargs, |
| 76 | ) |
| 77 | knn_outputs = self.knn_generator( |
| 78 | input_ids=input_ids, |
| 79 | attention_mask=attention_mask, |
| 80 | past_key_values=knn_past_key_values, |
| 81 | use_cache=use_cache, |
| 82 | **kwargs, |
| 83 | ) |
| 84 | |
| 85 | # Temperature on k‑NN logits only |
| 86 | logits_base = base_outputs.logits # (B, T, V) |
| 87 | logits_knn = knn_outputs.logits |
| 88 | if self.knn_temp != 1.0: |
| 89 | logits_knn = logits_knn / self.knn_temp |
| 90 | |
| 91 | # Convert to log‑probabilities first (numerically stable when fusing) |
| 92 | logp_base = F.log_softmax(logits_base, dim=-1) |
| 93 | logp_knn = F.log_softmax(logits_knn, dim=-1) |
| 94 | |
| 95 | logp_joint = torch.logaddexp( |
| 96 | logp_base + torch.log(torch.tensor(1.0 - self.lmbda, device=logp_base.device)), |
| 97 | logp_knn + torch.log(torch.tensor(self.lmbda, device=logp_base.device)), |
| 98 | ) |
| 99 | |
| 100 | return MemoryDecoderOutput( |
| 101 | logits=logp_joint, |
| 102 | past_key_values=base_outputs.past_key_values, |
| 103 | knn_past_key_values=knn_outputs.past_key_values, |
| 104 | hidden_states=None, |
| 105 | attentions=None |
| 106 | ) |
| 107 | |
| 108 | # ------------------------------------------------------------------ # |
| 109 | # 2. generate() |
no test coverage detected