Skip to content

Grouped Conv1d / Conv2d initialization uses ungrouped fan-in #4536

Description

@Evihut

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}")

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions