MCPcopy Create free account
hub / github.com/LUMIA-Group/MemoryDecoder / forward

Method forward

demo/memDec.py:57–106  ·  view source on GitHub ↗

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,
    )

Source from the content-addressed store, hash-verified

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()

Callers 1

generateMethod · 0.95

Calls 1

MemoryDecoderOutputClass · 0.85

Tested by

no test coverage detected