From 0b4f35f4e554c20ef69d91594c905b6ca2effcec Mon Sep 17 00:00:00 2001 From: Moinuddin Shaik Date: Tue, 22 Sep 2026 21:38:05 +0530 Subject: [PATCH] Fix grouped conv init fan_in --- python/mlx/nn/layers/convolution.py | 4 ++-- python/tests/test_nn.py | 18 ++++++++++++++++++ 2 files changed, 20 insertions(+), 2 deletions(-) diff --git a/python/mlx/nn/layers/convolution.py b/python/mlx/nn/layers/convolution.py index 40f34ff569..9241e2c148 100644 --- a/python/mlx/nn/layers/convolution.py +++ b/python/mlx/nn/layers/convolution.py @@ -51,7 +51,7 @@ def __init__( f"divisible by the number of groups ({groups})" ) - scale = math.sqrt(1 / (in_channels * kernel_size)) + scale = math.sqrt(1 / (in_channels // groups * kernel_size)) self.weight = mx.random.uniform( low=-scale, high=scale, @@ -132,7 +132,7 @@ def __init__( lambda x: (x, x) if isinstance(x, int) else x, (kernel_size, stride, padding), ) - scale = math.sqrt(1 / (in_channels * kernel_size[0] * kernel_size[1])) + scale = math.sqrt(1 / (in_channels // groups * kernel_size[0] * kernel_size[1])) self.weight = mx.random.uniform( low=-scale, high=scale, diff --git a/python/tests/test_nn.py b/python/tests/test_nn.py index 26b8fd1162..3bf7ade8ba 100644 --- a/python/tests/test_nn.py +++ b/python/tests/test_nn.py @@ -1,5 +1,6 @@ # Copyright © 2023-2024 Apple Inc. +import math import os import tempfile import unittest @@ -1030,6 +1031,23 @@ def test_conv2d(self): y = c(x) self.assertEqual(y.shape, (4, 7, 7, 8)) + def test_conv_grouped_init(self): + # weights are drawn from U(-s, s) with s = 1 / sqrt(fan_in), where + # fan_in only counts the input channels each group sees + mx.random.seed(0) + for groups in (1, 2, 4): + layers = [ + nn.Conv1d(64, 64, kernel_size=3, groups=groups), + nn.Conv2d(64, 64, kernel_size=3, groups=groups), + ] + for layer in layers: + w = layer.weight + fan_in = math.prod(w.shape[1:]) + bound = 1 / math.sqrt(fan_in) + w_max = mx.abs(w).max().item() + self.assertLessEqual(w_max, bound) + self.assertGreater(w_max, 0.95 * bound) + def test_conv_transpose_extra_repr(self): self.assertIn( "kernel_size=(3, 5)",