| 124 | return xs, kv_caches |
| 125 | |
| 126 | def forward_layers( |
| 127 | self, |
| 128 | xs: torch.Tensor, |
| 129 | att_mask: torch.Tensor, |
| 130 | pos_emb: torch.Tensor, |
| 131 | kv_caches: Optional[List[T_CACHE]] = None, |
| 132 | ) -> Tuple[torch.Tensor, Union[List[T_CACHE], None]]: |
| 133 | if self.training: |
| 134 | for (i, layer) in enumerate(self.decoders): |
| 135 | xs, _, _, _ = layer(xs, att_mask, pos_emb) |
| 136 | new_kv_caches = kv_caches |
| 137 | else: |
| 138 | assert kv_caches is not None |
| 139 | new_kv_caches = [] |
| 140 | for (i, layer) in enumerate(self.decoders): |
| 141 | xs, _, new_kv_cache, _ = layer(xs, |
| 142 | att_mask, |
| 143 | pos_emb, |
| 144 | att_cache=(kv_caches[i][0], |
| 145 | kv_caches[i][1])) |
| 146 | new_kv_caches.append(new_kv_cache) |
| 147 | |
| 148 | return xs, new_kv_caches |
| 149 | |
| 150 | @torch.jit.ignore(drop=True) |
| 151 | def forward_layers_checkpointed(self, xs: torch.Tensor, |