From 2f387010a3683e7863e4274796b0ba449103600b Mon Sep 17 00:00:00 2001 From: Shubhransh Gupta <54713516+shubhransh-gupta@users.noreply.github.com> Date: Fri, 28 Aug 2026 22:28:02 +0530 Subject: [PATCH] fix(ops): fix convolve padRight for even-sized kernels in same mode (#466) --- Source/MLX/Ops.swift | 2 +- Tests/MLXTests/OpsTests.swift | 28 ++++++++++++++++++++++++++++ 2 files changed, 29 insertions(+), 1 deletion(-) diff --git a/Source/MLX/Ops.swift b/Source/MLX/Ops.swift index b7e95a330..a76c65878 100644 --- a/Source/MLX/Ops.swift +++ b/Source/MLX/Ops.swift @@ -1003,7 +1003,7 @@ public func convolve( padding = weightSize / 2 } else { let padLeft = weightSize / 2 - let padRight = max(0, padLeft / 2 - 1) + let padRight = max(0, padLeft - 1) input = padded(input, widths: [0, [padLeft, padRight], 0], stream: stream) } diff --git a/Tests/MLXTests/OpsTests.swift b/Tests/MLXTests/OpsTests.swift index 46fd06cc2..4301ac048 100644 --- a/Tests/MLXTests/OpsTests.swift +++ b/Tests/MLXTests/OpsTests.swift @@ -113,4 +113,32 @@ class OpsTests: XCTestCase { XCTAssertNil(b2) } + func testConvolve() { + let a = MLXArray([1, 2, 3, 4, 5] as [Float]) + let vEven = MLXArray([1, 2, 3, 4] as [Float]) + let vOdd = MLXArray([1, 2, 3] as [Float]) + + // Even kernel size + let fullEven = convolve(a, vEven, mode: .full) + assertEqual(fullEven, MLXArray([1, 4, 10, 20, 30, 34, 31, 20] as [Float])) + + let validEven = convolve(a, vEven, mode: .valid) + assertEqual(validEven, MLXArray([20, 30] as [Float])) + + let sameEven = convolve(a, vEven, mode: .same) + XCTAssertEqual(sameEven.shape, [5]) + assertEqual(sameEven, MLXArray([4, 10, 20, 30, 34] as [Float])) + + // Odd kernel size + let fullOdd = convolve(a, vOdd, mode: .full) + assertEqual(fullOdd, MLXArray([1, 4, 10, 16, 22, 22, 15] as [Float])) + + let validOdd = convolve(a, vOdd, mode: .valid) + assertEqual(validOdd, MLXArray([10, 16, 22] as [Float])) + + let sameOdd = convolve(a, vOdd, mode: .same) + XCTAssertEqual(sameOdd.shape, [5]) + assertEqual(sameOdd, MLXArray([4, 10, 16, 22, 22] as [Float])) + } + }