Encode input sequences as chunks. Args: x: Encoder input features. (1, T_in, F) x_len: Encoder input features lengths. (1,) processed_frames: Number of frames already seen. left_context: Number of frames in left context. right_conte
(
self,
x: torch.Tensor,
x_len: torch.Tensor,
processed_frames: torch.tensor,
chunk_size: int = 16,
left_context: int = 32,
right_context: int = 0,
)
| 1255 | return x |
| 1256 | |
| 1257 | def chunk_forward( |
| 1258 | self, |
| 1259 | x: torch.Tensor, |
| 1260 | x_len: torch.Tensor, |
| 1261 | processed_frames: torch.tensor, |
| 1262 | chunk_size: int = 16, |
| 1263 | left_context: int = 32, |
| 1264 | right_context: int = 0, |
| 1265 | ) -> torch.Tensor: |
| 1266 | """Encode input sequences as chunks. |
| 1267 | Args: |
| 1268 | x: Encoder input features. (1, T_in, F) |
| 1269 | x_len: Encoder input features lengths. (1,) |
| 1270 | processed_frames: Number of frames already seen. |
| 1271 | left_context: Number of frames in left context. |
| 1272 | right_context: Number of frames in right context. |
| 1273 | Returns: |
| 1274 | x: Encoder outputs. (B, T_out, D_enc) |
| 1275 | """ |
| 1276 | mask = make_source_mask(x_len) |
| 1277 | x, mask = self.embed(x, mask, None) |
| 1278 | |
| 1279 | if left_context > 0: |
| 1280 | processed_mask = ( |
| 1281 | torch.arange(left_context, device=x.device).view(1, left_context).flip(1) |
| 1282 | ) |
| 1283 | processed_mask = processed_mask >= processed_frames |
| 1284 | mask = torch.cat([processed_mask, mask], dim=1) |
| 1285 | pos_enc = self.pos_enc(x, left_context=left_context) |
| 1286 | x = self.encoders.chunk_forward( |
| 1287 | x, |
| 1288 | pos_enc, |
| 1289 | mask, |
| 1290 | chunk_size=chunk_size, |
| 1291 | left_context=left_context, |
| 1292 | right_context=right_context, |
| 1293 | ) |
| 1294 | |
| 1295 | if right_context > 0: |
| 1296 | x = x[:, 0:-right_context, :] |
| 1297 | |
| 1298 | if self.time_reduction_factor > 1: |
| 1299 | x = x[:, :: self.time_reduction_factor, :] |
| 1300 | return x |
nothing calls this directly
no test coverage detected