Forward chunk. Args: xs_pad: TODO. ilens: TODO. cache: State cache dict for streaming inference. **kwargs: Additional keyword arguments.
(
self,
xs_pad: torch.Tensor,
ilens: torch.Tensor,
cache: dict = None,
**kwargs,
)
| 494 | return overlap_feats |
| 495 | |
| 496 | def forward_chunk( |
| 497 | self, |
| 498 | xs_pad: torch.Tensor, |
| 499 | ilens: torch.Tensor, |
| 500 | cache: dict = None, |
| 501 | **kwargs, |
| 502 | ): |
| 503 | """Forward chunk. |
| 504 | |
| 505 | Args: |
| 506 | xs_pad: TODO. |
| 507 | ilens: TODO. |
| 508 | cache: State cache dict for streaming inference. |
| 509 | **kwargs: Additional keyword arguments. |
| 510 | """ |
| 511 | if cache is None: |
| 512 | cache = {} |
| 513 | is_final = kwargs.get("is_final", False) |
| 514 | xs_pad *= self.output_size() ** 0.5 |
| 515 | if self.embed is None: |
| 516 | xs_pad = xs_pad |
| 517 | else: |
| 518 | xs_pad = self.embed(xs_pad, cache) |
| 519 | if cache["tail_chunk"]: |
| 520 | xs_pad = to_device(cache["feats"], device=xs_pad.device) |
| 521 | else: |
| 522 | xs_pad = self._add_overlap_chunk(xs_pad, cache) |
| 523 | if cache["opt"] is None: |
| 524 | cache_layer_num = len(self.encoders0) + len(self.encoders) |
| 525 | new_cache = [None] * cache_layer_num |
| 526 | else: |
| 527 | new_cache = cache["opt"] |
| 528 | |
| 529 | for layer_idx, encoder_layer in enumerate(self.encoders0): |
| 530 | encoder_outs = encoder_layer.forward_chunk( |
| 531 | xs_pad, new_cache[layer_idx], cache["chunk_size"], cache["encoder_chunk_look_back"] |
| 532 | ) |
| 533 | xs_pad, new_cache[0] = encoder_outs[0], encoder_outs[1] |
| 534 | |
| 535 | for layer_idx, encoder_layer in enumerate(self.encoders): |
| 536 | encoder_outs = encoder_layer.forward_chunk( |
| 537 | xs_pad, |
| 538 | new_cache[layer_idx + len(self.encoders0)], |
| 539 | cache["chunk_size"], |
| 540 | cache["encoder_chunk_look_back"], |
| 541 | ) |
| 542 | xs_pad, new_cache[layer_idx + len(self.encoders0)] = encoder_outs[0], encoder_outs[1] |
| 543 | |
| 544 | if self.normalize_before: |
| 545 | xs_pad = self.after_norm(xs_pad) |
| 546 | if cache["encoder_chunk_look_back"] > 0 or cache["encoder_chunk_look_back"] == -1: |
| 547 | cache["opt"] = new_cache |
| 548 | |
| 549 | return xs_pad, ilens, None |
nothing calls this directly
no test coverage detected