Describe the bug
Conv1d and Conv2d sample weights from U(-s, s) with
s = 1 / sqrt(in_channels * prod(kernel_size))
but allocate the last axis as in_channels // groups. The two expressions agree only when groups == 1. For any grouped convolution the sampling range is therefore sqrt(groups) too narrow.
# python/mlx/nn/layers/convolution.py (Conv1d)
scale = math.sqrt(1 / (in_channels * kernel_size))
self.weight = mx.random.uniform(
low=-scale,
high=scale,
shape=(out_channels, kernel_size, in_channels // groups),
)
Conv2d uses the same scale with kernel_size[0] * kernel_size[1].
The intended bound is 1 / sqrt(fan_in). For a grouped convolution,
fan_in = (in_channels // groups) * prod(kernel_size)
because each output channel only sees the channels in its own group. PyTorch includes this group factor in the corresponding bound.
To Reproduce
Script at the bottom, run against the released wheel with no local repo checkout. Deterministic under mx.random.seed(0) (verified by running twice and comparing output).
fan_in is taken from the allocated weight.shape, not recomputed from the formula under test.
layer groups weight shape fan_in 1/sqrt(fan_in) observed |w|max ratio
Conv1d 1 (64, 3, 64) 192 0.072169 0.072163 1.000
Conv1d 2 (64, 3, 32) 96 0.102062 0.072167 0.707
Conv1d 4 (64, 3, 16) 48 0.144338 0.072167 0.500
Conv1d 8 (64, 3, 8) 24 0.204124 0.072160 0.354
Conv1d 64 (64, 3, 1) 3 0.577350 0.071612 0.124
Conv2d 1 (64, 3, 3, 64) 576 0.041667 0.041665 1.000
Conv2d 2 (64, 3, 3, 32) 288 0.058926 0.041666 0.707
Conv2d 4 (64, 3, 3, 16) 144 0.083333 0.041662 0.500
Conv2d 8 (64, 3, 3, 8) 72 0.117851 0.041667 0.354
Conv2d 64 (64, 3, 3, 1) 9 0.333333 0.041620 0.125
The observed sampling bound stays at ~0.0722 (Conv1d) and ~0.0417 (Conv2d) for every groups. The correct bound grows as the weight narrows. The ratio tracks 1 / sqrt(groups) in both layers: 0.707, 0.500, 0.354, 0.125.
Impact
Grouped and depthwise stacks suffer severe activation-scale shrinkage. In an eight-layer depthwise Conv1d stack the output standard deviation falls to 6.466e-10, versus 1.254e-02 for the matching dense stack. That extra shrinkage can make optimization difficult, especially in architectures without normalization or residual paths.
dense, groups=1 std 1.000e+00 -> 1.254e-02
depthwise, groups=64 std 1.000e+00 -> 6.466e-10
depthwise × sqrt(groups) std 1.000e+00 -> 1.544e-02
Rescaling the depthwise weights by sqrt(groups) restores the activation scale to the same order of magnitude as the dense case, which confirms that the missing group factor accounts for the discrepancy in this experiment.
Why this looks unintentional
For groups == 1 the current formula coincides with 1 / sqrt(fan_in). Combined with the change history, that strongly suggests the group factor was omitted by accident:
Suggested fix
Mirror the weight-shape expression in the scale, two lines:
scale = math.sqrt(1 / (in_channels // groups * kernel_size))
scale = math.sqrt(1 / (in_channels // groups * kernel_size[0] * kernel_size[1]))
This leaves the default groups=1 behavior unchanged and does not affect existing serialized weights. It intentionally changes the initialization of newly created grouped convolutions, and therefore the training trajectory of a grouped model under a fixed seed.
Conv3d and the ConvTranspose{1,2,3}d layers have no groups parameter, so they are not affected.
I have this fix and a regression test ready, and am happy to open a PR.
Desktop
- mlx 0.32.2 (released wheel,
pip install mlx)
- Apple M4, macOS 26.6, arm64
- Python 3.12.12
Additional context
Full reproduction script:
import math, platform
import mlx.core as mx
import mlx.nn as nn
mx.random.seed(0)
print(f"mlx {mx.__version__} | Python {platform.python_version()} | "
f"macOS {platform.mac_ver()[0]} {platform.machine()}\n")
def row(c, groups, label):
shp = tuple(c.weight.shape)
fan_in = math.prod(shp[1:])
bound = fan_in ** -0.5
wmax = mx.abs(c.weight).max().item()
print(f"{label:<8} {groups:>6} {str(shp):>20} {fan_in:>7} "
f"{bound:>14.6f} {wmax:>15.6f} {wmax/bound:>7.3f}")
print(f"{'layer':<8} {'groups':>6} {'weight shape':>20} {'fan_in':>7} "
f"{'1/sqrt(fan_in)':>14} {'observed |w|max':>15} {'ratio':>7}")
for g in [1, 2, 4, 8, 64]:
row(nn.Conv1d(64, 64, 3, groups=g), g, "Conv1d")
for g in [1, 2, 4, 8, 64]:
row(nn.Conv2d(64, 64, 3, groups=g), g, "Conv2d")
x0 = mx.random.normal((8, 128, 64))
for g, tag in [(1, "dense, groups=1"), (64, "depthwise, groups=64")]:
x = x0
for _ in range(8):
x = nn.Conv1d(64, 64, 3, padding=1, groups=g)(x)
print(f" {tag:<24} std {float(x0.std().item()):.3e} -> {float(x.std().item()):.3e}")
x = x0
for _ in range(8):
c = nn.Conv1d(64, 64, 3, padding=1, groups=64)
c.weight = c.weight * math.sqrt(64)
x = c(x)
print(f" {'depthwise x sqrt(groups)':<24} std {float(x0.std().item()):.3e} -> "
f"{float(x.std().item()):.3e}")
Describe the bug
Conv1dandConv2dsample weights fromU(-s, s)withbut allocate the last axis as
in_channels // groups. The two expressions agree only whengroups == 1. For any grouped convolution the sampling range is thereforesqrt(groups)too narrow.Conv2duses the same scale withkernel_size[0] * kernel_size[1].The intended bound is
1 / sqrt(fan_in). For a grouped convolution,because each output channel only sees the channels in its own group. PyTorch includes this group factor in the corresponding bound.
To Reproduce
Script at the bottom, run against the released wheel with no local repo checkout. Deterministic under
mx.random.seed(0)(verified by running twice and comparing output).fan_inis taken from the allocatedweight.shape, not recomputed from the formula under test.The observed sampling bound stays at ~0.0722 (
Conv1d) and ~0.0417 (Conv2d) for everygroups. The correct bound grows as the weight narrows. The ratio tracks1 / sqrt(groups)in both layers: 0.707, 0.500, 0.354, 0.125.Impact
Grouped and depthwise stacks suffer severe activation-scale shrinkage. In an eight-layer depthwise
Conv1dstack the output standard deviation falls to6.466e-10, versus1.254e-02for the matching dense stack. That extra shrinkage can make optimization difficult, especially in architectures without normalization or residual paths.Rescaling the depthwise weights by
sqrt(groups)restores the activation scale to the same order of magnitude as the dense case, which confirms that the missing group factor accounts for the discrepancy in this experiment.Why this looks unintentional
For
groups == 1the current formula coincides with1 / sqrt(fan_in). Combined with the change history, that strongly suggests the group factor was omitted by accident:e6306cfee9(2023-11-29), beforegroupsexisted, whenin_channelswas the true fan-in.groupstoConv1d(2024-09-28) and Add groups in nn.Conv2d #1569 added it toConv2d(2024-11-07). Both changed the weight shape; neither updatedscale.Suggested fix
Mirror the weight-shape expression in the scale, two lines:
This leaves the default
groups=1behavior unchanged and does not affect existing serialized weights. It intentionally changes the initialization of newly created grouped convolutions, and therefore the training trajectory of a grouped model under a fixed seed.Conv3dand theConvTranspose{1,2,3}dlayers have nogroupsparameter, so they are not affected.I have this fix and a regression test ready, and am happy to open a PR.
Desktop
pip install mlx)Additional context
Full reproduction script: