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

Function _check_conv_cfg

axlearn/common/convolution.py:28–65  ·  view source on GitHub ↗
(
    *,
    window: Sequence[int],
    strides: Sequence[int],
    padding: ConvPaddingType,
    dilation: Optional[Sequence[int]],
    input_dim: int,
    output_dim: int,
    num_input_dim_groups: int,
)

Source from the content-addressed store, hash-verified

26
27# TODO(yuanliu939): Make this take `BaseConv.Config` directly.
28def _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
68class BaseConv(BaseLayer):

Calls

no outgoing calls

Tested by

no test coverage detected