from d2l import jax as d2l
from flax import nnx
import jax
from jax import numpy as jnpReal images have channels: RGB has 3, while deep CNN feature maps may have hundreds or thousands.
Going deeper, networks often trade spatial resolution for channel depth, representing more feature types at fewer locations.
This deck:
With c_i input channels, the kernel becomes c_i \times k_h \times k_w — a 2D filter per input channel. The output is the sum of per-channel cross-correlations:
Y = \sum_{c=1}^{c_i} X_c * K_c.
Two input channels: per-channel cross-correlation, then sum. (1{\cdot}1 + 2{\cdot}2 + 4{\cdot}3 + 5{\cdot}4) + (0{\cdot}0 + 1{\cdot}1 + 3{\cdot}2 + 4{\cdot}3) = 56.
Verify against the figure — same numbers:
Array([[ 56., 72.],
[104., 120.]], dtype=float32)
Each output channel comes from its own set of input-channel filters. Stack c_o such filter sets to get a 4-D kernel of shape c_o \times c_i \times k_h \times k_w:
\mathbf{Y}_j = \sum_{c=1}^{c_i} \mathbf{X}_c * \mathbf{K}_{j, c} \quad\text{for}\quad j = 1, \ldots, c_o.
Intuition: each of the c_o output channels is a different combination of inputs, learned to detect a different feature. Together, the channels form a learned feature representation.
Apply the multi-input-channel function c_o times and stack the results along a new leading axis:
Build a 3-output-channel kernel by stacking three offset copies:
(3, 2, 2, 2)
Array([[[ 56., 72.],
[104., 120.]],
[[ 76., 100.],
[148., 172.]],
[[ 96., 128.],
[192., 224.]]], dtype=float32)
A conv layer with c_o outputs, c_i inputs, and a k_h \times k_w kernel has
c_o \cdot c_i \cdot k_h \cdot k_w \;+\; c_o
learnable parameters. Standard sizes:
If input and output channel counts grow together, parameter count grows quadratically, which constrains how quickly networks can widen.
A 1 \times 1 kernel has no spatial extent beyond its current position.
Because it acts as a per-pixel fully connected layer across channels. At every spatial position, it computes a linear combination of the c_i input channels into the c_o output channels:
1×1 conv: 3 input channels × 2 output channels. Each output pixel = a 2×3 matrix-vector product on the input channel vector at that position.
At each pixel, the 1×1 conv applies the same c_o \times c_i matrix to the input channel vector. Reshape the spatial axes out and it’s a single matrix multiply:
Common uses include:
A dense conv connects every input channel to every output channel. Split the channels into g groups and convolve each group separately: parameters and compute drop by a factor of g (as in ResNeXt).
The extreme g = c_i = c_o is a depthwise convolution: one k \times k filter per channel, no channel mixing at all.
Dense conv mixes all input channels; depthwise filters each channel separately; pointwise 1×1 mixes them back.
Factor spatial filtering from channel mixing: depthwise k \times k, then pointwise 1 \times 1 (MobileNet, Xception).
\frac{\text{separable cost}}{\text{dense cost}} = \frac{1}{c_o} + \frac{1}{k^2} \approx \frac{1}{9} \;\text{ for } k = 3.
Parameter counts confirm it, 147k vs. 17.5k:
c_i, c_o, k = 128, 128, 3
X = jax.random.normal(d2l.get_key(), (1, 32, 32, c_i))
dense = nnx.Conv(c_i, c_o, kernel_size=(k, k), padding='SAME',
use_bias=False, rngs=nnx.Rngs(d2l.get_key()))
depthwise = nnx.Conv(c_i, c_i, kernel_size=(k, k), padding='SAME',
feature_group_count=c_i, use_bias=False,
rngs=nnx.Rngs(d2l.get_key()))
pointwise = nnx.Conv(c_i, c_o, kernel_size=(1, 1), use_bias=False,
rngs=nnx.Rngs(d2l.get_key()))
Y = depthwise(X)
assert dense(X).shape == pointwise(Y).shape
size = lambda model: sum(p.size for p in jax.tree_util.tree_leaves(
nnx.state(model, nnx.Param)))
p_dense = size(dense)
p_sep = size(depthwise) + size(pointwise)
p_dense, p_sep, p_dense / p_sep(147456, 17536, 8.408759124087592)
A h \times w image with k \times k kernel and c_i \to c_o channels takes
\mathcal{O}(h \cdot w \cdot k^2 \cdot c_i \cdot c_o)
operations. For a 256×256 image, 5×5 kernel, and 128→128 channels, counting multiplications and additions separately gives more than 53 billion operations for one layer.
Consequently: