(
self,
input: torch.Tensor,
att_mask: torch.Tensor,
input_position: Union[int, torch.Tensor] = 0,
kv_caches: Optional[List[T_CACHE]] = None,
)
| 104 | self.gradient_checkpointing = gradient_checkpointing |
| 105 | |
| 106 | def forward( |
| 107 | self, |
| 108 | input: torch.Tensor, |
| 109 | att_mask: torch.Tensor, |
| 110 | input_position: Union[int, torch.Tensor] = 0, |
| 111 | kv_caches: Optional[List[T_CACHE]] = None, |
| 112 | ) -> Tuple[torch.Tensor, Union[List[T_CACHE], None]]: |
| 113 | xs, pos_emb = self.pos_enc(input, offset=input_position) |
| 114 | if self.use_sdpa: |
| 115 | att_mask = mask_to_bias(att_mask, xs.dtype) |
| 116 | |
| 117 | if self.gradient_checkpointing and self.training: |
| 118 | xs = self.forward_layers_checkpointed(xs, att_mask, pos_emb) |
| 119 | else: |
| 120 | xs, kv_caches = self.forward_layers(xs, att_mask, pos_emb, |
| 121 | kv_caches) |
| 122 | if self.pre_norm and self.final_norm is not None: |
| 123 | xs = self.final_norm(xs) |
| 124 | return xs, kv_caches |
| 125 | |
| 126 | def forward_layers( |
| 127 | self, |
nothing calls this directly
no test coverage detected