(
*,
window: Sequence[int],
strides: Sequence[int],
padding: ConvPaddingType,
dilation: Optional[Sequence[int]],
input_dim: int,
output_dim: int,
num_input_dim_groups: int,
)
| 26 | |
| 27 | # TODO(yuanliu939): Make this take `BaseConv.Config` directly. |
| 28 | def _check_conv_cfg( |
| 29 | *, |
| 30 | window: Sequence[int], |
| 31 | strides: Sequence[int], |
| 32 | padding: ConvPaddingType, |
| 33 | dilation: Optional[Sequence[int]], |
| 34 | input_dim: int, |
| 35 | output_dim: int, |
| 36 | num_input_dim_groups: int, |
| 37 | ): |
| 38 | if any(w < 1 for w in window): |
| 39 | raise ValueError(f"window ({window}) must be a positive integer.") |
| 40 | |
| 41 | if any(s < 1 for s in strides): |
| 42 | raise ValueError(f"strides ({strides}) must be a positive integer.") |
| 43 | |
| 44 | if isinstance(padding, str): |
| 45 | if padding not in SUPPORT_CONV_PADDING: |
| 46 | raise ValueError(f"{padding} padding is not supported.") |
| 47 | else: |
| 48 | padding_flattened = jax.tree.leaves(padding) |
| 49 | if any(p < 0 for p in padding_flattened): |
| 50 | raise ValueError("Negative padding is not supported") |
| 51 | |
| 52 | if dilation is not None and any(d < 1 for d in dilation): |
| 53 | raise ValueError(f"dilation ({dilation}) must be a positive integer.") |
| 54 | |
| 55 | if input_dim % num_input_dim_groups != 0: |
| 56 | raise ValueError( |
| 57 | f"input_dim ({input_dim}) must be divisible by " |
| 58 | f"num_input_dim_groups({num_input_dim_groups})." |
| 59 | ) |
| 60 | |
| 61 | if output_dim % num_input_dim_groups != 0: |
| 62 | raise ValueError( |
| 63 | f"output_dim ({output_dim}) must be divisible by " |
| 64 | f"num_input_dim_groups({num_input_dim_groups})." |
| 65 | ) |
| 66 | |
| 67 | |
| 68 | class BaseConv(BaseLayer): |
no outgoing calls
no test coverage detected