MCPcopy Create free account
hub / github.com/apple/axlearn / forward

Method forward

axlearn/common/convolution.py:1807–1844  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

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:

Callers

nothing calls this directly

Calls 3

conv_paddingsMethod · 0.95
safe_notFunction · 0.90
conv_explicit_paddingFunction · 0.85

Tested by

no test coverage detected