MCPcopy Create free account
hub / github.com/Pints-AI/1.5-Pints / forward

Method forward

lit_gpt/model.py:255–280  ·  view source on GitHub ↗
(
        self,
        x: torch.Tensor,
        rope: RoPECache,
        max_seq_length: int,
        mask: Optional[torch.Tensor] = None,
        input_pos: Optional[torch.Tensor] = None,
        kv_cache: Optional[KVCache] = None,
    )

Source from the content-addressed store, hash-verified

253 self.config = config
254
255 def forward(
256 self,
257 x: torch.Tensor,
258 rope: RoPECache,
259 max_seq_length: int,
260 mask: Optional[torch.Tensor] = None,
261 input_pos: Optional[torch.Tensor] = None,
262 kv_cache: Optional[KVCache] = None,
263 ) -> Tuple[torch.Tensor, Optional[KVCache]]:
264 n_1 = self.norm_1(x)
265 h, new_kv_cache = self.attn(
266 n_1, rope, max_seq_length, mask, input_pos, kv_cache
267 )
268 if self.config.parallel_residual:
269 n_2 = n_1 if self.config.shared_attention_norm else self.norm_2(x)
270 x = x + h + self.mlp(n_2)
271 else:
272 if self.config.shared_attention_norm:
273 raise NotImplementedError(
274 'No checkpoint amongst the ones we support uses this configuration'
275 ' (non-parallel residual and shared attention norm).'
276 )
277
278 x = x + h
279 x = x + self.mlp(self.norm_2(x))
280 return x, new_kv_cache
281
282
283class CausalSelfAttention(nn.Module):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected