Stacks stride number of frames into one frame along the time axis. Args: inputs: Tensor of shape [batch, time, input_dim]. paddings: 0/1 boolean Tensor of shape [batch, time], paddings of the input sequences. Returns: stacked_inputs: Tensor of sh
(self, inputs: Tensor, *, paddings: Tensor)
| 1805 | padding: Union[tuple[int, int], Literal["SAME", "VALID", "CAUSAL"]] = "VALID" |
| 1806 | |
| 1807 | def forward(self, inputs: Tensor, *, paddings: Tensor) -> tuple[Tensor, Tensor]: |
| 1808 | """Stacks stride number of frames into one frame along the time axis. |
| 1809 | |
| 1810 | Args: |
| 1811 | inputs: Tensor of shape [batch, time, input_dim]. |
| 1812 | paddings: 0/1 boolean Tensor of shape [batch, time], paddings of the input sequences. |
| 1813 | |
| 1814 | Returns: |
| 1815 | stacked_inputs: Tensor of shape [batch, time // stride, input_dim * stride]. |
| 1816 | stacked_paddings: 0/1 boolean Tensor of shape [batch, time // stride]. An output frame |
| 1817 | is padding if at least one of the stacked input frames is padding. |
| 1818 | |
| 1819 | Raises: |
| 1820 | ValueError: If stride is <= 1. |
| 1821 | """ |
| 1822 | cfg = self.config |
| 1823 | if cfg.stride <= 1: |
| 1824 | raise ValueError(f"stride should be greater than 1, but got {cfg.stride}.") |
| 1825 | |
| 1826 | # For the last partial frame. |
| 1827 | inputs = inputs * safe_not(paddings)[:, :, None] |
| 1828 | |
| 1829 | padding = cfg.padding |
| 1830 | if isinstance(padding, str): |
| 1831 | padding = conv_explicit_padding( |
| 1832 | window=(cfg.stride,), strides=(cfg.stride,), padding=padding, dilation=(1,) |
| 1833 | )[0] |
| 1834 | inputs = jnp.pad(inputs, ((0, 0), padding, (0, 0)), constant_values=0) |
| 1835 | |
| 1836 | batch_size, seq_len, input_dim = inputs.shape |
| 1837 | output_length = seq_len // cfg.stride |
| 1838 | new_shape = [batch_size, output_length, input_dim * cfg.stride] |
| 1839 | # Stack inputs over the time dimension. |
| 1840 | stacked_inputs = jnp.reshape(inputs[:, : output_length * cfg.stride, :], new_shape) |
| 1841 | # An output frame is padding if at least one of the stacked input frames is padding. |
| 1842 | stacked_paddings = self.conv_paddings(paddings) |
| 1843 | stacked_inputs = stacked_inputs * safe_not(stacked_paddings)[:, :, None] |
| 1844 | return stacked_inputs, stacked_paddings |
| 1845 | |
| 1846 | @nowrap |
| 1847 | def conv_paddings(self, paddings: Tensor) -> Tensor: |
nothing calls this directly
no test coverage detected