diff --git a/.github/workflows/pull_request.yml b/.github/workflows/pull_request.yml index d8c76493c..751ff4493 100644 --- a/.github/workflows/pull_request.yml +++ b/.github/workflows/pull_request.yml @@ -115,6 +115,7 @@ jobs: run: | xcrun xctest ~/Library/Developer/Xcode/DerivedData/mlx-swift-*/Build/Products/Debug/CmlxTests.xctest xcrun xctest ~/Library/Developer/Xcode/DerivedData/mlx-swift-*/Build/Products/Debug/MLXTests.xctest + xcrun xctest ~/Library/Developer/Xcode/DerivedData/mlx-swift-*/Build/Products/Debug/MLXIntegrationTests.xctest - name: Run Distributed Tests SwiftPM (Xcode, macOS) shell: sh diff --git a/MAINTENANCE.md b/MAINTENANCE.md index f71c5b268..dbf4e1007 100644 --- a/MAINTENANCE.md +++ b/MAINTENANCE.md @@ -91,12 +91,7 @@ dependencies: [.product(name: "MLX", package: "mlx-swift"), .product(name: "MLXOptimizers", package: "mlx-swift")] ``` -10. Update `tools/generate_integration_tests.py` as needed - -``` -import MLXNN -@testable import MLXOptimizers -``` +10. Update `tools/integration_tests/cases.py` as needed, regenerate tests if needed 11. Update tests as needed diff --git a/Package.swift b/Package.swift index 1cb1bf138..f9bde7a86 100644 --- a/Package.swift +++ b/Package.swift @@ -405,6 +405,12 @@ let package = Package( "MLX", "MLXNN", "MLXOptimizers", ] ), + .testTarget( + name: "MLXIntegrationTests", + dependencies: [ + "MLX", "MLXNN", "MLXOptimizers", + ] + ), // ------ // Example programs diff --git a/Source/MLX/Documentation.docc/Articles/converting-python.md b/Source/MLX/Documentation.docc/Articles/converting-python.md index 6c94c240a..ddd45b344 100644 --- a/Source/MLX/Documentation.docc/Articles/converting-python.md +++ b/Source/MLX/Documentation.docc/Articles/converting-python.md @@ -107,8 +107,8 @@ Note: some of the symbols are not linkable. `cos` | ``MLXArray/cos(stream:)`` `cummax` | ``MLXArray/cummax(axis:reverse:inclusive:stream:)`` `cummin` | ``MLXArray/cummin(axis:reverse:inclusive:stream:)`` -`cumprod` | ``MLXArray/cumprod(axis:reverse:inclusive:stream:)`` -`cumsum` | ``MLXArray/cumsum(axis:reverse:inclusive:stream:)`` +`cumprod` | ``MLXArray/cumprod(axis:reverse:inclusive:dtype:stream:)`` +`cumsum` | ``MLXArray/cumsum(axis:reverse:inclusive:dtype:stream:)`` `exp` | ``MLXArray/exp(stream:)`` `flatten` | ``MLXArray/flattened(start:end:stream:)`` `log` | ``MLXArray/log(stream:)`` @@ -190,7 +190,7 @@ This is a mapping of `mx` free functions to their ``MLX`` counterparts. `identity` | ``MLXArray/identity(_:type:stream:)`` `less` | ``MLX/less(_:_:stream:)`` `less_equal` | ``MLX/lessEqual(_:_:stream:)`` -`linspace` | ``MLXArray/linspace(_:_:count:stream:)-(Int,Int,Int,StreamOrDevice)`` +`linspace` | ``MLXArray/linspace(_:_:count:dtype:stream:)-(Int,Int,Int,DType?,StreamOrDevice)`` `load` | ``MLX/loadArray(url:stream:)`` and ``MLX/loadArrays(url:stream:)`` `log` | ``MLX/log(_:stream:)`` `log10` | ``MLX/log10(_:stream:)`` @@ -210,7 +210,7 @@ This is a mapping of `mx` free functions to their ``MLX`` counterparts. `negative` | ``MLX/negative(_:stream:)`` `not_equal` | ``MLX/notEqual(_:_:stream:)`` `ones` | ``MLXArray/ones(_:type:stream:)`` -`ones_like` | ``MLXArray/ones(like:stream:)`` +`ones_like` | ``MLXArray/ones(like:dtype:stream:)`` `pad` | ``MLX/padded(_:width:mode:value:stream:)`` `partition` | ``MLX/partitioned(_:kth:axis:stream:)`` `power` | ``MLX/pow(_:_:stream:)-(MLXArray,MLXArray,_)`` @@ -255,4 +255,4 @@ This is a mapping of `mx` free functions to their ``MLX`` counterparts. `var` | ``MLX/variance(_:axes:keepDims:ddof:stream:)`` `where` | ``MLX/which(_:_:_:stream:)`` `zeros` | ``MLXArray/zeros(_:type:stream:)`` -`zeros_like` | ``MLXArray/zeros(like:stream:)`` +`zeros_like` | ``MLXArray/zeros(like:dtype:stream:)`` diff --git a/Source/MLX/Documentation.docc/Organization/cumulative.md b/Source/MLX/Documentation.docc/Organization/cumulative.md index b9c4aeb6c..541f23d13 100644 --- a/Source/MLX/Documentation.docc/Organization/cumulative.md +++ b/Source/MLX/Documentation.docc/Organization/cumulative.md @@ -25,9 +25,9 @@ These are available as both methods on `MLXArray` and free functions. They each - ``MLXArray/cummax(reverse:inclusive:stream:)`` - ``MLXArray/cummin(axis:reverse:inclusive:stream:)`` - ``MLXArray/cummin(reverse:inclusive:stream:)`` -- ``MLXArray/cumprod(axis:reverse:inclusive:stream:)`` +- ``MLXArray/cumprod(axis:reverse:inclusive:dtype:stream:)`` - ``MLXArray/cumprod(reverse:inclusive:stream:)`` -- ``MLXArray/cumsum(axis:reverse:inclusive:stream:)`` +- ``MLXArray/cumsum(axis:reverse:inclusive:dtype:stream:)`` - ``MLXArray/cumsum(reverse:inclusive:stream:)`` - ``MLXArray/logCumsumExp(axis:reverse:inclusive:stream:)`` - ``MLXArray/logCumsumExp(reverse:inclusive:stream:)`` diff --git a/Source/MLX/Documentation.docc/Organization/initialization.md b/Source/MLX/Documentation.docc/Organization/initialization.md index 447ba6ed2..dc3e6132f 100644 --- a/Source/MLX/Documentation.docc/Organization/initialization.md +++ b/Source/MLX/Documentation.docc/Organization/initialization.md @@ -264,8 +264,8 @@ there are specific initializers to request it: - ``MLXArray/full(_:values:type:stream:)`` - ``MLXArray/full(_:values:stream:)`` - ``MLXArray/identity(_:type:stream:)`` -- ``MLXArray/linspace(_:_:count:stream:)-(Int,Int,Int,StreamOrDevice)`` -- ``MLXArray/linspace(_:_:count:stream:)-(Double,Double,Int,StreamOrDevice)`` +- ``MLXArray/linspace(_:_:count:dtype:stream:)-(Int,Int,Int,DType?,StreamOrDevice)`` +- ``MLXArray/linspace(_:_:count:dtype:stream:)-(Double,Double,Int,DType?,StreamOrDevice)`` - ``MLXArray/repeated(_:count:axis:stream:)`` - ``MLXArray/repeated(_:count:stream:)`` - ``MLXArray/repeat(_:count:axis:stream:)`` @@ -282,8 +282,8 @@ there are specific initializers to request it: - ``MLX/full(_:values:type:stream:)`` - ``MLX/full(_:values:stream:)`` - ``MLX/identity(_:type:stream:)`` -- ``linspace(_:_:count:endpoint:stream:)-47tc1`` -- ``linspace(_:_:count:endpoint:stream:)-38sfd`` +- ``linspace(_:_:count:endpoint:dtype:stream:)-2b6eu`` +- ``linspace(_:_:count:endpoint:dtype:stream:)-8k1d2`` - ``MLXArray/repeated(_:count:axis:stream:)`` - ``MLXArray/repeated(_:count:stream:)`` - ``MLX/repeat(_:count:axis:stream:)`` diff --git a/Source/MLX/Documentation.docc/free-functions.md b/Source/MLX/Documentation.docc/free-functions.md index e0592dbec..c54ebddbf 100644 --- a/Source/MLX/Documentation.docc/free-functions.md +++ b/Source/MLX/Documentation.docc/free-functions.md @@ -106,8 +106,8 @@ operations as methods for convenience. - ``MLX/full(_:values:type:stream:)`` - ``MLX/full(_:values:stream:)`` - ``MLX/identity(_:type:stream:)`` -- ``linspace(_:_:count:endpoint:stream:)-47tc1`` -- ``linspace(_:_:count:endpoint:stream:)-38sfd`` +- ``linspace(_:_:count:endpoint:dtype:stream:)-2b6eu`` +- ``linspace(_:_:count:endpoint:dtype:stream:)-8k1d2`` - ``MLX/repeated(_:count:axis:stream:)`` - ``MLX/repeated(_:count:stream:)`` - ``MLX/repeat(_:count:axis:stream:)`` diff --git a/Source/MLX/Factory.swift b/Source/MLX/Factory.swift index 0abbae72d..0741e0c29 100644 --- a/Source/MLX/Factory.swift +++ b/Source/MLX/Factory.swift @@ -331,31 +331,34 @@ extension MLXArray { MLX.identity(n, dtype: dtype, stream: stream) } - /// Generate `num` evenly spaced numbers over interval `[start, stop]` for `BinaryInteger`. + /// Generate `count` evenly spaced numbers over interval `[start, stop]` for `BinaryInteger`. /// - /// Example: + /// The result is floating point (`float32` by default) even for integer + /// bounds -- see ``linspace(_:_:count:endpoint:dtype:stream:)-2b6eu``. /// /// ```swift - /// // Create a 50 element 1-D array with values from 0 to 50 - /// let r = MLXArray.linspace(0, 50) + /// // [0, 0.5, 1] as float32 + /// let r = MLXArray.linspace(0, 1, count: 3) /// ``` /// /// - Parameters: /// - start: start value /// - stop: stop value /// - count: number of samples + /// - dtype: dtype of the result, `float32` if not specified /// - stream: stream or device to evaluate on /// /// ### See Also /// - - /// - ``linspace(_:_:count:stream:)-92x6l`` + /// - ``linspace(_:_:count:dtype:stream:)-3fx01`` static public func linspace( - _ start: T, _ stop: T, count: Int = 50, stream: StreamOrDevice = .default + _ start: T, _ stop: T, count: Int = 50, dtype: DType? = nil, + stream: StreamOrDevice = .default ) -> MLXArray where T: BinaryInteger { - MLX.linspace(start, stop, count: count, stream: stream) + MLX.linspace(start, stop, count: count, dtype: dtype, stream: stream) } - /// Generate `num` evenly spaced numbers over interval `[start, stop]` for `BinaryFloatingPoint`. + /// Generate `count` evenly spaced numbers over interval `[start, stop]` for `BinaryFloatingPoint`. /// /// Example: /// @@ -368,15 +371,17 @@ extension MLXArray { /// - start: start value /// - stop: stop value /// - count: number of samples + /// - dtype: dtype of the result, derived from `T` if not specified /// - stream: stream or device to evaluate on /// /// ### See Also /// - - /// - ``linspace(_:_:count:stream:)-7m7eg`` + /// - ``linspace(_:_:count:dtype:stream:)-9yqai`` static public func linspace( - _ start: T, _ stop: T, count: Int = 50, stream: StreamOrDevice = .default + _ start: T, _ stop: T, count: Int = 50, dtype: DType? = nil, + stream: StreamOrDevice = .default ) -> MLXArray where T: BinaryFloatingPoint { - MLX.linspace(start, stop, count: count, stream: stream) + MLX.linspace(start, stop, count: count, dtype: dtype, stream: stream) } /// Generate values in the half-open interval `[0, stop)`. @@ -1007,13 +1012,18 @@ public func identity(_ n: Int, dtype: DType, stream: StreamOrDevice = .default) return MLXArray(result) } -/// Generate `num` evenly spaced numbers over interval `[start, stop]`. +/// Generate `count` evenly spaced numbers over interval `[start, stop]`. /// -/// Example: +/// The result is floating point (`float32` by default), matching python's +/// `mx.linspace`, even for integer bounds: evenly spaced values between two +/// integers are generally not integers. Pass `dtype:` for anything else: /// /// ```swift -/// // Create a 50 element 1-D array with values from 0 to 50 -/// let r = MLXArray.linspace(0, 50) +/// // [0, 0.5, 1] as float32 -- *not* [0, 0, 1] as int64 +/// let r = MLXArray.linspace(0, 1, count: 3) +/// +/// // opt in to an integer (truncating) result +/// let i = MLXArray.linspace(0, 10, count: 6, dtype: .int32) /// ``` /// /// - Parameters: @@ -1021,25 +1031,31 @@ public func identity(_ n: Int, dtype: DType, stream: StreamOrDevice = .default) /// - stop: stop value /// - count: number of samples /// - endpoint: if `true` then the endpoint is the last sample, if `false` it is a half-open interval +/// - dtype: dtype of the result, `float32` if not specified /// - stream: stream or device to evaluate on /// /// ### See Also /// - -/// - ``linspace(_:_:count:endpoint:stream:)`` +/// - ``linspace(_:_:count:endpoint:dtype:stream:)-2b6eu`` public func linspace( _ start: T, _ stop: T, count: Int = 50, endpoint: Bool = true, + dtype: DType? = nil, stream: StreamOrDevice = .default ) -> MLXArray where T: BinaryInteger { var result = mlx_array_new() mlx_linspace_endpoint( - &result, Double(start), Double(stop), count.int32, endpoint, T.dtype.cmlxDtype, stream.ctx) + &result, Double(start), Double(stop), count.int32, endpoint, + (dtype ?? .float32).cmlxDtype, stream.ctx) return MLXArray(result) } -/// Generate `num` evenly spaced numbers over interval `[start, stop]`. +/// Generate `count` evenly spaced numbers over interval `[start, stop]`. /// -/// Example: +/// The result dtype follows the bounds (`Float` -> `float32`, +/// `Float16` -> `float16`) except that `Double` produces `float32`, matching +/// ``MLXArray/init(_:)`` and python's `mx.linspace` -- `float64` is not +/// supported on the GPU. Pass `dtype:` to be explicit. /// /// ```swift /// // Create a 50 element 1-D array with values from 0 to 1 @@ -1051,19 +1067,24 @@ public func linspace( /// - stop: stop value /// - count: number of samples /// - endpoint: if `true` then the endpoint is the last sample, if `false` it is a half-open interval +/// - dtype: dtype of the result, derived from `T` if not specified /// - stream: stream or device to evaluate on /// /// ### See Also /// - -/// - ``linspace(_:_:count:endpoint:stream:)`` +/// - ``linspace(_:_:count:endpoint:dtype:stream:)-8k1d2`` public func linspace( _ start: T, _ stop: T, count: Int = 50, endpoint: Bool = true, + dtype: DType? = nil, stream: StreamOrDevice = .default ) -> MLXArray where T: BinaryFloatingPoint { + // Double.dtype is float64, but we do not automatically promote to float64 + // (see HasDType conformances) and it is not available on the GPU + let resolved = dtype ?? (T.dtype == .float64 ? .float32 : T.dtype) var result = mlx_array_new() mlx_linspace_endpoint( - &result, Double(start), Double(stop), count.int32, endpoint, T.dtype.cmlxDtype, stream.ctx) + &result, Double(start), Double(stop), count.int32, endpoint, resolved.cmlxDtype, stream.ctx) return MLXArray(result) } diff --git a/Source/MLX/Ops+Array.swift b/Source/MLX/Ops+Array.swift index b21d857e1..a1cba744f 100644 --- a/Source/MLX/Ops+Array.swift +++ b/Source/MLX/Ops+Array.swift @@ -605,7 +605,7 @@ public func cummin( /// ### See Also /// - /// - ``cumprod(_:reverse:inclusive:dtype:stream:)`` -/// - ``MLXArray/cumprod(axis:reverse:inclusive:stream:)`` +/// - ``MLXArray/cumprod(axis:reverse:inclusive:dtype:stream:)`` public func cumprod( _ array: MLXArray, axis: Int, reverse: Bool = false, inclusive: Bool = true, dtype: DType? = nil, @@ -629,7 +629,7 @@ public func cumprod( /// ### See Also /// - /// - ``cumprod(_:axis:reverse:inclusive:dtype:stream:)`` -/// - ``MLXArray/cumprod(axis:reverse:inclusive:stream:)`` +/// - ``MLXArray/cumprod(axis:reverse:inclusive:dtype:stream:)`` public func cumprod( _ array: MLXArray, reverse: Bool = false, inclusive: Bool = true, dtype: DType? = nil, @@ -653,7 +653,7 @@ public func cumprod( /// ### See Also /// - /// - ``cumsum(_:reverse:inclusive:dtype:stream:)`` -/// - ``MLXArray/cumsum(axis:reverse:inclusive:stream:)`` +/// - ``MLXArray/cumsum(axis:reverse:inclusive:dtype:stream:)`` public func cumsum( _ array: MLXArray, axis: Int, reverse: Bool = false, inclusive: Bool = true, dtype: DType? = nil, @@ -677,7 +677,7 @@ public func cumsum( /// ### See Also /// - /// - ``cumsum(_:axis:reverse:inclusive:dtype:stream:)`` -/// - ``MLXArray/cumsum(axis:reverse:inclusive:stream:)`` +/// - ``MLXArray/cumsum(axis:reverse:inclusive:dtype:stream:)`` public func cumsum( _ array: MLXArray, reverse: Bool = false, inclusive: Bool = true, dtype: DType? = nil, diff --git a/Source/MLX/Ops.swift b/Source/MLX/Ops.swift index 0c689303b..5a0da9e1a 100644 --- a/Source/MLX/Ops.swift +++ b/Source/MLX/Ops.swift @@ -1002,6 +1002,9 @@ public func convolve( if weightSize % 2 == 1 { padding = weightSize / 2 } else { + // even sized weights use asymmetric padding -- this must match + // python's `mx.convolve()` so that the result is the centered + // `input.size` window of the full convolution let padLeft = weightSize / 2 let padRight = max(0, padLeft - 1) @@ -2184,7 +2187,7 @@ public func multiply( /// - public func nanToNum( _ array: MLXArray, - nan: Float = 0, posInf: Float? = 0, negInf: Float? = 0, + nan: Float = 0, posInf: Float? = nil, negInf: Float? = nil, stream: StreamOrDevice = .default ) -> MLXArray { let posInf = mlx_optional_float(value: posInf ?? 0, has_value: posInf != nil) @@ -3110,7 +3113,7 @@ public func tanh(_ array: MLXArray, stream: StreamOrDevice = .default) -> MLXArr /// - Parameters: /// - a: input array /// - b: input array -/// - axes: sum over the last `axes` dimensions +/// - axes: sum over the last `axes` dimensions of `a` and the first `axes` of `b` /// - stream: stream or device to evaluate on /// - Returns: tensor dot product /// @@ -3118,7 +3121,7 @@ public func tanh(_ array: MLXArray, stream: StreamOrDevice = .default) -> MLXArr /// - /// - ``tensordot(_:_:axes:stream:)-(MLXArray,MLXArray,Int,StreamOrDevice)`` public func tensordot( - _ a: MLXArray, _ b: MLXArray, axes: Int = 1, stream: StreamOrDevice = .default + _ a: MLXArray, _ b: MLXArray, axes: Int = 2, stream: StreamOrDevice = .default ) -> MLXArray { var result = mlx_array_new() mlx_tensordot_axis(&result, a.ctx, b.ctx, axes.int32, stream.ctx) diff --git a/Source/MLXNN/Pooling.swift b/Source/MLXNN/Pooling.swift index c978176db..1110f318f 100644 --- a/Source/MLXNN/Pooling.swift +++ b/Source/MLXNN/Pooling.swift @@ -41,9 +41,10 @@ open class Pool: Module, UnaryLayer { // Apply padding if any padding value is greater than 0 if padding.contains(where: { $0 > 0 }) { - // batch and channel dimension get no padding - let padWidths: [IntOrPair] = - [0, 0] + padding.map { .init($0) } + [0, 0] + // one width per dimension: batch and channel get no padding, each + // spatial dimension is padded symmetrically (`IntOrPair` is already + // the (before, after) pair for a single dimension) + let padWidths: [IntOrPair] = [0] + padding.map { IntOrPair($0) } + [0] let paddingValue = paddingValue.asMLXArray(dtype: input.dtype) input = padded(input, widths: padWidths, mode: .constant, value: paddingValue) } diff --git a/Source/MLXNN/PositionalEncoding.swift b/Source/MLXNN/PositionalEncoding.swift index bf690feba..2dca4ea1d 100644 --- a/Source/MLXNN/PositionalEncoding.swift +++ b/Source/MLXNN/PositionalEncoding.swift @@ -156,10 +156,30 @@ final public class ALiBi: Module { public override init() { } + /// The per-head slopes, matching python's `ALiBi.create_alibi_slope()`. + /// + /// For a power of two head count the slopes are `2^(-8i/n)` for `i` in + /// `1...n`. Otherwise python uses the slopes of the next power of two *below* + /// `n` and pads them with every other slope of the next power of two *above*, + /// which is not the same as extending the geometric series. static func alibiSlope(numHeads: Int) -> MLXArray { - let x = pow(pow(2, 8), (1 / Float(numHeads))) - let out = pow(x, -MLXArray(1 ..< (numHeads + 1))) - return out.expandedDimensions(axes: [-1, -2]) + func slopes(_ n: Int) -> [Float] { + let log2n = Foundation.log2(Double(n)) + if log2n == log2n.rounded(.down) { + let start = Foundation.pow(2.0, -Foundation.pow(2.0, 3 - log2n)) + return (1 ... n).map { Float(Foundation.pow(start, Double($0))) } + } + + let closestPowerOf2 = Int(Foundation.pow(2.0, log2n.rounded(.down))) + let interleaved = slopes(2 * closestPowerOf2) + .enumerated() + .filter { $0.offset.isMultiple(of: 2) } + .map { $0.element } + .prefix(n - closestPowerOf2) + return slopes(closestPowerOf2) + interleaved + } + + return MLXArray(slopes(numHeads)).expandedDimensions(axes: [-1, -2]) } static func alibiMatrix(key: Key) -> MLXArray { @@ -167,8 +187,10 @@ final public class ALiBi: Module { return value } + // x1 is a column and x2 a row so that the difference is the (q, k) + // distance matrix -- python: `x1[:, None] - x2[None, :]` let x1 = MLXArray(key.offset ..< key.qSequenceLength).expandedDimensions(axis: 1) - let x2 = MLXArray(0 ..< key.kSequenceLength).expandedDimensions(axis: 1) + let x2 = MLXArray(0 ..< key.kSequenceLength).expandedDimensions(axis: 0) let distanceMatrix = -abs(expandedDimensions((x1 - x2), axes: [0, 1])) let slope = alibiSlope(numHeads: key.numHeads) diff --git a/Source/MLXNN/Recurrent.swift b/Source/MLXNN/Recurrent.swift index 59a87607d..ea3ca857a 100644 --- a/Source/MLXNN/Recurrent.swift +++ b/Source/MLXNN/Recurrent.swift @@ -160,8 +160,12 @@ open class GRU: Module { var n = x_n[.ellipsis, index, 0...] if hidden != nil { - // Note: xProj_n was computed earlier + // Note: hProj_n was computed earlier and already includes bhn n = n + r * hProj_n + } else if let bhn { + // matches python: without an incoming hidden state the hidden + // bias is still applied, gated by r + n = n + r * bhn } n = tanh(n) diff --git a/Source/MLXOptimizers/Optimizers.swift b/Source/MLXOptimizers/Optimizers.swift index a01083b85..42f247517 100644 --- a/Source/MLXOptimizers/Optimizers.swift +++ b/Source/MLXOptimizers/Optimizers.swift @@ -9,6 +9,17 @@ import MLXNN /// ### See Also /// - /// - ``OptimizerBase`` +// Note on scalar precision: these optimizers store their hyperparameters as `Float` +// while python stores them as double, so coefficients derived from them differ in +// the last bits -- `1 - Float(0.95)` is 0.050000012, while python's `1 - 0.95` +// rounds to 0.05. Widening at the point of use (`1 - Double(momentum)`) cannot +// recover python's value: the constant was already rounded when it was stored. +// +// The 2.4e-7 difference is invisible in the smooth optimizers, but Muon's +// Newton-Schulz iteration grows it to ~2e-4 relative over a few steps, which is why +// the generated `Muon` cases compare with a looser tolerance than the rest. See +// `todos/muon-04-pending-p3-float-hyperparameters.md`. + public protocol Optimizer: Updatable, Evaluatable { /// Apply the gradients to the parameters of the model and update the model with the new parameters. @@ -561,17 +572,17 @@ open class Lion: OptimizerBaseArrayState { /// The learning rate public var learningRate: Float - /// The coefficients used for computing running averages of the gradient and its square - public var betas: (Float, Float) = (0.9, 0.999) + /// The coefficients used for computing the gradient momentum and update direction + public var betas: (Float, Float) = (0.9, 0.99) /// The weight decay public var weightDecay: Float = 0.0 /// Initialize the optimizer. /// - Parameters: /// - learningRate: the learning rate - /// - betas: coefficients used for computing running averages of the gradient and its square + /// - betas: coefficients used for computing the gradient momentum and update direction /// - weightDecay:the weight decay - public init(learningRate: Float, betas: (Float, Float) = (0.9, 0.999), weightDecay: Float = 0.0) + public init(learningRate: Float, betas: (Float, Float) = (0.9, 0.99), weightDecay: Float = 0.0) { self.learningRate = learningRate self.betas = betas @@ -710,7 +721,9 @@ open class Adafactor: OptimizerBase { func approvateExpMovingAverage(expAvgSqRow: MLXArray, expAvgSqCol: MLXArray) -> MLXArray { let rFactor = rsqrt(expAvgSqRow / mean(expAvgSqRow, axis: -1, keepDims: true)) let cFactor = rsqrt(expAvgSqCol) - return matmul(rFactor.expandedDimensions(axis: -1), cFactor.expandedDimensions(axis: 0)) + // broadcast rather than matmul so this also works for parameters with more + // than two dimensions + return rFactor.expandedDimensions(axis: -1) * cFactor.expandedDimensions(axis: -2) } override open func applySingle(gradient: MLXArray, parameter: MLXArray, state: State) -> ( @@ -730,8 +743,15 @@ open class Adafactor: OptimizerBase { var update = square(gradient) + eps.0 if factored { - var expAvgSqRow = state.expAvgSqRow! - var expAvgSqCol = state.expAvgSqCol! + // the state is created from the gradient shape so that it always agrees + // with the branch taken here, even if a parameter's rank changed + let rowShape = Array(gradientShape.dropLast()) + let columnShape = Array(gradientShape.dropLast(2)) + [gradientShape.last!] + + var expAvgSqRow = + state.expAvgSqRow ?? MLXArray.zeros(rowShape, dtype: gradient.dtype) + var expAvgSqCol = + state.expAvgSqCol ?? MLXArray.zeros(columnShape, dtype: gradient.dtype) expAvgSqRow = (beta2 * expAvgSqRow) + (1 - beta2) * mean(update, axis: -1) expAvgSqCol = (beta2 * expAvgSqCol) + (1 - beta2) * mean(update, axis: -2) @@ -742,7 +762,7 @@ open class Adafactor: OptimizerBase { update = approvateExpMovingAverage(expAvgSqRow: expAvgSqRow, expAvgSqCol: expAvgSqCol) update = update * gradient } else { - var expAvgSq = state.expAvgSq! + var expAvgSq = state.expAvgSq ?? MLXArray.zeros(like: gradient) expAvgSq = (beta2 * expAvgSq) + (1 - beta2) * update state.expAvgSq = expAvgSq update = rsqrt(expAvgSq) * gradient @@ -752,7 +772,7 @@ open class Adafactor: OptimizerBase { update = learningRate * update if let beta1 { - var expAvg = state.expAvg! + var expAvg = state.expAvg ?? MLXArray.zeros(like: gradient) expAvg = (beta1 * expAvg) + (1 - beta1) * update state.expAvg = expAvg update = expAvg @@ -811,19 +831,19 @@ open class Muon: OptimizerBaseArrayState { } /// Orthogonalize a 2D matrix via a quintic Newton-Schulz iteration. - private func zeropowerViaNewtonSchulz5(_ input: MLXArray, steps: Int) -> MLXArray { + func zeropowerViaNewtonSchulz5(_ input: MLXArray, steps: Int) -> MLXArray { precondition(input.ndim == 2, "Newton-Schulz iteration expects a 2D array") let (a, b, c): (Float, Float, Float) = (3.4445, -4.7750, 2.0315) let transposeNeeded = input.dim(-2) > input.dim(-1) var X = transposeNeeded ? input.transposed(1, 0) : input // Frobenius-normalize so the iteration converges. - X = X / (sqrt((X * X).sum(keepDims: true)) + 1e-7) + X = X / (MLX.norm(X, keepDims: true) + 1e-7) for _ in 0 ..< steps { let A = matmul(X, X.transposed(1, 0)) - let B = b * A + c * matmul(A, A) - X = a * X + matmul(B, X) + let B = addMM(b * A, A, A, alpha: c, beta: 1.0) + X = addMM(a * X, B, X, alpha: 1.0, beta: 1.0) } return transposeNeeded ? X.transposed(1, 0) : X diff --git a/Tests/MLXIntegrationTests/Generated/GeneratedActivationsTests.swift b/Tests/MLXIntegrationTests/Generated/GeneratedActivationsTests.swift new file mode 100644 index 000000000..6bb65b882 --- /dev/null +++ b/Tests/MLXIntegrationTests/Generated/GeneratedActivationsTests.swift @@ -0,0 +1,823 @@ +// Copyright © 2026 Apple Inc. +// +// GENERATED by tools/integration_tests -- DO NOT EDIT. +// +// Values in this file were produced by python mlx and are compared against +// the Swift API. Regenerate with: +// +// python3 tools/integration_tests/generate.py +// +// python mlx: 0.32.2.dev20260910+1f8e74e3f +// vendored mlx (Cmlx): 0.32.2 +// generator revision: 2 +// cases: 34 + +import Foundation +import MLX +import MLXNN +import Testing + +@Suite("generated: Activations") +struct GeneratedActivationsTests { + + @Test("relu") + func test_relu() throws { + try withIntegrationState(seed: 67500) { + let x = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.relu(x) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.5200212001800537, + minimum: 0.0, + maximum: 2.232034921646118, + absoluteSum: 6.2402544021606445, + positionChecksum: 3.4608532587687173, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [0.0, 0.0, 0.0, 0.8907164931297302, 0.0, 0.36071768403053284]), + tolerance: .float32) + } + } + + @Test("reluSquared") + func test_reluSquared() throws { + try withIntegrationState(seed: 41758) { + let x = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.reluSquared(x) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.453746497631073, + minimum: 0.0, + maximum: 3.3485403060913086, + absoluteSum: 5.444957733154297, + positionChecksum: 2.6932687759399414, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.0, 0.5844548344612122, 0.09588181227445602, 0.10871507227420807, 0.0, + 0.499411016702652, + ]), + tolerance: .float32) + } + } + + @Test("relu6") + func test_relu6() throws { + try withIntegrationState(seed: 12149) { + let x = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.relu6(x) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.240987628698349, + minimum: 0.0, + maximum: 0.8727357983589172, + absoluteSum: 2.8918514251708984, + positionChecksum: 2.2303806940714517, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.0, 0.0, 0.0, 0.2970517873764038, 0.8727357983589172, 0.5916443467140198, + ]), + tolerance: .float32) + } + } + + @Test("leakyRelu") + func test_leakyRelu() throws { + try withIntegrationState(seed: 56128) { + let x = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.leakyRelu(x) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 1.0438215732574463, + minimum: -0.018068933859467506, + maximum: 2.521155834197998, + absoluteSum: 12.575902938842773, + positionChecksum: 7.75894292195638, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.7122242450714111, 1.8448302745819092, 0.9383899569511414, + 1.2284408807754517, -0.018068933859467506, 2.521155834197998, + ]), + tolerance: .float32) + } + } + + @Test("leakyRelu/slope") + func test_leakyRelu_slope() throws { + try withIntegrationState(seed: 39880) { + let x = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.leakyRelu(x, negativeSlope: 0.2) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.4364432692527771, + minimum: -0.2560635805130005, + maximum: 1.8702387809753418, + absoluteSum: 6.699383735656738, + positionChecksum: 4.133757909138997, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.6765661239624023, 1.2138452529907227, 0.13987021148204803, + -0.16138756275177002, 1.8702387809753418, 1.159955382347107, + ]), + tolerance: .float32) + } + } + + @Test("elu") + func test_elu() throws { + try withIntegrationState(seed: 86398) { + let x = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.elu(x) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.25546324253082275, + minimum: -0.8628273606300354, + maximum: 2.130934000015259, + absoluteSum: 7.796996116638184, + positionChecksum: 3.408879597981771, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.1739104986190796, 2.130934000015259, -0.8628273606300354, + -0.36768800020217896, 0.5139061212539673, 0.48894003033638, + ]), + tolerance: .float32) + } + } + + @Test("elu/alpha") + func test_elu_alpha() throws { + try withIntegrationState(seed: 33060) { + let x = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.elu(x, alpha: 0.5) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.1415794938802719, + minimum: -0.3538806438446045, + maximum: 1.1850383281707764, + absoluteSum: 4.625144004821777, + positionChecksum: 2.805095672607422, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.09659192711114883, -0.1328100562095642, 0.11225724220275879, + 0.22676324844360352, -0.3352320194244385, -0.26999032497406006, + ]), + tolerance: .float32) + } + } + + @Test("celu") + func test_celu() throws { + try withIntegrationState(seed: 19670) { + let x = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.celu(x) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.14471694827079773, + minimum: -0.7049131393432617, + maximum: 1.0946556329727173, + absoluteSum: 5.854050159454346, + positionChecksum: 2.7480808893839517, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.045488059520721436, 1.0946556329727173, 0.7886189222335815, + -0.6293155550956726, 0.3050518333911896, 0.0885806754231453, + ]), + tolerance: .float32) + } + } + + @Test("celu/alpha") + func test_celu_alpha() throws { + try withIntegrationState(seed: 5167) { + let x = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.celu(x, alpha: 0.5) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.14974400401115417, + minimum: -0.43982020020484924, + maximum: 1.1591838598251343, + absoluteSum: 5.7410783767700195, + positionChecksum: 3.1327006022135415, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.43982020020484924, -0.4007747173309326, -0.206155925989151, + 0.9863490462303162, -0.006057232618331909, 0.0833292305469513, + ]), + tolerance: .float32) + } + } + + @Test("silu") + func test_silu() throws { + try withIntegrationState(seed: 84525) { + let x = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.silu(x) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.1968996822834015, + minimum: -0.2784386873245239, + maximum: 1.4008158445358276, + absoluteSum: 5.122735500335693, + positionChecksum: 2.5903452237447104, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.2193196713924408, 1.4008158445358276, 0.881088376045227, + -0.2784386873245239, 0.7538488507270813, 0.07409554719924927, + ]), + tolerance: .float32) + } + } + + @Test("mish") + func test_mish() throws { + try withIntegrationState(seed: 7302) { + let x = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.mish(x) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.21375694870948792, + minimum: -0.3031258285045624, + maximum: 1.1241786479949951, + absoluteSum: 5.914010047912598, + positionChecksum: 3.54515806833903, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.27689889073371887, -0.29029443860054016, 1.1014254093170166, + 1.1241786479949951, 0.8475918173789978, 0.5795539617538452, + ]), + tolerance: .float32) + } + } + + @Test("selu") + func test_selu() throws { + try withIntegrationState(seed: 76809) { + let x = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.selu(x) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.6044207811355591, + minimum: -1.5526567697525024, + maximum: 0.91837078332901, + absoluteSum: 10.77687931060791, + positionChecksum: 5.267512957255046, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -1.5526567697525024, -1.3209490776062012, -0.4850633144378662, + -0.5401492118835449, 0.6920045018196106, -1.2968387603759766, + ]), + tolerance: .float32) + } + } + + @Test("softplus") + func test_softplus() throws { + try withIntegrationState(seed: 6943) { + let x = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.softplus(x) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.8977645635604858, + minimum: 0.26141318678855896, + maximum: 2.1883463859558105, + absoluteSum: 10.773174285888672, + positionChecksum: 6.323604583740234, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.26141318678855896, 1.2181404829025269, 0.3854416310787201, + 1.010901689529419, 2.1883463859558105, 0.720491886138916, + ]), + tolerance: .float32) + } + } + + @Test("softsign") + func test_softsign() throws { + try withIntegrationState(seed: 40624) { + let x = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.softsign(x) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.004115648567676544, + minimum: -0.6314003467559814, + maximum: 0.7147443294525146, + absoluteSum: 5.524916648864746, + positionChecksum: 3.1572577158610025, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.6314003467559814, 0.5352221131324768, -0.5489830374717712, + -0.627332329750061, 0.7147443294525146, 0.6036931872367859, + ]), + tolerance: .float32) + } + } + + @Test("softshrink") + func test_softshrink() throws { + try withIntegrationState(seed: 92275) { + let x = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.softshrink(x) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.22316786646842957, + minimum: -1.119591474533081, + maximum: 0.3092571496963501, + absoluteSum: 3.999713897705078, + positionChecksum: 1.6630719502766926, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.3092571496963501, -0.43543922901153564, -0.818618893623352, + -0.6015747785568237, 0.041445255279541016, 0.0, + ]), + tolerance: .float32) + } + } + + @Test("softshrink/lambda") + func test_softshrink_lambda() throws { + try withIntegrationState(seed: 6070) { + let x = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.softshrink(x, lambda: 0.2) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.05193863809108734, + minimum: -0.8008320927619934, + maximum: 1.2317386865615845, + absoluteSum: 6.1652302742004395, + positionChecksum: 3.3861312866210938, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.768548309803009, -0.7094013690948486, -0.5380290746688843, + -0.8008320927619934, 0.23463128507137299, 1.2317386865615845, + ]), + tolerance: .float32) + } + } + + @Test("softmin") + func test_softmin() throws { + try withIntegrationState(seed: 54376) { + let x = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.softmin(x) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.3333333432674408, + minimum: 0.08368699252605438, + maximum: 0.767266571521759, + absoluteSum: 4.0, + positionChecksum: 2.2911632855733237, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.08813121169805527, 0.7124828100204468, 0.13925348222255707, + 0.4346010982990265, 0.08368699252605438, 0.5157842040061951, + ]), + tolerance: .float32) + } + } + + @Test("softmin/axis") + func test_softmin_axis() throws { + try withIntegrationState(seed: 147) { + let x = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.softmin(x, axis: 0) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.25, + minimum: 0.055289801210165024, + maximum: 0.43014487624168396, + absoluteSum: 3.0, + positionChecksum: 1.5460936228434246, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.4214111864566803, 0.2248254418373108, 0.15415097773075104, + 0.19324195384979248, 0.33697277307510376, 0.18245795369148254, + ]), + tolerance: .float32) + } + } + + @Test("logSoftmax") + func test_logSoftmax() throws { + try withIntegrationState(seed: 36714) { + let x = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.logSoftmax(x) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -1.3329758644104004, + minimum: -2.9164342880249023, + maximum: -0.40699923038482666, + absoluteSum: 15.995709419250488, + positionChecksum: 8.133668263753256, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.9382022023200989, -2.536550998687744, -1.00897216796875, + -2.07830810546875, -1.2845144271850586, -0.8200601935386658, + ]), + tolerance: .float32) + } + } + + @Test("logSoftmax/axis") + func test_logSoftmax_axis() throws { + try withIntegrationState(seed: 19083) { + let x = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.logSoftmax(x, axis: 0) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -2.090636730194092, + minimum: -5.523003101348877, + maximum: -0.09577543288469315, + absoluteSum: 25.08763885498047, + positionChecksum: 13.576393127441406, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.9566798806190491, -3.1437976360321045, -1.5162262916564941, + -0.5962467789649963, -2.607832908630371, -0.09577543288469315, + ]), + tolerance: .float32) + } + } + + @Test("logSigmoid") + func test_logSigmoid() throws { + try withIntegrationState(seed: 13659) { + let x = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.logSigmoid(x) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.6867945194244385, + minimum: -1.328415870666504, + maximum: -0.23934398591518402, + absoluteSum: 8.241534233093262, + positionChecksum: 4.763319969177246, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.44154447317123413, -1.2001237869262695, -0.7010847926139832, + -0.5049616694450378, -0.43374401330947876, -1.30607008934021, + ]), + tolerance: .float32) + } + } + + @Test("gelu") + func test_gelu() throws { + try withIntegrationState(seed: 56993) { + let x = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.gelu(x) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.23518109321594238, + minimum: -0.15614983439445496, + maximum: 1.5465519428253174, + absoluteSum: 4.124262809753418, + positionChecksum: 2.6090475718180337, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.29401835799217224, -0.015457117930054665, 1.0918797254562378, + 0.12314192205667496, -0.02627631649374962, 1.5465519428253174, + ]), + tolerance: .float32) + } + } + + @Test("geluApproximate") + func test_geluApproximate() throws { + try withIntegrationState(seed: 59006) { + let x = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.geluApproximate(x) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.34252291917800903, + minimum: -0.14033935964107513, + maximum: 1.349687933921814, + absoluteSum: 4.6187591552734375, + positionChecksum: 2.154048442840576, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.14033935964107513, 1.349687933921814, -0.08664838969707489, + 0.11771057546138763, -0.02725459821522236, 0.6522086262702942, + ]), + tolerance: .float32) + } + } + + @Test("geluFastApproximate") + func test_geluFastApproximate() throws { + try withIntegrationState(seed: 45811) { + let x = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.geluFastApproximate(x) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.7524160742759705, + minimum: -0.15481270849704742, + maximum: 3.242194414138794, + absoluteSum: 9.772090911865234, + positionChecksum: 4.8879391352335615, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.75855553150177, -0.14043818414211273, 0.5213433504104614, + -0.15481270849704742, 0.8235732316970825, -0.07629833370447159, + ]), + tolerance: .float32) + } + } + + @Test("glu") + func test_glu() throws { + try withIntegrationState(seed: 49808) { + let x = MLXRandom.normal([4, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.glu(x) + expectSummary( + result, + ArraySummary( + shape: [4, 4], + dtype: .float32, + mean: -0.15694786608219147, + minimum: -1.4327794313430786, + maximum: 0.9524335861206055, + absoluteSum: 7.32803201675415, + positionChecksum: 4.967343807220459, + sampleIndices: [0, 3, 6, 9, 12, 15], + samples: [ + 0.06294449418783188, 0.1666058599948883, 0.3972969651222229, + -1.0184520483016968, -0.4588458836078644, -1.4327794313430786, + ]), + tolerance: .float32) + } + } + + @Test("glu/axis") + func test_glu_axis() throws { + try withIntegrationState(seed: 29911) { + let x = MLXRandom.normal([4, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.glu(x, axis: 0) + expectSummary( + result, + ArraySummary( + shape: [2, 8], + dtype: .float32, + mean: 0.322890043258667, + minimum: -0.396503210067749, + maximum: 1.6453715562820435, + absoluteSum: 7.251438617706299, + positionChecksum: 4.466485977172852, + sampleIndices: [0, 3, 6, 9, 12, 15], + samples: [ + 0.4837617576122284, 0.14164136350154877, -0.20633184909820557, + 0.618834376335144, -0.396503210067749, 0.431894451379776, + ]), + tolerance: .float32) + } + } + + @Test("step") + func test_step() throws { + try withIntegrationState(seed: 79720) { + let x = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.step(x) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .int32, + mean: 0.5, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 6.0, + positionChecksum: 3.6666666666666665, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [1.0, 0.0, 0.0, 0.0, 1.0, 1.0]), + tolerance: .exact) + } + } + + @Test("step/threshold") + func test_step_threshold() throws { + try withIntegrationState(seed: 47231) { + let x = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.step(x, threshold: 0.5) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .int32, + mean: 0.4166666865348816, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 5.0, + positionChecksum: 2.3333333333333335, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [1.0, 0.0, 0.0, 0.0, 0.0, 1.0]), + tolerance: .exact) + } + } + + @Test("hardSwish") + func test_hardSwish() throws { + try withIntegrationState(seed: 54778) { + let x = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.hardSwish(x) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.6743363738059998, + minimum: -0.36365342140197754, + maximum: 2.4174962043762207, + absoluteSum: 10.894828796386719, + positionChecksum: 5.196015357971191, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 1.9947257041931152, 1.401398777961731, 1.191110610961914, + -0.35835549235343933, 0.4512682259082794, -0.31963229179382324, + ]), + tolerance: .float32) + } + } + + @Test("hardTanH") + func test_hardTanH() throws { + try withIntegrationState(seed: 64404) { + let x = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.hardTanH(x) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.1564551740884781, + minimum: -1.0, + maximum: 1.0, + absoluteSum: 7.426859378814697, + positionChecksum: 3.8016363779703775, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -1.0, 1.0, -0.22063256800174713, -0.4806191623210907, 0.5961291193962097, + -1.0, + ]), + tolerance: .float32) + } + } + + @Test("hardTanH/range") + func test_hardTanH_range() throws { + try withIntegrationState(seed: 12317) { + let x = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.hardTanH(x, min: -0.5, max: 0.5) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.24953168630599976, + minimum: -0.5, + maximum: 0.5, + absoluteSum: 5.2974853515625, + positionChecksum: 2.741623560587565, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [-0.5, 0.5, 0.38924914598464966, 0.5, 0.4365149736404419, 0.5]), + tolerance: .float32) + } + } + + @Test("hardShrink") + func test_hardShrink() throws { + try withIntegrationState(seed: 98167) { + let x = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.hardShrink(x) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.34890827536582947, + minimum: -1.311409831047058, + maximum: 1.3795651197433472, + absoluteSum: 6.809719085693359, + positionChecksum: 3.9342784881591797, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [0.0, 0.0, 0.0, 1.3795651197433472, 0.0, 1.1961815357208252]), + tolerance: .float32) + } + } + + @Test("hardShrink/lambda") + func test_hardShrink_lambda() throws { + try withIntegrationState(seed: 17487) { + let x = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.hardShrink(x, lambda: 0.2) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.1098986268043518, + minimum: -0.9246480464935303, + maximum: 0.6220173239707947, + absoluteSum: 3.96665096282959, + positionChecksum: 1.78379487991333, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.9246480464935303, 0.4655166268348694, 0.0, 0.5469300150871277, 0.0, + 0.4056254029273987, + ]), + tolerance: .float32) + } + } + + @Test("prelu") + func test_prelu() throws { + try withIntegrationState(seed: 57294) { + let x = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let alpha = MLXRandom.uniform(low: 0.1, high: 0.5, [3], dtype: .float32) + let result = MLXNN.prelu(x, alpha: alpha) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.3869323432445526, + minimum: -0.9682585597038269, + maximum: 1.647214412689209, + absoluteSum: 9.230525970458984, + positionChecksum: 5.1564515431722, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 1.647214412689209, 1.027597427368164, -0.08731298893690109, + 0.215767964720726, -0.9682585597038269, 1.128369688987732, + ]), + tolerance: .float32) + } + } +} diff --git a/Tests/MLXIntegrationTests/Generated/GeneratedBinaryTests.swift b/Tests/MLXIntegrationTests/Generated/GeneratedBinaryTests.swift new file mode 100644 index 000000000..72978092f --- /dev/null +++ b/Tests/MLXIntegrationTests/Generated/GeneratedBinaryTests.swift @@ -0,0 +1,1462 @@ +// Copyright © 2026 Apple Inc. +// +// GENERATED by tools/integration_tests -- DO NOT EDIT. +// +// Values in this file were produced by python mlx and are compared against +// the Swift API. Regenerate with: +// +// python3 tools/integration_tests/generate.py +// +// python mlx: 0.32.2.dev20260910+1f8e74e3f +// vendored mlx (Cmlx): 0.32.2 +// generator revision: 2 +// cases: 61 + +import Foundation +import MLX +import Testing + +@Suite("generated: Binary") +struct GeneratedBinaryTests { + + @Test("add") + func test_add() throws { + try withIntegrationState(seed: 31173) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.add(a, b) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.6918743848800659, + minimum: -4.831423282623291, + maximum: 1.7133017778396606, + absoluteSum: 15.499208450317383, + positionChecksum: 8.161814371744791, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 1.7133017778396606, 1.263370156288147, -0.09468770027160645, + -2.1408917903900146, -1.1286695003509521, -0.5542188882827759, + ]), + tolerance: .float32) + } + } + + @Test("subtract") + func test_subtract() throws { + try withIntegrationState(seed: 70937) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.subtract(a, b) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.1385997086763382, + minimum: -3.5849361419677734, + maximum: 2.4345836639404297, + absoluteSum: 16.477638244628906, + positionChecksum: 10.1959228515625, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.21980679035186768, 2.4345836639404297, -1.0636004209518433, + 0.7404163479804993, -3.5849361419677734, 2.1673178672790527, + ]), + tolerance: .float32) + } + } + + @Test("multiply") + func test_multiply() throws { + try withIntegrationState(seed: 763) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.multiply(a, b) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.15953750908374786, + minimum: -0.6936066150665283, + maximum: 1.4158040285110474, + absoluteSum: 3.8533198833465576, + positionChecksum: 1.7182307243347168, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.3082238733768463, 0.13711540400981903, 1.4158040285110474, + -0.004460824187844992, 0.039156656712293625, -0.10085374861955643, + ]), + tolerance: .float32) + } + } + + @Test("divide") + func test_divide() throws { + try withIntegrationState(seed: 52989) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.divide(a, b) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 13.941320419311523, + minimum: -3.5286707878112793, + maximum: 129.2261199951172, + absoluteSum: 190.4403076171875, + positionChecksum: 118.93863932291667, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.3581501543521881, -2.9087166786193848, 0.7051337957382202, + -2.402435779571533, 0.1394377052783966, -0.26220136880874634, + ]), + tolerance: .float32) + } + } + + @Test("remainder") + func test_remainder() throws { + try withIntegrationState(seed: 90342) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.remainder(a, b) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.006593585014343262, + minimum: -0.6975749731063843, + maximum: 1.1502771377563477, + absoluteSum: 5.422736167907715, + positionChecksum: 3.33150323232015, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.02850806713104248, -0.6895581483840942, -0.6510635614395142, + 0.3228166699409485, 0.7109462022781372, 1.1502771377563477, + ]), + tolerance: .float32) + } + } + + @Test("maximum") + func test_maximum() throws { + try withIntegrationState(seed: 93290) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.maximum(a, b) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.325461745262146, + minimum: -0.5284785032272339, + maximum: 1.6095561981201172, + absoluteSum: 7.125810623168945, + positionChecksum: 3.293111801147461, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.08404814451932907, 1.6095561981201172, -0.43784892559051514, + -0.5284785032272339, 0.484352171421051, 0.12096934765577316, + ]), + tolerance: .float32) + } + } + + @Test("minimum") + func test_minimum() throws { + try withIntegrationState(seed: 90053) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.minimum(a, b) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.22598682343959808, + minimum: -1.4718194007873535, + maximum: 1.3531242609024048, + absoluteSum: 7.797120094299316, + positionChecksum: 3.343695322672526, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.8319122791290283, -1.4718194007873535, 0.052562080323696136, + -0.15580348670482635, -0.6957205533981323, 0.44752341508865356, + ]), + tolerance: .float32) + } + } + + @Test("logAddExp") + func test_logAddExp() throws { + try withIntegrationState(seed: 53415) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.logAddExp(a, b) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.949323832988739, + minimum: -0.1343468427658081, + maximum: 1.801002025604248, + absoluteSum: 11.660578727722168, + positionChecksum: 6.266410827636719, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 1.801002025604248, 0.9974111318588257, 1.0150268077850342, + 1.2767889499664307, 0.6966230273246765, 1.4964959621429443, + ]), + tolerance: .float32) + } + } + + @Test("atan2") + func test_atan2() throws { + try withIntegrationState(seed: 42823) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.atan2(a, b) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.7183784246444702, + minimum: -2.9523682594299316, + maximum: 2.919264316558838, + absoluteSum: 23.06583023071289, + positionChecksum: 13.557057698567709, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -2.835385322570801, 0.07711521536111832, -2.8996195793151855, + -1.5876164436340332, 2.8714325428009033, -1.5721657276153564, + ]), + tolerance: .float32) + } + } + + @Test("floorDivide") + func test_floorDivide() throws { + try withIntegrationState(seed: 33967) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.floorDivide(a, b) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 2.9166667461395264, + minimum: -3.0, + maximum: 37.0, + absoluteSum: 51.0, + positionChecksum: 28.416666666666668, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [0.0, 1.0, 0.0, -3.0, 1.0, -3.0]), + tolerance: .float32) + } + } + + @Test("pow") + func test_pow() throws { + try withIntegrationState(seed: 1911) { + let a = MLXRandom.uniform(low: 0.1, high: 2.0, [4, 3], dtype: .float32) + let b = MLXRandom.uniform(low: 0.1, high: 2.0, [4, 3], dtype: .float32) + let result = MLX.pow(a, b) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 1.0408190488815308, + minimum: 0.028537485748529434, + maximum: 2.14806866645813, + absoluteSum: 12.489828109741211, + positionChecksum: 6.593130747477214, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.673996090888977, 0.39860963821411133, 1.0955862998962402, + 1.98665189743042, 0.6895277500152588, 0.877358078956604, + ]), + tolerance: .float32) + } + } + + @Test("equal") + func test_equal() throws { + try withIntegrationState(seed: 99685) { + let a = MLXRandom.randInt(low: 0, high: 4, [4, 3], type: Int32.self) + let b = MLXRandom.randInt(low: 0, high: 4, [4, 3], type: Int32.self) + let result = MLX.equal(a, b) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .bool, + mean: 0.4166666865348816, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 5.0, + positionChecksum: 2.9166666666666665, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [0.0, 0.0, 1.0, 0.0, 1.0, 0.0]), + tolerance: .exact) + } + } + + @Test("notEqual") + func test_notEqual() throws { + try withIntegrationState(seed: 8645) { + let a = MLXRandom.randInt(low: 0, high: 4, [4, 3], type: Int32.self) + let b = MLXRandom.randInt(low: 0, high: 4, [4, 3], type: Int32.self) + let result = MLX.notEqual(a, b) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .bool, + mean: 0.75, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 9.0, + positionChecksum: 4.25, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [1.0, 1.0, 1.0, 0.0, 1.0, 0.0]), + tolerance: .exact) + } + } + + @Test("less") + func test_less() throws { + try withIntegrationState(seed: 85782) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.less(a, b) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .bool, + mean: 0.5, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 6.0, + positionChecksum: 3.9166666666666665, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [0.0, 1.0, 0.0, 0.0, 1.0, 1.0]), + tolerance: .exact) + } + } + + @Test("lessEqual") + func test_lessEqual() throws { + try withIntegrationState(seed: 77066) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.lessEqual(a, b) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .bool, + mean: 0.5833333730697632, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 7.0, + positionChecksum: 4.333333333333333, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [1.0, 0.0, 1.0, 1.0, 1.0, 1.0]), + tolerance: .exact) + } + } + + @Test("greater") + func test_greater() throws { + try withIntegrationState(seed: 90710) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.greater(a, b) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .bool, + mean: 0.4166666865348816, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 5.0, + positionChecksum: 2.9166666666666665, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [0.0, 0.0, 0.0, 1.0, 1.0, 0.0]), + tolerance: .exact) + } + } + + @Test("greaterEqual") + func test_greaterEqual() throws { + try withIntegrationState(seed: 85655) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.greaterEqual(a, b) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .bool, + mean: 0.5833333730697632, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 7.0, + positionChecksum: 3.6666666666666665, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [1.0, 1.0, 0.0, 1.0, 0.0, 1.0]), + tolerance: .exact) + } + } + + @Test("operator/add") + func test_operator_add() throws { + try withIntegrationState(seed: 77360) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = a + b + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.27254313230514526, + minimum: -1.8845726251602173, + maximum: 3.533946990966797, + absoluteSum: 11.555704116821289, + positionChecksum: 7.568548838297526, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -1.0339982509613037, 0.43266522884368896, 0.7669238448143005, + 0.01592520996928215, 1.4233887195587158, 3.533946990966797, + ]), + tolerance: .float32) + } + } + + @Test("operator/add/scalarRHS") + func test_operator_add_scalarRHS() throws { + try withIntegrationState(seed: 96149) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = a + 1.3 + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 1.0331014394760132, + minimum: -0.7769765853881836, + maximum: 2.1503243446350098, + absoluteSum: 13.951169967651367, + positionChecksum: 8.28534189860026, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 1.6562983989715576, -0.7769765853881836, 1.057753086090088, + 0.4848131537437439, 0.8572401404380798, 2.1503243446350098, + ]), + tolerance: .float32) + } + } + + @Test("operator/add/scalarLHS") + func test_operator_add_scalarLHS() throws { + try withIntegrationState(seed: 12303) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = 0.5 + a + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.53053218126297, + minimum: -0.9197744131088257, + maximum: 2.8648841381073, + absoluteSum: 10.549100875854492, + positionChecksum: 4.844537417093913, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.174974262714386, -0.5325756072998047, 1.936285376548767, + 0.10657274723052979, -0.9197744131088257, 0.23554858565330505, + ]), + tolerance: .float32) + } + } + + @Test("operator/subtract") + func test_operator_subtract() throws { + try withIntegrationState(seed: 60591) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = a - b + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.27403706312179565, + minimum: -0.7568876147270203, + maximum: 2.6599271297454834, + absoluteSum: 8.948841094970703, + positionChecksum: 3.8554064432779946, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.2547779083251953, 1.4729726314544678, 2.6599271297454834, + -0.31046876311302185, 0.3260926604270935, 0.31457972526550293, + ]), + tolerance: .float32) + } + } + + @Test("operator/subtract/scalarRHS") + func test_operator_subtract_scalarRHS() throws { + try withIntegrationState(seed: 28905) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = a - 1.3 + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.9529048800468445, + minimum: -2.7124485969543457, + maximum: 0.6620986461639404, + absoluteSum: 12.759056091308594, + positionChecksum: 8.024199803670248, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.5383754968643188, -1.0335915088653564, -1.7431178092956543, + -2.7124485969543457, -1.635654091835022, -0.5485337972640991, + ]), + tolerance: .float32) + } + } + + @Test("operator/subtract/scalarLHS") + func test_operator_subtract_scalarLHS() throws { + try withIntegrationState(seed: 24115) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = 0.5 - a + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.36531418561935425, + minimum: -1.047974944114685, + maximum: 2.1034936904907227, + absoluteSum: 9.299480438232422, + positionChecksum: 5.230242093404134, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.08025974035263062, 0.06461766362190247, -0.45031672716140747, + 0.8556313514709473, -0.4064660668373108, -1.047974944114685, + ]), + tolerance: .float32) + } + } + + @Test("operator/multiply") + func test_operator_multiply() throws { + try withIntegrationState(seed: 27085) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = a * b + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.4524042308330536, + minimum: -2.1970603466033936, + maximum: 0.457118958234787, + absoluteSum: 6.941291809082031, + positionChecksum: 3.4856224060058594, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -2.1970603466033936, 0.11475061625242233, 0.457118958234787, + -1.40084969997406, -0.4138683080673218, -1.1669623851776123, + ]), + tolerance: .float32) + } + } + + @Test("operator/multiply/scalarRHS") + func test_operator_multiply_scalarRHS() throws { + try withIntegrationState(seed: 63195) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = a * 1.3 + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.30748769640922546, + minimum: -2.6694865226745605, + maximum: 1.9966527223587036, + absoluteSum: 14.375575065612793, + positionChecksum: 7.712694803873698, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -2.6694865226745605, -1.180345058441162, -0.7085902690887451, + -0.6169244647026062, 1.6982543468475342, -1.4327366352081299, + ]), + tolerance: .float32) + } + } + + @Test("operator/multiply/scalarLHS") + func test_operator_multiply_scalarLHS() throws { + try withIntegrationState(seed: 35649) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = 0.5 * a + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.09928220510482788, + minimum: -0.9127927422523499, + maximum: 0.8308178186416626, + absoluteSum: 4.5851569175720215, + positionChecksum: 2.197216033935547, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.9127927422523499, 0.20200762152671814, -0.110417440533638, + -0.5224007964134216, 0.8308178186416626, 0.3921858072280884, + ]), + tolerance: .float32) + } + } + + @Test("operator/divide") + func test_operator_divide() throws { + try withIntegrationState(seed: 60155) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = a / b + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 1.8858397006988525, + minimum: -4.235711097717285, + maximum: 17.544374465942383, + absoluteSum: 42.191917419433594, + positionChecksum: 24.856480916341145, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.19931618869304657, -0.5708990693092346, 0.011940893717110157, + 0.45701560378074646, -0.8684737086296082, -4.235711097717285, + ]), + tolerance: .float32) + } + } + + @Test("operator/divide/scalarRHS") + func test_operator_divide_scalarRHS() throws { + try withIntegrationState(seed: 40534) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = a / 1.3 + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.1341741383075714, + minimum: -0.765953540802002, + maximum: 1.2086530923843384, + absoluteSum: 4.993619918823242, + positionChecksum: 2.275047938028971, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.0909951850771904, 1.2086530923843384, 0.4493958055973053, + 0.1280558556318283, 0.15634208917617798, -0.08264032751321793, + ]), + tolerance: .float32) + } + } + + @Test("operator/divide/scalarLHS") + func test_operator_divide_scalarLHS() throws { + try withIntegrationState(seed: 24204) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = 0.5 / a + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -10.988043785095215, + minimum: -135.88998413085938, + maximum: 4.149685382843018, + absoluteSum: 146.1427764892578, + positionChecksum: 62.945159912109375, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.8811553716659546, 0.2845892608165741, -135.88998413085938, + 4.149685382843018, 0.4647601842880249, -0.7214249968528748, + ]), + tolerance: .float32) + } + } + + @Test("operator/remainder") + func test_operator_remainder() throws { + try withIntegrationState(seed: 71211) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = a % b + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.05340499430894852, + minimum: -0.9634010791778564, + maximum: 1.3972878456115723, + absoluteSum: 3.9187231063842773, + positionChecksum: 1.8389933904012044, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.9634010791778564, 0.14139050245285034, -0.22529059648513794, + 0.050104208290576935, 0.3663478493690491, 0.0027291178703308105, + ]), + tolerance: .float32) + } + } + + @Test("operator/remainder/scalarRHS") + func test_operator_remainder_scalarRHS() throws { + try withIntegrationState(seed: 35295) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = a % 1.3 + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.5924494862556458, + minimum: 0.051908016204833984, + maximum: 1.2263563871383667, + absoluteSum: 7.10939359664917, + positionChecksum: 4.6820723215738935, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.293434739112854, 0.06582850962877274, 1.0644726753234863, + 0.3147084712982178, 1.1621785163879395, 0.9618409872055054, + ]), + tolerance: .float32) + } + } + + @Test("operator/remainder/scalarLHS") + func test_operator_remainder_scalarLHS() throws { + try withIntegrationState(seed: 2437) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = 0.5 % a + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.20201009511947632, + minimum: -0.4174903631210327, + maximum: 0.5, + absoluteSum: 3.881199598312378, + positionChecksum: 2.143406550089518, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.07111084461212158, 0.5, -0.06144183874130249, 0.5, 0.5, + -0.4174903631210327, + ]), + tolerance: .float32) + } + } + + @Test("operator/pow") + func test_operator_pow() throws { + try withIntegrationState(seed: 55330) { + let a = MLXRandom.uniform(low: 0.1, high: 2.0, [4, 3], dtype: .float32) + let b = MLXRandom.uniform(low: 0.1, high: 2.0, [4, 3], dtype: .float32) + let result = a ** b + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.8067885637283325, + minimum: 0.041118938475847244, + maximum: 1.7501370906829834, + absoluteSum: 9.681462287902832, + positionChecksum: 4.114236831665039, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 1.7501370906829834, 1.371971845626831, 0.3343576490879059, + 0.29617589712142944, 0.8186249136924744, 0.041118938475847244, + ]), + tolerance: .float32) + } + } + + @Test("operator/negate") + func test_operator_negate() throws { + try withIntegrationState(seed: 27407) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = -a + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.36295005679130554, + minimum: -1.8532614707946777, + maximum: 0.9766664505004883, + absoluteSum: 9.706311225891113, + positionChecksum: 6.226847966512044, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.22385217249393463, 0.6678629517555237, -1.2412484884262085, + -1.682862401008606, -0.23273582756519318, -1.8532614707946777, + ]), + tolerance: .float32) + } + } + + @Test("operator/equal") + func test_operator_equal() throws { + try withIntegrationState(seed: 5151) { + let a = MLXRandom.randInt(low: 0, high: 4, [4, 3], type: Int32.self) + let b = MLXRandom.randInt(low: 0, high: 4, [4, 3], type: Int32.self) + let result = a .== b + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .bool, + mean: 0.3333333432674408, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 4.0, + positionChecksum: 1.8333333333333333, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [1.0, 0.0, 0.0, 0.0, 1.0, 0.0]), + tolerance: .exact) + } + } + + @Test("operator/notEqual") + func test_operator_notEqual() throws { + try withIntegrationState(seed: 22291) { + let a = MLXRandom.randInt(low: 0, high: 4, [4, 3], type: Int32.self) + let b = MLXRandom.randInt(low: 0, high: 4, [4, 3], type: Int32.self) + let result = a .!= b + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .bool, + mean: 0.6666666865348816, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 8.0, + positionChecksum: 3.8333333333333335, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [0.0, 1.0, 1.0, 1.0, 0.0, 0.0]), + tolerance: .exact) + } + } + + @Test("operator/less") + func test_operator_less() throws { + try withIntegrationState(seed: 59592) { + let a = MLXRandom.randInt(low: 0, high: 4, [4, 3], type: Int32.self) + let b = MLXRandom.randInt(low: 0, high: 4, [4, 3], type: Int32.self) + let result = a .< b + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .bool, + mean: 0.25, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 3.0, + positionChecksum: 2.25, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [0.0, 0.0, 0.0, 0.0, 1.0, 0.0]), + tolerance: .exact) + } + } + + @Test("operator/lessEqual") + func test_operator_lessEqual() throws { + try withIntegrationState(seed: 52583) { + let a = MLXRandom.randInt(low: 0, high: 4, [4, 3], type: Int32.self) + let b = MLXRandom.randInt(low: 0, high: 4, [4, 3], type: Int32.self) + let result = a .<= b + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .bool, + mean: 0.5833333730697632, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 7.0, + positionChecksum: 3.75, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [0.0, 1.0, 0.0, 1.0, 1.0, 0.0]), + tolerance: .exact) + } + } + + @Test("operator/greater") + func test_operator_greater() throws { + try withIntegrationState(seed: 67466) { + let a = MLXRandom.randInt(low: 0, high: 4, [4, 3], type: Int32.self) + let b = MLXRandom.randInt(low: 0, high: 4, [4, 3], type: Int32.self) + let result = a .> b + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .bool, + mean: 0.4166666865348816, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 5.0, + positionChecksum: 3.5, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [0.0, 0.0, 0.0, 0.0, 0.0, 1.0]), + tolerance: .exact) + } + } + + @Test("operator/greaterEqual") + func test_operator_greaterEqual() throws { + try withIntegrationState(seed: 84422) { + let a = MLXRandom.randInt(low: 0, high: 4, [4, 3], type: Int32.self) + let b = MLXRandom.randInt(low: 0, high: 4, [4, 3], type: Int32.self) + let result = a .>= b + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .bool, + mean: 0.5, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 6.0, + positionChecksum: 3.1666666666666665, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [0.0, 1.0, 0.0, 1.0, 1.0, 0.0]), + tolerance: .exact) + } + } + + @Test("logicalAnd") + func test_logicalAnd() throws { + try withIntegrationState(seed: 77847) { + let a = MLXRandom.bernoulli(0.5, [4, 3]) + let b = MLXRandom.bernoulli(0.5, [4, 3]) + let result = MLX.logicalAnd(a, b) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .bool, + mean: 0.3333333432674408, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 4.0, + positionChecksum: 1.5833333333333333, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [1.0, 0.0, 1.0, 0.0, 0.0, 0.0]), + tolerance: .exact) + } + } + + @Test("logicalOr") + func test_logicalOr() throws { + try withIntegrationState(seed: 96553) { + let a = MLXRandom.bernoulli(0.5, [4, 3]) + let b = MLXRandom.bernoulli(0.5, [4, 3]) + let result = MLX.logicalOr(a, b) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .bool, + mean: 0.75, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 9.0, + positionChecksum: 4.75, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [1.0, 1.0, 1.0, 1.0, 0.0, 1.0]), + tolerance: .exact) + } + } + + @Test("logicalXor") + func test_logicalXor() throws { + try withIntegrationState(seed: 20568) { + let a = MLXRandom.bernoulli(0.5, [4, 3]) + let b = MLXRandom.bernoulli(0.5, [4, 3]) + let result = MLX.logicalXor(a, b) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .bool, + mean: 0.5, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 6.0, + positionChecksum: 3.3333333333333335, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [0.0, 0.0, 1.0, 0.0, 0.0, 0.0]), + tolerance: .exact) + } + } + + @Test("which") + func test_which() throws { + try withIntegrationState(seed: 39197) { + let mask = MLXRandom.bernoulli(0.5, [4, 3]) + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.which(mask, a, b) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.08233775198459625, + minimum: -0.9875338673591614, + maximum: 2.0078234672546387, + absoluteSum: 10.938934326171875, + positionChecksum: 5.835131963094075, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.5632016062736511, 1.5790736675262451, 0.684048056602478, + 2.0078234672546387, 1.3224455118179321, -0.5925998687744141, + ]), + tolerance: .float32) + } + } + + @Test("bitwiseAnd") + func test_bitwiseAnd() throws { + try withIntegrationState(seed: 18363) { + let a = MLXRandom.randInt(low: 0, high: 16, [4, 3], type: Int32.self) + let b = MLXRandom.randInt(low: 0, high: 16, [4, 3], type: Int32.self) + let result = MLX.bitwiseAnd(a, b) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .int32, + mean: 3.8333334922790527, + minimum: 0.0, + maximum: 12.0, + absoluteSum: 46.0, + positionChecksum: 25.333333333333332, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [8.0, 5.0, 0.0, 10.0, 0.0, 12.0]), + tolerance: .exact) + } + } + + @Test("bitwiseOr") + func test_bitwiseOr() throws { + try withIntegrationState(seed: 78878) { + let a = MLXRandom.randInt(low: 0, high: 16, [4, 3], type: Int32.self) + let b = MLXRandom.randInt(low: 0, high: 16, [4, 3], type: Int32.self) + let result = MLX.bitwiseOr(a, b) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .int32, + mean: 10.166666984558105, + minimum: 6.0, + maximum: 15.0, + absoluteSum: 122.0, + positionChecksum: 63.166666666666664, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [11.0, 11.0, 6.0, 14.0, 8.0, 6.0]), + tolerance: .exact) + } + } + + @Test("bitwiseXOr") + func test_bitwiseXOr() throws { + try withIntegrationState(seed: 66230) { + let a = MLXRandom.randInt(low: 0, high: 16, [4, 3], type: Int32.self) + let b = MLXRandom.randInt(low: 0, high: 16, [4, 3], type: Int32.self) + let result = MLX.bitwiseXOr(a, b) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .int32, + mean: 7.0, + minimum: 0.0, + maximum: 14.0, + absoluteSum: 84.0, + positionChecksum: 39.916666666666664, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [12.0, 10.0, 11.0, 4.0, 4.0, 9.0]), + tolerance: .exact) + } + } + + @Test("bitwiseInvert") + func test_bitwiseInvert() throws { + try withIntegrationState(seed: 53813) { + let a = MLXRandom.randInt(low: 0, high: 16, [4, 3], type: Int32.self) + let result = MLX.bitwiseInvert(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .int32, + mean: -7.9166669845581055, + minimum: -16.0, + maximum: -1.0, + absoluteSum: 95.0, + positionChecksum: 61.5, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [-8.0, -1.0, -4.0, -10.0, -16.0, -16.0]), + tolerance: .exact) + } + } + + @Test("leftShift") + func test_leftShift() throws { + try withIntegrationState(seed: 83785) { + let a = MLXRandom.randInt(low: 0, high: 16, [4, 3], type: Int32.self) + let b = MLXRandom.randInt(low: 0, high: 4, [4, 3], type: Int32.self) + let result = MLX.leftShift(a, b) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .int32, + mean: 32.41666793823242, + minimum: 0.0, + maximum: 120.0, + absoluteSum: 389.0, + positionChecksum: 275.0833333333333, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [16.0, 40.0, 9.0, 0.0, 16.0, 120.0]), + tolerance: .exact) + } + } + + @Test("rightShift") + func test_rightShift() throws { + try withIntegrationState(seed: 20116) { + let a = MLXRandom.randInt(low: 0, high: 1024, [4, 3], type: Int32.self) + let b = MLXRandom.randInt(low: 0, high: 4, [4, 3], type: Int32.self) + let result = MLX.rightShift(a, b) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .int32, + mean: 284.5, + minimum: 24.0, + maximum: 856.0, + absoluteSum: 3414.0, + positionChecksum: 1920.3333333333333, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [201.0, 287.0, 512.0, 51.0, 224.0, 302.0]), + tolerance: .exact) + } + } + + @Test("matmul") + func test_matmul() throws { + try withIntegrationState(seed: 69847) { + let a = MLXRandom.normal([10, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([8, 13], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.matmul(a, b) + expectSummary( + result, + ArraySummary( + shape: [10, 13], + dtype: .float32, + mean: 0.013026426546275616, + minimum: -7.829680442810059, + maximum: 5.517763614654541, + absoluteSum: 251.1970977783203, + positionChecksum: 140.54157151442308, + sampleIndices: [0, 26, 52, 77, 103, 129], + samples: [ + -1.5777220726013184, -2.814626455307007, -0.9406680464744568, + 1.512544870376587, 1.1238691806793213, 2.8315815925598145, + ]), + tolerance: .float32) + } + } + + @Test("matmul/batched") + func test_matmul_batched() throws { + try withIntegrationState(seed: 58020) { + let a = MLXRandom.normal([2, 4, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([2, 8, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.matmul(a, b) + expectSummary( + result, + ArraySummary( + shape: [2, 4, 3], + dtype: .float32, + mean: 0.560451865196228, + minimum: -4.739752292633057, + maximum: 9.252123832702637, + absoluteSum: 51.17353057861328, + positionChecksum: 27.05833689371745, + sampleIndices: [0, 5, 9, 14, 18, 23], + samples: [ + 0.6251500844955444, 0.8671976327896118, 9.252123832702637, + 0.1872236281633377, 1.1555585861206055, -2.105651617050171, + ]), + tolerance: .float32) + } + } + + @Test("inner") + func test_inner() throws { + try withIntegrationState(seed: 36443) { + let a = MLXRandom.normal([12], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([12], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.inner(a, b) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: -2.8012592792510986, + minimum: -2.8012592792510986, + maximum: -2.8012592792510986, + absoluteSum: 2.8012592792510986, + positionChecksum: 2.8012592792510986, + sampleIndices: [0], + samples: [-2.8012592792510986]), + tolerance: .float32) + } + } + + @Test("outer") + func test_outer() throws { + try withIntegrationState(seed: 99139) { + let a = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([4], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.outer(a, b) + expectSummary( + result, + ArraySummary( + shape: [5, 4], + dtype: .float32, + mean: 0.2968282103538513, + minimum: -1.7715340852737427, + maximum: 3.274460792541504, + absoluteSum: 17.70854377746582, + positionChecksum: 9.23790283203125, + sampleIndices: [0, 4, 8, 11, 15, 19], + samples: [ + -0.5336867570877075, -1.4807971715927124, -0.4454834759235382, + 0.15995460748672485, 0.636084258556366, -0.22449640929698944, + ]), + tolerance: .float32) + } + } + + @Test("vecdot") + func test_vecdot() throws { + try withIntegrationState(seed: 57547) { + let a = MLXRandom.normal([4, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([4, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.vecdot(a, b) + expectSummary( + result, + ArraySummary( + shape: [4], + dtype: .float32, + mean: 0.39991307258605957, + minimum: 0.07599055767059326, + maximum: 1.119821548461914, + absoluteSum: 1.5996522903442383, + positionChecksum: 1.2174617052078247, + sampleIndices: [0, 1, 2, 3], + samples: [ + 0.08565317094326019, 0.07599055767059326, 1.119821548461914, + 0.31818699836730957, + ]), + tolerance: .float32) + } + } + + @Test("tensordot") + func test_tensordot() throws { + try withIntegrationState(seed: 41209) { + let a = MLXRandom.normal([2, 3, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([4, 3, 2], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.tensordot(a, b, axes: ([1, 2], [1, 0])) + expectSummary( + result, + ArraySummary( + shape: [2, 2], + dtype: .float32, + mean: 2.566420793533325, + minimum: -2.453749656677246, + maximum: 8.781932830810547, + absoluteSum: 15.173182487487793, + positionChecksum: 11.776962280273438, + sampleIndices: [0, 1, 2, 3], + samples: [ + -2.453749656677246, 2.286130666732788, 1.6513692140579224, + 8.781932830810547, + ]), + tolerance: .float32) + } + } + + @Test("addMM") + func test_addMM() throws { + try withIntegrationState(seed: 78623) { + let c = MLXRandom.normal([4, 5], dtype: .float32, loc: 0.0, scale: 1.0) + let a = MLXRandom.normal([4, 6], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([6, 5], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.addMM(c, a, b, alpha: 0.5, beta: 2.0) + expectSummary( + result, + ArraySummary( + shape: [4, 5], + dtype: .float32, + mean: 0.18949465453624725, + minimum: -2.853945732116699, + maximum: 4.330568313598633, + absoluteSum: 26.563777923583984, + positionChecksum: 16.041049194335937, + sampleIndices: [0, 4, 8, 11, 15, 19], + samples: [ + -1.6629900932312012, 0.6326870918273926, -0.28035566210746765, + 4.330568313598633, 0.5440598726272583, -1.714925765991211, + ]), + tolerance: .float32) + } + } + + @Test("kron") + func test_kron() throws { + try withIntegrationState(seed: 42558) { + let a = MLXRandom.normal([2, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([3, 2], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.kron(a, b) + expectSummary( + result, + ArraySummary( + shape: [6, 6], + dtype: .float32, + mean: -0.23324330151081085, + minimum: -1.7474167346954346, + maximum: 1.0404057502746582, + absoluteSum: 22.564327239990234, + positionChecksum: 11.563710530598959, + sampleIndices: [0, 7, 14, 21, 28, 35], + samples: [ + 0.4093790054321289, 0.27812692523002625, -1.7474167346954346, + -0.03940131887793541, -0.668895959854126, 0.8408368229866028, + ]), + tolerance: .float32) + } + } + + @Test("norm/frobeniusKeepDims") + func test_norm_frobeniusKeepDims() throws { + try withIntegrationState(seed: 80992) { + let a = MLXRandom.normal([3, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.norm(a, keepDims: true) + expectSummary( + result, + ArraySummary( + shape: [1, 1], + dtype: .float32, + mean: 3.2592105865478516, + minimum: 3.2592105865478516, + maximum: 3.2592105865478516, + absoluteSum: 3.2592105865478516, + positionChecksum: 3.2592105865478516, + sampleIndices: [0], + samples: [3.2592105865478516]), + tolerance: .float32) + } + } + + @Test("addMM/newtonSchulzForm") + func test_addMM_newtonSchulzForm() throws { + try withIntegrationState(seed: 34985) { + let a = MLXRandom.normal([3, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.addMM(-4.7750 * a, a, a, alpha: 2.0315, beta: 1.0) + expectSummary( + result, + ArraySummary( + shape: [3, 3], + dtype: .float32, + mean: -0.7809702157974243, + minimum: -4.185337066650391, + maximum: 4.572365760803223, + absoluteSum: 24.614940643310547, + positionChecksum: 14.121938069661459, + sampleIndices: [0, 2, 3, 5, 6, 8], + samples: [ + -0.531244158744812, 4.220737934112549, -2.525264263153076, + -0.26388388872146606, -1.9338984489440918, -4.185337066650391, + ]), + tolerance: .float32) + } + } + + @Test("einsum") + func test_einsum() throws { + try withIntegrationState(seed: 27180) { + let a = MLXRandom.normal([4, 5], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([5, 6], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.einsum("ij,jk->ik", a, b) + expectSummary( + result, + ArraySummary( + shape: [4, 6], + dtype: .float32, + mean: 0.5131218433380127, + minimum: -6.676657199859619, + maximum: 8.000497817993164, + absoluteSum: 71.23182678222656, + positionChecksum: 36.53941090901693, + sampleIndices: [0, 5, 9, 14, 18, 23], + samples: [ + 2.163822889328003, -0.3911140263080597, 1.3212900161743164, + 8.000497817993164, -1.8788937330245972, -1.330705165863037, + ]), + tolerance: .float32) + } + } +} diff --git a/Tests/MLXIntegrationTests/Generated/GeneratedConvolutionTests.swift b/Tests/MLXIntegrationTests/Generated/GeneratedConvolutionTests.swift new file mode 100644 index 000000000..078667b15 --- /dev/null +++ b/Tests/MLXIntegrationTests/Generated/GeneratedConvolutionTests.swift @@ -0,0 +1,422 @@ +// Copyright © 2026 Apple Inc. +// +// GENERATED by tools/integration_tests -- DO NOT EDIT. +// +// Values in this file were produced by python mlx and are compared against +// the Swift API. Regenerate with: +// +// python3 tools/integration_tests/generate.py +// +// python mlx: 0.32.2.dev20260910+1f8e74e3f +// vendored mlx (Cmlx): 0.32.2 +// generator revision: 2 +// cases: 16 + +import Foundation +import MLX +import Testing + +@Suite("generated: Convolution") +struct GeneratedConvolutionTests { + + @Test("convolve/full") + func test_convolve_full() throws { + try withIntegrationState(seed: 4719) { + let a = MLXRandom.normal([20], dtype: .float32, loc: 0.0, scale: 1.0) + let v = MLXRandom.normal([4], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.convolve(a, v) + expectSummary( + result, + ArraySummary( + shape: [23], + dtype: .float32, + mean: -0.010141642764210701, + minimum: -2.3259332180023193, + maximum: 1.7424554824829102, + absoluteSum: 19.981769561767578, + positionChecksum: 10.247398044752037, + sampleIndices: [0, 4, 9, 13, 18, 22], + samples: [ + -0.008841410279273987, 1.1080988645553589, 0.9621633291244507, + 1.7424554824829102, -0.30217692255973816, 0.002387824235484004, + ]), + tolerance: .float32) + } + } + + @Test("convolve/same") + func test_convolve_same() throws { + try withIntegrationState(seed: 68267) { + let a = MLXRandom.normal([20], dtype: .float32, loc: 0.0, scale: 1.0) + let v = MLXRandom.normal([4], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.convolve(a, v, mode: .same) + expectSummary( + result, + ArraySummary( + shape: [20], + dtype: .float32, + mean: 0.4013185203075409, + minimum: -3.1303353309631348, + maximum: 4.640427112579346, + absoluteSum: 29.873620986938477, + positionChecksum: 15.805455017089844, + sampleIndices: [0, 4, 8, 11, 15, 19], + samples: [ + 0.08655143529176712, 0.08424151688814163, -3.1303353309631348, + 2.3075802326202393, -1.9558295011520386, -0.9330341219902039, + ]), + tolerance: .float32) + } + } + + @Test("convolve/valid") + func test_convolve_valid() throws { + try withIntegrationState(seed: 33328) { + let a = MLXRandom.normal([20], dtype: .float32, loc: 0.0, scale: 1.0) + let v = MLXRandom.normal([4], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.convolve(a, v, mode: .valid) + expectSummary( + result, + ArraySummary( + shape: [17], + dtype: .float32, + mean: 0.18561282753944397, + minimum: -3.5847108364105225, + maximum: 4.357660293579102, + absoluteSum: 23.426000595092773, + positionChecksum: 14.26570488424862, + sampleIndices: [0, 3, 6, 10, 13, 16], + samples: [ + 1.5556355714797974, 0.9753752946853638, 0.8864786624908447, + 0.2802148759365082, 2.622657299041748, 2.514362677175086e-05, + ]), + tolerance: .float32) + } + } + + @Test("conv1d") + func test_conv1d() throws { + try withIntegrationState(seed: 33875) { + let input = MLXRandom.normal([2, 10, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let weight = MLXRandom.normal([3, 3, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.conv1d(input, weight) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 3], + dtype: .float32, + mean: -0.23364683985710144, + minimum: -7.760387897491455, + maximum: 6.223398685455322, + absoluteSum: 115.83242797851562, + positionChecksum: 60.60389709472656, + sampleIndices: [0, 9, 19, 28, 38, 47], + samples: [ + -2.8744194507598877, -2.188898801803589, -2.1922268867492676, + 3.848215341567993, 3.4207470417022705, -3.5831971168518066, + ]), + tolerance: .float32) + } + } + + @Test("conv1d/stridePadding") + func test_conv1d_stridePadding() throws { + try withIntegrationState(seed: 94979) { + let input = MLXRandom.normal([2, 10, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let weight = MLXRandom.normal([3, 3, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.conv1d(input, weight, stride: 2, padding: 1) + expectSummary( + result, + ArraySummary( + shape: [2, 5, 3], + dtype: .float32, + mean: 0.45748910307884216, + minimum: -7.211019039154053, + maximum: 8.565747261047363, + absoluteSum: 72.53246307373047, + positionChecksum: 44.807047526041664, + sampleIndices: [0, 6, 12, 17, 23, 29], + samples: [ + 0.2506880760192871, -1.96531081199646, -0.8448853492736816, + -1.1593079566955566, -1.14354407787323, 4.385194778442383, + ]), + tolerance: .float32) + } + } + + @Test("conv1d/dilation") + func test_conv1d_dilation() throws { + try withIntegrationState(seed: 17243) { + let input = MLXRandom.normal([2, 10, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let weight = MLXRandom.normal([3, 3, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.conv1d(input, weight, dilation: 2) + expectSummary( + result, + ArraySummary( + shape: [2, 6, 3], + dtype: .float32, + mean: -0.8824627995491028, + minimum: -6.105772972106934, + maximum: 6.817114353179932, + absoluteSum: 88.25929260253906, + positionChecksum: 49.157989501953125, + sampleIndices: [0, 7, 14, 21, 28, 35], + samples: [ + -1.936726450920105, -1.7665845155715942, -0.4554101824760437, + -2.9440197944641113, 0.016120949760079384, 5.547804355621338, + ]), + tolerance: .float32) + } + } + + @Test("conv1d/groups") + func test_conv1d_groups() throws { + try withIntegrationState(seed: 44073) { + let input = MLXRandom.normal([2, 10, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let weight = MLXRandom.normal([4, 3, 2], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.conv1d(input, weight, groups: 2) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 4], + dtype: .float32, + mean: -0.29602718353271484, + minimum: -6.234485149383545, + maximum: 7.6336469650268555, + absoluteSum: 137.21255493164062, + positionChecksum: 70.91708374023438, + sampleIndices: [0, 13, 25, 38, 50, 63], + samples: [ + -0.6050371527671814, 2.1364376544952393, 1.7200886011123657, + -1.4107182025909424, -2.151890754699707, -5.122685432434082, + ]), + tolerance: .float32) + } + } + + @Test("conv2d") + func test_conv2d() throws { + try withIntegrationState(seed: 54064) { + let input = MLXRandom.normal([2, 8, 8, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let weight = MLXRandom.normal([4, 3, 3, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.conv2d(input, weight) + expectSummary( + result, + ArraySummary( + shape: [2, 6, 6, 4], + dtype: .float32, + mean: 0.11624614149332047, + minimum: -11.582426071166992, + maximum: 12.33521842956543, + absoluteSum: 929.9988403320312, + positionChecksum: 462.14453125, + sampleIndices: [0, 57, 115, 172, 230, 287], + samples: [ + -2.1473069190979004, 2.8220767974853516, 4.751214027404785, + -3.128023624420166, -0.647574782371521, -7.356082916259766, + ]), + tolerance: .float32) + } + } + + @Test("conv2d/stridePadding") + func test_conv2d_stridePadding() throws { + try withIntegrationState(seed: 57937) { + let input = MLXRandom.normal([2, 8, 8, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let weight = MLXRandom.normal([4, 3, 3, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.conv2d(input, weight, stride: [2, 1], padding: [1, 0]) + expectSummary( + result, + ArraySummary( + shape: [2, 4, 6, 4], + dtype: .float32, + mean: 0.16354301571846008, + minimum: -11.90718936920166, + maximum: 14.7648286819458, + absoluteSum: 820.2205200195312, + positionChecksum: 408.4796549479167, + sampleIndices: [0, 38, 76, 115, 153, 191], + samples: [ + -6.419956207275391, 3.057616710662842, -3.531296968460083, + 7.2112555503845215, -1.983046531677246, 2.6352462768554688, + ]), + tolerance: .float32) + } + } + + @Test("conv3d") + func test_conv3d() throws { + try withIntegrationState(seed: 25745) { + let input = MLXRandom.normal([1, 4, 6, 6, 2], dtype: .float32, loc: 0.0, scale: 1.0) + let weight = MLXRandom.normal([3, 2, 3, 3, 2], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.conv3d(input, weight) + expectSummary( + result, + ArraySummary( + shape: [1, 3, 4, 4, 3], + dtype: .float32, + mean: -0.19438402354717255, + minimum: -21.18064308166504, + maximum: 13.98067569732666, + absoluteSum: 585.3264770507812, + positionChecksum: 298.1985134548611, + sampleIndices: [0, 29, 57, 86, 114, 143], + samples: [ + -1.989408254623413, -4.56206750869751, -4.980064392089844, + 0.4489351511001587, 1.2831811904907227, -0.33633241057395935, + ]), + tolerance: .float32) + } + } + + @Test("convTransposed1d") + func test_convTransposed1d() throws { + try withIntegrationState(seed: 56758) { + let input = MLXRandom.normal([2, 8, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let weight = MLXRandom.normal([3, 3, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.convTransposed1d(input, weight) + expectSummary( + result, + ArraySummary( + shape: [2, 10, 3], + dtype: .float32, + mean: -0.08723472058773041, + minimum: -7.693572521209717, + maximum: 7.150848865509033, + absoluteSum: 127.65473937988281, + positionChecksum: 62.71107177734375, + sampleIndices: [0, 12, 24, 35, 47, 59], + samples: [ + 0.22783879935741425, -2.8541452884674072, 1.6900525093078613, + -1.3801987171173096, -0.30918294191360474, 0.5386923551559448, + ]), + tolerance: .float32) + } + } + + @Test("convTransposed1d/stride") + func test_convTransposed1d_stride() throws { + try withIntegrationState(seed: 7340) { + let input = MLXRandom.normal([2, 8, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let weight = MLXRandom.normal([3, 3, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.convTransposed1d( + input, weight, stride: 2, padding: 1, outputPadding: 1) + expectSummary( + result, + ArraySummary( + shape: [2, 16, 3], + dtype: .float32, + mean: 0.20248651504516602, + minimum: -6.641228199005127, + maximum: 7.307250499725342, + absoluteSum: 178.95211791992188, + positionChecksum: 95.32716878255208, + sampleIndices: [0, 19, 38, 57, 76, 95], + samples: [ + 2.1703476905822754, 0.1148005798459053, 1.0188913345336914, + 1.8733576536178589, -2.3983073234558105, 1.2071207761764526, + ]), + tolerance: .float32) + } + } + + @Test("convTransposed2d") + func test_convTransposed2d() throws { + try withIntegrationState(seed: 60757) { + let input = MLXRandom.normal([2, 6, 6, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let weight = MLXRandom.normal([4, 3, 3, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.convTransposed2d(input, weight) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 8, 4], + dtype: .float32, + mean: -0.15208807587623596, + minimum: -13.422320365905762, + maximum: 12.339730262756348, + absoluteSum: 1462.8289794921875, + positionChecksum: 722.030029296875, + sampleIndices: [0, 102, 204, 307, 409, 511], + samples: [ + -1.8898229598999023, 6.999209403991699, 3.9277143478393555, + -0.8446962237358093, 0.1216912716627121, -0.5072753429412842, + ]), + tolerance: .float32) + } + } + + @Test("convTransposed3d") + func test_convTransposed3d() throws { + try withIntegrationState(seed: 37364) { + let input = MLXRandom.normal([1, 4, 4, 4, 2], dtype: .float32, loc: 0.0, scale: 1.0) + let weight = MLXRandom.normal([3, 2, 2, 2, 2], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.convTransposed3d(input, weight) + expectSummary( + result, + ArraySummary( + shape: [1, 5, 5, 5, 3], + dtype: .float32, + mean: 0.038671448826789856, + minimum: -10.21535873413086, + maximum: 12.558921813964844, + absoluteSum: 831.7537841796875, + positionChecksum: 421.350625, + sampleIndices: [0, 75, 150, 224, 299, 374], + samples: [ + 1.1298370361328125, -0.9043198823928833, 1.8892734050750732, + -0.6862713098526001, 0.47907447814941406, -0.20345965027809143, + ]), + tolerance: .float32) + } + } + + @Test("convGeneral") + func test_convGeneral() throws { + try withIntegrationState(seed: 91113) { + let input = MLXRandom.normal([2, 8, 8, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let weight = MLXRandom.normal([4, 3, 3, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.convGeneral(input, weight, strides: 2, padding: 1) + expectSummary( + result, + ArraySummary( + shape: [2, 4, 4, 4], + dtype: .float32, + mean: 0.563105046749115, + minimum: -9.814929008483887, + maximum: 15.443587303161621, + absoluteSum: 483.46710205078125, + positionChecksum: 256.685791015625, + sampleIndices: [0, 25, 51, 76, 102, 127], + samples: [ + 4.08807897567749, 0.49885979294776917, 5.748270511627197, + -0.7634090781211853, -2.1041526794433594, 2.326464891433716, + ]), + tolerance: .float32) + } + } + + @Test("convGeneral/flip") + func test_convGeneral_flip() throws { + try withIntegrationState(seed: 61813) { + let input = MLXRandom.normal([2, 8, 8, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let weight = MLXRandom.normal([4, 3, 3, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.convGeneral(input, weight, kernelDilation: 2, flip: true) + expectSummary( + result, + ArraySummary( + shape: [2, 4, 4, 4], + dtype: .float32, + mean: 1.1288611888885498, + minimum: -10.111289978027344, + maximum: 13.187017440795898, + absoluteSum: 500.216796875, + positionChecksum: 253.00144958496094, + sampleIndices: [0, 25, 51, 76, 102, 127], + samples: [ + -0.9194361567497253, 0.8226830959320068, 13.187017440795898, + -2.3197848796844482, -3.078174352645874, -1.5299072265625, + ]), + tolerance: .float32) + } + } +} diff --git a/Tests/MLXIntegrationTests/Generated/GeneratedDefaultsTests.swift b/Tests/MLXIntegrationTests/Generated/GeneratedDefaultsTests.swift new file mode 100644 index 000000000..b915c6b39 --- /dev/null +++ b/Tests/MLXIntegrationTests/Generated/GeneratedDefaultsTests.swift @@ -0,0 +1,1410 @@ +// Copyright © 2026 Apple Inc. +// +// GENERATED by tools/integration_tests -- DO NOT EDIT. +// +// Values in this file were produced by python mlx and are compared against +// the Swift API. Regenerate with: +// +// python3 tools/integration_tests/generate.py +// +// python mlx: 0.32.2.dev20260910+1f8e74e3f +// vendored mlx (Cmlx): 0.32.2 +// generator revision: 2 +// cases: 59 + +import Foundation +import MLX +import MLXNN +import Testing + +@Suite("generated: Defaults") +struct GeneratedDefaultsTests { + + @Test("nanToNum") + func test_nanToNum() throws { + // python replaces +/-inf with the dtype max/min unless told otherwise + try withIntegrationState(seed: 12047) { + let a = MLXArray([1.0, -2.5, Float.infinity, -Float.infinity, Float.nan, 0.0]) + let result = MLX.nanToNum(a) + expectSummary( + result, + ArraySummary( + shape: [6], + dtype: .float32, + mean: 0.0, + minimum: -3.4028234663852886e+38, + maximum: 3.4028234663852886e+38, + absoluteSum: Double.infinity, + positionChecksum: Double.infinity, + sampleIndices: [0, 1, 2, 3, 4, 5], + samples: [1.0, -2.5, 3.4028234663852886e+38, -3.4028234663852886e+38, 0.0, 0.0]), + tolerance: .float32) + } + } + + @Test("tensordot") + func test_tensordot() throws { + // axes defaults to 2 (as in numpy), not 1 + try withIntegrationState(seed: 43747) { + let a = MLXRandom.normal([2, 3, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([3, 4, 5], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.tensordot(a, b) + expectSummary( + result, + ArraySummary( + shape: [2, 5], + dtype: .float32, + mean: -0.1383628100156784, + minimum: -6.713634967803955, + maximum: 6.5990705490112305, + absoluteSum: 30.51371955871582, + positionChecksum: 18.927197265625, + sampleIndices: [0, 2, 4, 5, 7, 9], + samples: [ + 0.017482735216617584, -2.9454421997070312, 6.5990705490112305, + 0.03260062262415886, 0.44432538747787476, 2.5363638401031494, + ]), + tolerance: .float32) + } + } + + @Test("isClose/within") + func test_isClose_within() throws { + // 1e-6 apart: inside the default rtol 1e-5 / atol 1e-8 + try withIntegrationState(seed: 9065) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let b = a + 1e-6 + let result = MLX.isClose(a, b) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .bool, + mean: 0.8333333730697632, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 10.0, + positionChecksum: 5.333333333333333, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [1.0, 1.0, 1.0, 1.0, 0.0, 1.0]), + tolerance: .exact) + } + } + + @Test("isClose/outside") + func test_isClose_outside() throws { + try withIntegrationState(seed: 50854) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let b = a + 1e-3 + let result = MLX.isClose(a, b) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .bool, + mean: 0.0, + minimum: 0.0, + maximum: 0.0, + absoluteSum: 0.0, + positionChecksum: 0.0, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [0.0, 0.0, 0.0, 0.0, 0.0, 0.0]), + tolerance: .exact) + } + } + + @Test("allClose") + func test_allClose() throws { + try withIntegrationState(seed: 52631) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let b = a + 1e-6 + let result = MLX.allClose(a, b) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .bool, + mean: 0.0, + minimum: 0.0, + maximum: 0.0, + absoluteSum: 0.0, + positionChecksum: 0.0, + sampleIndices: [0], + samples: [0.0]), + tolerance: .exact) + } + } + + @Test("arrayEqual") + func test_arrayEqual() throws { + try withIntegrationState(seed: 12896) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.arrayEqual(a, a) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .bool, + mean: 1.0, + minimum: 1.0, + maximum: 1.0, + absoluteSum: 1.0, + positionChecksum: 1.0, + sampleIndices: [0], + samples: [1.0]), + tolerance: .exact) + } + } + + @Test("diff") + func test_diff() throws { + try withIntegrationState(seed: 78249) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.diff(a) + expectSummary( + result, + ArraySummary( + shape: [4, 2], + dtype: .float32, + mean: 0.16116583347320557, + minimum: -1.240124225616455, + maximum: 1.5920885801315308, + absoluteSum: 7.024305820465088, + positionChecksum: 3.702437162399292, + sampleIndices: [0, 1, 3, 4, 6, 7], + samples: [ + 0.5840742588043213, -1.240124225616455, -0.21392442286014557, + 1.5920885801315308, 0.6998119950294495, -0.25859102606773376, + ]), + tolerance: .float32) + } + } + + @Test("trace") + func test_trace() throws { + try withIntegrationState(seed: 60095) { + let a = MLXRandom.normal([4, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.trace(a) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 2.147444009780884, + minimum: 2.147444009780884, + maximum: 2.147444009780884, + absoluteSum: 2.147444009780884, + positionChecksum: 2.147444009780884, + sampleIndices: [0], + samples: [2.147444009780884]), + tolerance: .float32) + } + } + + @Test("diag") + func test_diag() throws { + try withIntegrationState(seed: 94808) { + let a = MLXRandom.normal([4, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.diag(a) + expectSummary( + result, + ArraySummary( + shape: [4], + dtype: .float32, + mean: -0.4410656690597534, + minimum: -1.6250548362731934, + maximum: 0.6272233128547668, + absoluteSum: 3.3691916465759277, + positionChecksum: 1.827361822128296, + sampleIndices: [0, 1, 2, 3], + samples: [ + -1.6250548362731934, 0.17524121701717377, -0.9416722655296326, + 0.6272233128547668, + ]), + tolerance: .float32) + } + } + + @Test("diagonal") + func test_diagonal() throws { + try withIntegrationState(seed: 71392) { + let a = MLXRandom.normal([4, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.diagonal(a) + expectSummary( + result, + ArraySummary( + shape: [4], + dtype: .float32, + mean: 0.10765820741653442, + minimum: -1.8548674583435059, + maximum: 1.9325915575027466, + absoluteSum: 4.152509689331055, + positionChecksum: 3.0919158458709717, + sampleIndices: [0, 1, 2, 3], + samples: [ + -0.0060711028054356575, 1.9325915575027466, 0.35897988080978394, + -1.8548674583435059, + ]), + tolerance: .float32) + } + } + + @Test("flipped") + func test_flipped() throws { + try withIntegrationState(seed: 2426) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.flipped(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.3812035024166107, + minimum: -0.9314361214637756, + maximum: 2.043067693710327, + absoluteSum: 8.247232437133789, + positionChecksum: 5.1659698486328125, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.5396879315376282, -0.9314361214637756, 0.1392711102962494, + 1.5242443084716797, 2.043067693710327, 0.7740973234176636, + ]), + tolerance: .float32) + } + } + + @Test("round") + func test_round() throws { + try withIntegrationState(seed: 4970) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 10.0) + let result = MLX.round(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 2.3333334922790527, + minimum: -13.0, + maximum: 17.0, + absoluteSum: 88.0, + positionChecksum: 49.166666666666664, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [-13.0, 9.0, -3.0, 17.0, 10.0, 12.0]), + tolerance: .float32) + } + } + + @Test("softmax") + func test_softmax() throws { + try withIntegrationState(seed: 65331) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.softmax(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.0833333358168602, + minimum: 0.019660301506519318, + maximum: 0.1961202472448349, + absoluteSum: 1.0, + positionChecksum: 0.5366625785827637, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.056501682847738266, 0.023698141798377037, 0.14301876723766327, + 0.019660301506519318, 0.09527640044689178, 0.10355699062347412, + ]), + tolerance: .float32) + } + } + + @Test("cumsum") + func test_cumsum() throws { + try withIntegrationState(seed: 57443) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.cumsum(a) + expectSummary( + result, + ArraySummary( + shape: [12], + dtype: .float32, + mean: 0.0727943405508995, + minimum: -2.4656498432159424, + maximum: 1.888854742050171, + absoluteSum: 12.072671890258789, + positionChecksum: 7.7668304443359375, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.4540281295776367, 1.1351778507232666, 0.9642502069473267, + -2.4656498432159424, -0.36840271949768066, 1.888854742050171, + ]), + tolerance: .float32) + } + } + + @Test("cumprod") + func test_cumprod() throws { + try withIntegrationState(seed: 75562) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.cumprod(a) + expectSummary( + result, + ArraySummary( + shape: [12], + dtype: .float32, + mean: 0.08276982605457306, + minimum: -0.6033686399459839, + maximum: 0.7661811709403992, + absoluteSum: 2.5519943237304688, + positionChecksum: 0.796099583307902, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.1745307743549347, -0.6033686399459839, 0.7661811709403992, + -0.0003647314733825624, -0.0005849742447026074, -0.0001368979865219444, + ]), + tolerance: .float32) + } + } + + @Test("cummax") + func test_cummax() throws { + try withIntegrationState(seed: 27207) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.cummax(a) + expectSummary( + result, + ArraySummary( + shape: [12], + dtype: .float32, + mean: 1.3060426712036133, + minimum: 0.6872203350067139, + maximum: 1.429807186126709, + absoluteSum: 15.67251205444336, + positionChecksum: 9.108100255330404, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.6872203350067139, 1.429807186126709, 1.429807186126709, 1.429807186126709, + 1.429807186126709, 1.429807186126709, + ]), + tolerance: .float32) + } + } + + @Test("cummin") + func test_cummin() throws { + try withIntegrationState(seed: 84926) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.cummin(a) + expectSummary( + result, + ArraySummary( + shape: [12], + dtype: .float32, + mean: -1.6473878622055054, + minimum: -2.1973390579223633, + maximum: -0.8774563074111938, + absoluteSum: 19.768653869628906, + positionChecksum: 12.632850646972656, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.8774563074111938, -0.8774563074111938, -0.8774563074111938, + -2.1973390579223633, -2.1973390579223633, -2.1973390579223633, + ]), + tolerance: .float32) + } + } + + @Test("logCumsumExp") + func test_logCumsumExp() throws { + try withIntegrationState(seed: 43574) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.logCumsumExp(a) + expectSummary( + result, + ArraySummary( + shape: [12], + dtype: .float32, + mean: 2.4215335845947266, + minimum: -0.11741061508655548, + maximum: 3.1633148193359375, + absoluteSum: 29.29322052001953, + positionChecksum: 18.49832534790039, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.11741061508655548, 2.4433906078338623, 2.7169978618621826, + 2.9348068237304688, 3.073072671890259, 3.1633148193359375, + ]), + tolerance: .float32) + } + } + + @Test("median") + func test_median() throws { + try withIntegrationState(seed: 92274) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.median(a) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: -0.07507183402776718, + minimum: -0.07507183402776718, + maximum: -0.07507183402776718, + absoluteSum: 0.07507183402776718, + positionChecksum: 0.07507183402776718, + sampleIndices: [0], + samples: [-0.07507183402776718]), + tolerance: .float32) + } + } + + @Test("variance") + func test_variance() throws { + try withIntegrationState(seed: 5488) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.variance(a) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 0.4012840986251831, + minimum: 0.4012840986251831, + maximum: 0.4012840986251831, + absoluteSum: 0.4012840986251831, + positionChecksum: 0.4012840986251831, + sampleIndices: [0], + samples: [0.4012840986251831]), + tolerance: .float32) + } + } + + @Test("std") + func test_std() throws { + try withIntegrationState(seed: 52117) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.std(a) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 0.7837796211242676, + minimum: 0.7837796211242676, + maximum: 0.7837796211242676, + absoluteSum: 0.7837796211242676, + positionChecksum: 0.7837796211242676, + sampleIndices: [0], + samples: [0.7837796211242676]), + tolerance: .float32) + } + } + + @Test("hadamardTransform") + func test_hadamardTransform() throws { + try withIntegrationState(seed: 17021) { + let a = MLXRandom.normal([4, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.hadamardTransform(a) + expectSummary( + result, + ArraySummary( + shape: [4, 16], + dtype: .float32, + mean: -0.1470964550971985, + minimum: -2.6109731197357178, + maximum: 1.7692186832427979, + absoluteSum: 53.5502815246582, + positionChecksum: 25.162355422973633, + sampleIndices: [0, 13, 25, 38, 50, 63], + samples: [ + 1.0269194841384888, -1.6071994304656982, -1.6351640224456787, + -1.1893904209136963, 1.2207064628601074, 1.7337404489517212, + ]), + tolerance: .float32) + } + } + + @Test("linspace") + func test_linspace() throws { + try withIntegrationState(seed: 63792) { + let result = MLX.linspace(0.0, 1.0) + expectSummary( + result, + ArraySummary( + shape: [50], + dtype: .float32, + mean: 0.5, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 25.0, + positionChecksum: 17.0, + sampleIndices: [0, 10, 20, 29, 39, 49], + samples: [ + 0.0, 0.20408163964748383, 0.40816327929496765, 0.5918367505073547, + 0.795918345451355, 1.0, + ]), + tolerance: .float32) + } + } + + @Test("eye") + func test_eye() throws { + try withIntegrationState(seed: 94476) { + let result = MLX.eye(4) + expectSummary( + result, + ArraySummary( + shape: [4, 4], + dtype: .float32, + mean: 0.25, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 4.0, + positionChecksum: 2.125, + sampleIndices: [0, 3, 6, 9, 12, 15], + samples: [1.0, 0.0, 0.0, 0.0, 0.0, 1.0]), + tolerance: .float32) + } + } + + @Test("tri") + func test_tri() throws { + // python requires m and k; Swift defaults them to n and 0 + try withIntegrationState(seed: 81451) { + let result = MLX.tri(4) + expectSummary( + result, + ArraySummary( + shape: [4, 4], + dtype: .float32, + mean: 0.625, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 10.0, + positionChecksum: 6.25, + sampleIndices: [0, 3, 6, 9, 12, 15], + samples: [1.0, 0.0, 0.0, 1.0, 1.0, 1.0]), + tolerance: .float32) + } + } + + @Test("convolve") + func test_convolve() throws { + try withIntegrationState(seed: 62456) { + let a = MLXRandom.normal([20], dtype: .float32, loc: 0.0, scale: 1.0) + let v = MLXRandom.normal([4], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.convolve(a, v) + expectSummary( + result, + ArraySummary( + shape: [23], + dtype: .float32, + mean: 0.08946806937456131, + minimum: -1.803636908531189, + maximum: 1.8454562425613403, + absoluteSum: 19.04773712158203, + positionChecksum: 8.943277110224185, + sampleIndices: [0, 4, 9, 13, 18, 22], + samples: [ + 1.0646476745605469, 1.7893160581588745, -0.8670757412910461, + -0.4499398171901703, 1.8454562425613403, -0.021289121359586716, + ]), + tolerance: .float32) + } + } + + @Test("conv1d") + func test_conv1d() throws { + try withIntegrationState(seed: 88096) { + let input = MLXRandom.normal([2, 10, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let weight = MLXRandom.normal([3, 3, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.conv1d(input, weight) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 3], + dtype: .float32, + mean: 0.5933355093002319, + minimum: -7.674775123596191, + maximum: 7.810962677001953, + absoluteSum: 154.97207641601562, + positionChecksum: 74.69058227539062, + sampleIndices: [0, 9, 19, 28, 38, 47], + samples: [ + 1.575302243232727, -0.19398178160190582, 4.896422863006592, + 7.810962677001953, 2.7325944900512695, 3.176520586013794, + ]), + tolerance: .float32) + } + } + + @Test("quantized") + func test_quantized() throws { + // group size and bit count default in the C++ core (64 / 4) + try withIntegrationState(seed: 5234) { + let w = MLXRandom.uniform(low: -1.0, high: 1.0, [64, 128], dtype: .float32) + let result = MLX.quantized(w).wq + expectSummary( + result, + ArraySummary( + shape: [64, 16], + dtype: .uint32, + mean: 2255552512.0, + minimum: 10032196.0, + maximum: 4293439744.0, + absoluteSum: 2309685772288.0, + positionChecksum: 1166251655168.0, + sampleIndices: [0, 205, 409, 614, 818, 1023], + samples: [ + 1967645312.0, 1215345152.0, 1610332928.0, 788708928.0, 3749942784.0, + 3667559680.0, + ]), + tolerance: .exact) + } + } + + @Test("dequantized") + func test_dequantized() throws { + try withIntegrationState(seed: 20669) { + let w = MLXRandom.uniform(low: -1.0, high: 1.0, [64, 128], dtype: .float32) + let q = MLX.quantized(w) + let result = MLX.dequantized(q.wq, scales: q.scales, biases: q.biases) + expectSummary( + result, + ArraySummary( + shape: [64, 128], + dtype: .float32, + mean: -0.0056404173374176025, + minimum: -0.9999048709869385, + maximum: 0.9997067451477051, + absoluteSum: 4079.2216796875, + positionChecksum: 2048.1728515625, + sampleIndices: [0, 1638, 3276, 4915, 6553, 8191], + samples: [ + 0.7425611615180969, -0.24511511623859406, -0.9868167042732239, + -0.12466387450695038, -0.6139729022979736, -0.24921728670597076, + ]), + tolerance: .float32) + } + } + + @Test("quantizedMM") + func test_quantizedMM() throws { + try withIntegrationState(seed: 17188) { + let w = MLXRandom.uniform(low: -1.0, high: 1.0, [64, 128], dtype: .float32) + let x = MLXRandom.normal([8, 128], dtype: .float32, loc: 0.0, scale: 1.0) + let q = MLX.quantized(w) + let result = MLX.quantizedMM(x, q.wq, scales: q.scales, biases: q.biases) + expectSummary( + result, + ArraySummary( + shape: [8, 64], + dtype: .float32, + mean: -0.008240580558776855, + minimum: -17.94013214111328, + maximum: 18.256729125976562, + absoluteSum: 2450.4580078125, + positionChecksum: 1212.8326416015625, + sampleIndices: [0, 102, 204, 307, 409, 511], + samples: [ + 4.712850093841553, 6.662078857421875, 6.3404364585876465, 5.768905162811279, + -6.76831579208374, -5.105240821838379, + ]), + tolerance: .float32) + } + } + + @Test("norm") + func test_norm() throws { + try withIntegrationState(seed: 51906) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.norm(a, stream: .cpu) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 2.6667308807373047, + minimum: 2.6667308807373047, + maximum: 2.6667308807373047, + absoluteSum: 2.6667308807373047, + positionChecksum: 2.6667308807373047, + sampleIndices: [0], + samples: [2.6667308807373047]), + tolerance: .float32) + } + } + + @Test("cholesky") + func test_cholesky() throws { + // upper defaults to false on both sides + try withIntegrationState(seed: 65482) { + let a = MLXRandom.normal([4, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let spd = MLX.matmul(a, a.T) + 4.0 * MLX.eye(4) + let result = MLX.cholesky(spd, stream: .cpu) + expectSummary( + result, + ArraySummary( + shape: [4, 4], + dtype: .float32, + mean: 0.6790149211883545, + minimum: -1.1485399007797241, + maximum: 3.0164363384246826, + absoluteSum: 13.845975875854492, + positionChecksum: 8.125425338745117, + sampleIndices: [0, 3, 6, 9, 12, 15], + samples: [ + 3.0164363384246826, 0.0, 0.0, -0.2918112277984619, 0.553765058517456, + 2.652482271194458, + ]), + tolerance: .float32) + } + } + + @Test("triInv") + func test_triInv() throws { + try withIntegrationState(seed: 32119) { + let a = MLXRandom.normal([4, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let spd = MLX.matmul(a, a.T) + 4.0 * MLX.eye(4) + let lower = MLX.tril(spd) + let result = MLX.triInv(lower, stream: .cpu) + expectSummary( + result, + ArraySummary( + shape: [4, 4], + dtype: .float32, + mean: 0.021889742463827133, + minimum: -0.040400367230176926, + maximum: 0.14794215559959412, + absoluteSum: 0.6479233503341675, + positionChecksum: 0.3984556496143341, + sampleIndices: [0, 3, 6, 9, 12, 15], + samples: [ + 0.09358751028776169, 0.0, 0.0, -0.03134113550186157, 0.005859235301613808, + 0.14794215559959412, + ]), + tolerance: .float32) + } + } + + @Test("solveTriangular") + func test_solveTriangular() throws { + try withIntegrationState(seed: 82366) { + let a = MLXRandom.normal([4, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let spd = MLX.matmul(a, a.T) + 4.0 * MLX.eye(4) + let lower = MLX.tril(spd) + let b = MLXRandom.normal([4, 2], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.solveTriangular(lower, b, stream: .cpu) + expectSummary( + result, + ArraySummary( + shape: [4, 2], + dtype: .float32, + mean: 0.007373856380581856, + minimum: -0.15618708729743958, + maximum: 0.1315477341413498, + absoluteSum: 0.5120813250541687, + positionChecksum: 0.23803149163722992, + sampleIndices: [0, 1, 3, 4, 6, 7], + samples: [ + -0.15618708729743958, -0.065582275390625, 0.1315477341413498, + 0.03207559138536453, 0.0042139156721532345, 0.09015661478042603, + ]), + tolerance: .float32) + } + } + + @Test("fft") + func test_fft() throws { + try withIntegrationState(seed: 13721) { + let r = MLXRandom.normal([100], dtype: .float32, loc: 0.0, scale: 1.0) + let i = MLXRandom.normal([100], dtype: .float32, loc: 0.0, scale: 1.0) + let c = r + i.asImaginary() + let result = MLX.fft(c) + expectSummary( + result.realPart(), + ArraySummary( + shape: [100], + dtype: .float32, + mean: 0.5256693959236145, + minimum: -20.80100440979004, + maximum: 19.87451171875, + absoluteSum: 795.9764404296875, + positionChecksum: 385.171640625, + sampleIndices: [0, 20, 40, 59, 79, 99], + samples: [ + -1.5068063735961914, -8.961091041564941, -7.156008243560791, + 8.837318420410156, 0.48036813735961914, 0.8878612518310547, + ]), + tolerance: .float32) + expectSummary( + result.imaginaryPart(), + ArraySummary( + shape: [100], + dtype: .float32, + mean: -0.07993728667497635, + minimum: -23.762022018432617, + maximum: 17.19780921936035, + absoluteSum: 733.4420776367188, + positionChecksum: 353.8171875, + sampleIndices: [0, 20, 40, 59, 79, 99], + samples: [ + 8.099128723144531, -16.915706634521484, 0.5834121704101562, + 5.633106231689453, 4.721193313598633, -21.837533950805664, + ]), + tolerance: .float32) + } + } + + @Test("fftfreq") + func test_fftfreq() throws { + try withIntegrationState(seed: 87337) { + let result = MLX.fftfreq(16) + expectSummary( + result, + ArraySummary( + shape: [16], + dtype: .float32, + mean: -0.03125, + minimum: -0.5, + maximum: 0.4375, + absoluteSum: 4.0, + positionChecksum: 2.25, + sampleIndices: [0, 3, 6, 9, 12, 15], + samples: [0.0, 0.1875, 0.375, -0.4375, -0.25, -0.0625]), + tolerance: .float32) + } + } + + @Test("gumbel") + func test_gumbel() throws { + try withIntegrationState(seed: 53461) { + let result = MLXRandom.gumbel() + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 1.186912178993225, + minimum: 1.186912178993225, + maximum: 1.186912178993225, + absoluteSum: 1.186912178993225, + positionChecksum: 1.186912178993225, + sampleIndices: [0], + samples: [1.186912178993225]), + tolerance: .float32) + } + } + + @Test("uniform") + func test_uniform() throws { + try withIntegrationState(seed: 83902) { + let result = MLXRandom.uniform() + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 0.9644234776496887, + minimum: 0.9644234776496887, + maximum: 0.9644234776496887, + absoluteSum: 0.9644234776496887, + positionChecksum: 0.9644234776496887, + sampleIndices: [0], + samples: [0.9644234776496887]), + tolerance: .float32) + } + } + + @Test("normal") + func test_normal() throws { + try withIntegrationState(seed: 67758) { + let result = MLXRandom.normal() + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: -0.3701243996620178, + minimum: -0.3701243996620178, + maximum: -0.3701243996620178, + absoluteSum: 0.3701243996620178, + positionChecksum: 0.3701243996620178, + sampleIndices: [0], + samples: [-0.3701243996620178]), + tolerance: .float32) + } + } + + @Test("bernoulli") + func test_bernoulli() throws { + try withIntegrationState(seed: 34396) { + let result = MLXRandom.bernoulli() + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .bool, + mean: 1.0, + minimum: 1.0, + maximum: 1.0, + absoluteSum: 1.0, + positionChecksum: 1.0, + sampleIndices: [0], + samples: [1.0]), + tolerance: .exact) + } + } + + @Test("sorted/flat") + func test_sorted_flat() throws { + try withIntegrationState(seed: 45340) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.sorted(a) + expectSummary( + result, + ArraySummary( + shape: [12], + dtype: .float32, + mean: 0.7008785009384155, + minimum: -0.6916800737380981, + maximum: 1.6399438381195068, + absoluteSum: 9.793901443481445, + positionChecksum: 6.6583404541015625, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.6916800737380981, 0.2463919222354889, 0.49870675802230835, + 0.9339690804481506, 1.1501116752624512, 1.6399438381195068, + ]), + tolerance: .float32) + } + } + + @Test("argSort/flat") + func test_argSort_flat() throws { + try withIntegrationState(seed: 29628) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.argSort(a) + expectSummary( + result, + ArraySummary( + shape: [12], + dtype: .uint32, + mean: 5.5, + minimum: 0.0, + maximum: 11.0, + absoluteSum: 66.0, + positionChecksum: 31.833333333333332, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [9.0, 6.0, 11.0, 1.0, 0.0, 4.0]), + tolerance: .exact) + } + } + + @Test("partitioned/flat") + func test_partitioned_flat() throws { + try withIntegrationState(seed: 94638) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.partitioned(a, kth: 3) + expectSummary( + result, + ArraySummary( + shape: [12], + dtype: .float32, + mean: -0.2733670473098755, + minimum: -1.770876169204712, + maximum: 1.3596465587615967, + absoluteSum: 10.031713485717773, + positionChecksum: 4.404526074727376, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -1.770876169204712, -1.5603525638580322, -0.5617737770080566, + 0.21052542328834534, 0.6110134124755859, 1.3596465587615967, + ]), + tolerance: .float32) + } + } + + @Test("argPartition/flat") + func test_argPartition_flat() throws { + try withIntegrationState(seed: 12369) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.argPartition(a, kth: 3) + expectSummary( + result, + ArraySummary( + shape: [12], + dtype: .uint32, + mean: 5.5, + minimum: 0.0, + maximum: 11.0, + absoluteSum: 66.0, + positionChecksum: 35.166666666666664, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [11.0, 9.0, 4.0, 7.0, 10.0, 8.0]), + tolerance: .exact) + } + } + + @Test("top/flat") + func test_top_flat() throws { + try withIntegrationState(seed: 32885) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.top(a, k: 3) + expectSummary( + result, + ArraySummary( + shape: [3], + dtype: .float32, + mean: 0.8019466400146484, + minimum: 0.5089910626411438, + maximum: 1.005195140838623, + absoluteSum: 2.4058399200439453, + positionChecksum: 1.7692945798238118, + sampleIndices: [0, 1, 2], + samples: [0.5089910626411438, 0.8916536569595337, 1.005195140838623]), + tolerance: .float32) + } + } + + @Test("roll/flat") + func test_roll_flat() throws { + try withIntegrationState(seed: 90193) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.roll(a, shift: 2) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.5598515272140503, + minimum: -1.5043145418167114, + maximum: 2.463324546813965, + absoluteSum: 10.23849868774414, + positionChecksum: 4.870436032613118, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.2558262050151825, 1.3071829080581665, 0.6110653281211853, + 1.2561404705047607, 0.022526530548930168, 0.05659160390496254, + ]), + tolerance: .float32) + } + } + + @Test("takeAlong/flat") + func test_takeAlong_flat() throws { + try withIntegrationState(seed: 42978) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let indices = MLXRandom.randInt(low: 0, high: 12, [4], type: Int32.self) + let result = MLX.takeAlong(a, indices) + expectSummary( + result, + ArraySummary( + shape: [4], + dtype: .float32, + mean: -0.19436131417751312, + minimum: -0.8136857151985168, + maximum: 0.2381412237882614, + absoluteSum: 1.560187578201294, + positionChecksum: 1.216184139251709, + sampleIndices: [0, 1, 2, 3], + samples: [ + 0.2381412237882614, 0.1532299816608429, -0.35513073205947876, + -0.8136857151985168, + ]), + tolerance: .float32) + } + } + + @Test("repeated/flat") + func test_repeated_flat() throws { + try withIntegrationState(seed: 68684) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.repeated(a, count: 2) + expectSummary( + result, + ArraySummary( + shape: [24], + dtype: .float32, + mean: -0.0766555666923523, + minimum: -0.9837857484817505, + maximum: 0.9884175062179565, + absoluteSum: 12.13071060180664, + positionChecksum: 6.469824473063151, + sampleIndices: [0, 5, 9, 14, 18, 23], + samples: [ + -0.3888154923915863, 0.35114219784736633, -0.4268342852592468, + -0.750239372253418, -0.3761099576950073, 0.05179966613650322, + ]), + tolerance: .float32) + } + } + + @Test("crossEntropy") + func test_crossEntropy() throws { + try withIntegrationState(seed: 58839) { + let logits = MLXRandom.normal([4, 5], dtype: .float32, loc: 0.0, scale: 1.0) + let targets = MLXRandom.randInt(low: 0, high: 5, [4], type: Int32.self) + let result = MLXNN.crossEntropy(logits: logits, targets: targets) + expectSummary( + result, + ArraySummary( + shape: [4], + dtype: .float32, + mean: 1.5683512687683105, + minimum: 0.43761518597602844, + maximum: 3.6751675605773926, + absoluteSum: 6.273405075073242, + positionChecksum: 4.924343585968018, + sampleIndices: [0, 1, 2, 3], + samples: [ + 1.1801968812942505, 0.43761518597602844, 0.9804257154464722, + 3.6751675605773926, + ]), + tolerance: .float32) + } + } + + @Test("binaryCrossEntropy") + func test_binaryCrossEntropy() throws { + try withIntegrationState(seed: 37138) { + let logits = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let targets = MLXRandom.bernoulli(0.5, [4, 3]) + let result = MLXNN.binaryCrossEntropy(logits: logits, targets: targets.asType(.float32)) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 0.8657888174057007, + minimum: 0.8657888174057007, + maximum: 0.8657888174057007, + absoluteSum: 0.8657888174057007, + positionChecksum: 0.8657888174057007, + sampleIndices: [0], + samples: [0.8657888174057007]), + tolerance: .float32) + } + } + + @Test("l1Loss") + func test_l1Loss() throws { + try withIntegrationState(seed: 94773) { + let predictions = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let targets = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.l1Loss(predictions: predictions, targets: targets) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 0.5470814108848572, + minimum: 0.5470814108848572, + maximum: 0.5470814108848572, + absoluteSum: 0.5470814108848572, + positionChecksum: 0.5470814108848572, + sampleIndices: [0], + samples: [0.5470814108848572]), + tolerance: .float32) + } + } + + @Test("mseLoss") + func test_mseLoss() throws { + try withIntegrationState(seed: 57835) { + let predictions = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let targets = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.mseLoss(predictions: predictions, targets: targets) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 1.724809169769287, + minimum: 1.724809169769287, + maximum: 1.724809169769287, + absoluteSum: 1.724809169769287, + positionChecksum: 1.724809169769287, + sampleIndices: [0], + samples: [1.724809169769287]), + tolerance: .float32) + } + } + + @Test("smoothL1Loss") + func test_smoothL1Loss() throws { + try withIntegrationState(seed: 44045) { + let predictions = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let targets = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.smoothL1Loss(predictions: predictions, targets: targets) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 0.5390201807022095, + minimum: 0.5390201807022095, + maximum: 0.5390201807022095, + absoluteSum: 0.5390201807022095, + positionChecksum: 0.5390201807022095, + sampleIndices: [0], + samples: [0.5390201807022095]), + tolerance: .float32) + } + } + + @Test("huberLoss") + func test_huberLoss() throws { + try withIntegrationState(seed: 20289) { + let inputs = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let targets = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.huberLoss(inputs: inputs, targets: targets) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.933381199836731, + minimum: 0.017336202785372734, + maximum: 2.534982681274414, + absoluteSum: 11.200573921203613, + positionChecksum: 5.508398691813151, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.6122961044311523, 0.14958451688289642, 2.534982681274414, + 0.14333206415176392, 1.719494104385376, 0.608657717704773, + ]), + tolerance: .float32) + } + } + + @Test("logCoshLoss") + func test_logCoshLoss() throws { + try withIntegrationState(seed: 65708) { + let inputs = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let targets = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.logCoshLoss(inputs: inputs, targets: targets) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.6543585062026978, + minimum: 0.025879621505737305, + maximum: 1.854630708694458, + absoluteSum: 7.852302074432373, + positionChecksum: 4.758116404215495, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.09128165245056152, 1.542405366897583, 0.11331796646118164, + 0.13976502418518066, 1.8090453147888184, 0.6400725841522217, + ]), + tolerance: .float32) + } + } + + @Test("hingeLoss") + func test_hingeLoss() throws { + try withIntegrationState(seed: 32168) { + let inputs = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let signs = MLXRandom.bernoulli(0.5, [4, 3]) + let targets = MLX.which(signs, 1.0, -1.0) + let result = MLXNN.hingeLoss(inputs: inputs, targets: targets) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.990768313407898, + minimum: 0.0, + maximum: 2.4934892654418945, + absoluteSum: 11.889219284057617, + positionChecksum: 6.786125183105469, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 1.3253191709518433, 1.9193181991577148, 0.0, 0.9395096898078918, + 0.37627631425857544, 0.0, + ]), + tolerance: .float32) + } + } + + @Test("nllLoss") + func test_nllLoss() throws { + try withIntegrationState(seed: 78313) { + let logits = MLXRandom.normal([4, 5], dtype: .float32, loc: 0.0, scale: 1.0) + let inputs = MLX.log(MLX.softmax(logits, axis: -1)) + let targets = MLXRandom.randInt(low: 0, high: 5, [4], type: Int32.self) + let result = MLXNN.nllLoss(inputs: inputs, targets: targets) + expectSummary( + result, + ArraySummary( + shape: [4], + dtype: .float32, + mean: 1.7705587148666382, + minimum: 0.7432368993759155, + maximum: 3.2677173614501953, + absoluteSum: 7.082234859466553, + positionChecksum: 5.156996726989746, + sampleIndices: [0, 1, 2, 3], + samples: [ + 0.8151545524597168, 2.2561261653900146, 0.7432368993759155, + 3.2677173614501953, + ]), + tolerance: .float32) + } + } + + @Test("tripletLoss") + func test_tripletLoss() throws { + try withIntegrationState(seed: 18126) { + let anchors = MLXRandom.normal([4, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let positives = MLXRandom.normal([4, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let negatives = MLXRandom.normal([4, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.tripletLoss( + anchors: anchors, positives: positives, negatives: negatives) + expectSummary( + result, + ArraySummary( + shape: [4], + dtype: .float32, + mean: 1.5873920917510986, + minimum: 1.2556252479553223, + maximum: 1.9852931499481201, + absoluteSum: 6.3495683670043945, + positionChecksum: 3.7734334468841553, + sampleIndices: [0, 1, 2, 3], + samples: [ + 1.9852931499481201, 1.5465178489685059, 1.2556252479553223, + 1.5621321201324463, + ]), + tolerance: .float32) + } + } + + @Test("cosineSimilarityLoss") + func test_cosineSimilarityLoss() throws { + try withIntegrationState(seed: 76691) { + let x1 = MLXRandom.normal([4, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let x2 = MLXRandom.normal([4, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.cosineSimilarityLoss(x1: x1, x2: x2) + expectSummary( + result, + ArraySummary( + shape: [4], + dtype: .float32, + mean: 0.3369714617729187, + minimum: 0.12991397082805634, + maximum: 0.6953044533729553, + absoluteSum: 1.3478858470916748, + positionChecksum: 0.698710560798645, + sampleIndices: [0, 1, 2, 3], + samples: [ + 0.3417125642299652, 0.6953044533729553, 0.1809549331665039, + 0.12991397082805634, + ]), + tolerance: .float32) + } + } +} diff --git a/Tests/MLXIntegrationTests/Generated/GeneratedElementwiseTests.swift b/Tests/MLXIntegrationTests/Generated/GeneratedElementwiseTests.swift new file mode 100644 index 000000000..8dd22458a --- /dev/null +++ b/Tests/MLXIntegrationTests/Generated/GeneratedElementwiseTests.swift @@ -0,0 +1,1517 @@ +// Copyright © 2026 Apple Inc. +// +// GENERATED by tools/integration_tests -- DO NOT EDIT. +// +// Values in this file were produced by python mlx and are compared against +// the Swift API. Regenerate with: +// +// python3 tools/integration_tests/generate.py +// +// python mlx: 0.32.2.dev20260910+1f8e74e3f +// vendored mlx (Cmlx): 0.32.2 +// generator revision: 2 +// cases: 63 + +import Foundation +import MLX +import Testing + +@Suite("generated: Elementwise") +struct GeneratedElementwiseTests { + + @Test("exp") + func test_exp() throws { + try withIntegrationState(seed: 47954) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.exp(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 1.510894775390625, + minimum: 0.2611827552318573, + maximum: 7.178291320800781, + absoluteSum: 18.1307373046875, + positionChecksum: 7.237987518310547, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 7.178291320800781, 1.276309847831726, 0.298198401927948, 0.2611827552318573, + 0.6586772799491882, 0.5715979933738708, + ]), + tolerance: .float32) + } + } + + @Test("expm1") + func test_expm1() throws { + try withIntegrationState(seed: 78136) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.expm1(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.3363969922065735, + minimum: -0.7551273703575134, + maximum: 1.9503211975097656, + absoluteSum: 8.781084060668945, + positionChecksum: 4.586121241251628, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.3306465744972229, 1.373231053352356, -0.3045533299446106, + -0.6832998991012573, -0.6291795969009399, 0.5119333863258362, + ]), + tolerance: .float32) + } + } + + @Test("log") + func test_log() throws { + try withIntegrationState(seed: 57228) { + let a = MLXRandom.uniform(low: 0.1, high: 2.0, [4, 3], dtype: .float32) + let result = MLX.log(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.13923785090446472, + minimum: -0.6309344172477722, + maximum: 0.6703048944473267, + absoluteSum: 4.917095184326172, + positionChecksum: 2.5226616859436035, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.6309344172477722, 0.15799346566200256, 0.2860988676548004, + -0.1460823118686676, -0.28112390637397766, 0.6090521812438965, + ]), + tolerance: .float32) + } + } + + @Test("log2") + func test_log2() throws { + try withIntegrationState(seed: 96094) { + let a = MLXRandom.uniform(low: 0.1, high: 2.0, [4, 3], dtype: .float32) + let result = MLX.log2(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.3897838592529297, + minimum: -3.062103271484375, + maximum: 0.6787765026092529, + absoluteSum: 11.070028305053711, + positionChecksum: 5.862569173177083, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.607052743434906, -1.283917784690857, -1.0272319316864014, + 0.2580295503139496, -3.062103271484375, 0.6499488949775696, + ]), + tolerance: .float32) + } + } + + @Test("log10") + func test_log10() throws { + try withIntegrationState(seed: 4788) { + let a = MLXRandom.uniform(low: 0.1, high: 2.0, [4, 3], dtype: .float32) + let result = MLX.log10(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.07758500427007675, + minimum: -0.7380918264389038, + maximum: 0.2524550259113312, + absoluteSum: 3.0936691761016846, + positionChecksum: 1.3356029192606609, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.24614253640174866, -0.3498530387878418, 0.16428828239440918, + 0.1663184016942978, 0.16893142461776733, 0.05258811637759209, + ]), + tolerance: .float32) + } + } + + @Test("log1p") + func test_log1p() throws { + try withIntegrationState(seed: 75204) { + let a = MLXRandom.uniform(low: 0.1, high: 2.0, [4, 3], dtype: .float32) + let result = MLX.log1p(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.7237419486045837, + minimum: 0.18861305713653564, + maximum: 1.0836082696914673, + absoluteSum: 8.684903144836426, + positionChecksum: 4.419689814249675, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.8431658744812012, 0.45729416608810425, 1.0836082696914673, + 1.0338786840438843, 0.3002285659313202, 0.9232745170593262, + ]), + tolerance: .float32) + } + } + + @Test("sqrt") + func test_sqrt() throws { + try withIntegrationState(seed: 65948) { + let a = MLXRandom.uniform(low: 0.1, high: 2.0, [4, 3], dtype: .float32) + let result = MLX.sqrt(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 1.0848369598388672, + minimum: 0.5140633583068848, + maximum: 1.4113861322402954, + absoluteSum: 13.018043518066406, + positionChecksum: 6.730450948079427, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 1.4113861322402954, 1.3208208084106445, 0.5140633583068848, + 0.8169410824775696, 0.6033309102058411, 1.0939743518829346, + ]), + tolerance: .float32) + } + } + + @Test("rsqrt") + func test_rsqrt() throws { + try withIntegrationState(seed: 35677) { + let a = MLXRandom.uniform(low: 0.1, high: 2.0, [4, 3], dtype: .float32) + let result = MLX.rsqrt(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 1.1557154655456543, + minimum: 0.7132951021194458, + maximum: 3.140223264694214, + absoluteSum: 13.868584632873535, + positionChecksum: 7.229054133097331, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.7132951021194458, 2.035323143005371, 0.964241623878479, 0.896981954574585, + 0.7307946681976318, 0.9336780309677124, + ]), + tolerance: .float32) + } + } + + @Test("square") + func test_square() throws { + try withIntegrationState(seed: 72464) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.square(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.695485532283783, + minimum: 0.014689148403704166, + maximum: 4.136717319488525, + absoluteSum: 8.345826148986816, + positionChecksum: 6.028196334838867, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.020748626440763474, 0.4422823488712311, 0.02497928962111473, + 0.11787982285022736, 4.136717319488525, 0.20451615750789642, + ]), + tolerance: .float32) + } + } + + @Test("abs") + func test_abs() throws { + try withIntegrationState(seed: 54927) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.abs(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.7679557800292969, + minimum: 0.05212230235338211, + maximum: 1.779400110244751, + absoluteSum: 9.215469360351562, + positionChecksum: 5.223146438598633, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.2723303735256195, 0.05212230235338211, 0.7042833566665649, + 0.4646765887737274, 0.6566354632377625, 1.0102406740188599, + ]), + tolerance: .float32) + } + } + + @Test("negative") + func test_negative() throws { + try withIntegrationState(seed: 65451) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.negative(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.40208250284194946, + minimum: -2.357377290725708, + maximum: 1.9783403873443604, + absoluteSum: 11.139066696166992, + positionChecksum: 6.509375254313151, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.32914111018180847, 0.6736426949501038, -0.43585315346717834, + -0.034667275846004486, -2.357377290725708, 0.7370271682739258, + ]), + tolerance: .float32) + } + } + + @Test("sign") + func test_sign() throws { + try withIntegrationState(seed: 14106) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.sign(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.1666666716337204, + minimum: -1.0, + maximum: 1.0, + absoluteSum: 12.0, + positionChecksum: 6.5, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [-1.0, -1.0, 1.0, -1.0, 1.0, -1.0]), + tolerance: .float32) + } + } + + @Test("floor") + func test_floor() throws { + try withIntegrationState(seed: 48869) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.floor(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.4166666865348816, + minimum: -2.0, + maximum: 1.0, + absoluteSum: 11.0, + positionChecksum: 7.083333333333333, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [-1.0, 1.0, 0.0, 1.0, -1.0, -1.0]), + tolerance: .float32) + } + } + + @Test("ceil") + func test_ceil() throws { + try withIntegrationState(seed: 65987) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.ceil(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.8333333730697632, + minimum: -1.0, + maximum: 3.0, + absoluteSum: 14.0, + positionChecksum: 7.25, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [3.0, -0.0, 1.0, -1.0, 1.0, -1.0]), + tolerance: .float32) + } + } + + @Test("trunc") + func test_trunc() throws { + try withIntegrationState(seed: 30046) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.trunc(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.25, + minimum: -0.0, + maximum: 1.0, + absoluteSum: 3.0, + positionChecksum: 1.0833333333333333, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [1.0, 0.0, 0.0, 1.0, 0.0, 0.0]), + tolerance: .float32) + } + } + + @Test("reciprocal") + func test_reciprocal() throws { + try withIntegrationState(seed: 56536) { + let a = MLXRandom.uniform(low: 0.5, high: 2.0, [4, 3], dtype: .float32) + let result = MLX.reciprocal(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 1.0375549793243408, + minimum: 0.5049313306808472, + maximum: 1.9941542148590088, + absoluteSum: 12.450658798217773, + positionChecksum: 6.7019303639729815, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.8576971888542175, 1.7296605110168457, 1.9941542148590088, + 0.6068732738494873, 1.437481164932251, 1.684537410736084, + ]), + tolerance: .float32) + } + } + + @Test("sin") + func test_sin() throws { + try withIntegrationState(seed: 22723) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.sin(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.026744266971945763, + minimum: -0.9563843011856079, + maximum: 0.8906667828559875, + absoluteSum: 7.560203552246094, + positionChecksum: 4.530291557312012, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.3366437554359436, -0.10839032381772995, -0.41767552495002747, + -0.9191461205482483, -0.8813959360122681, 0.8294884562492371, + ]), + tolerance: .float32) + } + } + + @Test("cos") + func test_cos() throws { + try withIntegrationState(seed: 79724) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.cos(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.6566232442855835, + minimum: -0.3807412087917328, + maximum: 0.9957085251808167, + absoluteSum: 8.786455154418945, + positionChecksum: 4.586637496948242, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.8214612603187561, -0.3807412087917328, 0.9957085251808167, + 0.9071047306060791, 0.9621671438217163, -0.07274708896875381, + ]), + tolerance: .float32) + } + } + + @Test("tan") + func test_tan() throws { + try withIntegrationState(seed: 53934) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.tan(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.5176577568054199, + minimum: -0.9707298874855042, + maximum: 3.0964441299438477, + absoluteSum: 11.18362045288086, + positionChecksum: 4.061483383178711, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.44313281774520874, 3.0964441299438477, -0.4185173213481903, + -0.9707298874855042, -0.04070381820201874, -0.15154463052749634, + ]), + tolerance: .float32) + } + } + + @Test("sinh") + func test_sinh() throws { + try withIntegrationState(seed: 99750) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.sinh(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.2838282585144043, + minimum: -3.0325629711151123, + maximum: 1.4805724620819092, + absoluteSum: 14.353582382202148, + positionChecksum: 5.958340962727864, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 1.3117921352386475, -2.4681921005249023, 1.1744270324707031, + -0.6643975973129272, -0.37228044867515564, -0.08928694576025009, + ]), + tolerance: .float32) + } + } + + @Test("cosh") + func test_cosh() throws { + try withIntegrationState(seed: 31959) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.cosh(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 1.7174079418182373, + minimum: 1.0034161806106567, + maximum: 5.466761589050293, + absoluteSum: 20.60889434814453, + positionChecksum: 12.778172810872396, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 1.0733435153961182, 1.0399240255355835, 1.1779524087905884, + 2.703171968460083, 1.1084223985671997, 1.0640565156936646, + ]), + tolerance: .float32) + } + } + + @Test("tanh") + func test_tanh() throws { + try withIntegrationState(seed: 35559) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.tanh(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.09210336208343506, + minimum: -0.8818567395210266, + maximum: 0.9250870943069458, + absoluteSum: 6.192976474761963, + positionChecksum: 3.220587412516276, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.581309974193573, 0.3747791647911072, 0.3662618398666382, + 0.20108503103256226, -0.43759405612945557, 0.9250870943069458, + ]), + tolerance: .float32) + } + } + + @Test("asin") + func test_asin() throws { + try withIntegrationState(seed: 21542) { + let a = MLXRandom.uniform(low: -0.9, high: 0.9, [4, 3], dtype: .float32) + let result = MLX.asin(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.08627458661794662, + minimum: -0.46629178524017334, + maximum: 0.8269193768501282, + absoluteSum: 4.095224380493164, + positionChecksum: 2.3354612986246743, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.6206415295600891, 0.10426806658506393, -0.29076746106147766, + -0.33669838309288025, 0.8269193768501282, 0.12828460335731506, + ]), + tolerance: .float32) + } + } + + @Test("acos") + func test_acos() throws { + try withIntegrationState(seed: 78281) { + let a = MLXRandom.uniform(low: -0.9, high: 0.9, [4, 3], dtype: .float32) + let result = MLX.acos(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 1.745958685874939, + minimum: 0.736549437046051, + maximum: 2.466407060623169, + absoluteSum: 20.95150375366211, + positionChecksum: 12.173113505045572, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 1.836718201637268, 1.2751185894012451, 2.083163261413574, 1.961542010307312, + 2.466407060623169, 2.0229902267456055, + ]), + tolerance: .float32) + } + } + + @Test("atan") + func test_atan() throws { + try withIntegrationState(seed: 90539) { + let a = MLXRandom.uniform(low: -0.9, high: 0.9, [4, 3], dtype: .float32) + let result = MLX.atan(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.046083711087703705, + minimum: -0.7164856195449829, + maximum: 0.5717471241950989, + absoluteSum: 4.243737697601318, + positionChecksum: 2.1660304069519043, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.33865731954574585, 0.058168165385723114, 0.495836466550827, + 0.28327685594558716, -0.7164856195449829, 0.06804904341697693, + ]), + tolerance: .float32) + } + } + + @Test("asinh") + func test_asinh() throws { + try withIntegrationState(seed: 18773) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.asinh(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.17652754485607147, + minimum: -1.0286109447479248, + maximum: 1.2197173833847046, + absoluteSum: 6.706778049468994, + positionChecksum: 3.318317731221517, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.5257062315940857, -0.8774291276931763, 0.18421567976474762, + 0.0898388996720314, 0.9146624207496643, 0.02543237805366516, + ]), + tolerance: .float32) + } + } + + @Test("acosh") + func test_acosh() throws { + try withIntegrationState(seed: 95620) { + let a = MLXRandom.uniform(low: 1.0, high: 3.0, [4, 3], dtype: .float32) + let result = MLX.acosh(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 1.4068176746368408, + minimum: 0.5282204747200012, + maximum: 1.7533022165298462, + absoluteSum: 16.881811141967773, + positionChecksum: 9.038970947265625, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 1.674250841140747, 1.7533022165298462, 1.330007791519165, 1.561944842338562, + 1.1493288278579712, 1.735342264175415, + ]), + tolerance: .float32) + } + } + + @Test("atanh") + func test_atanh() throws { + try withIntegrationState(seed: 17268) { + let a = MLXRandom.uniform(low: -0.9, high: 0.9, [4, 3], dtype: .float32) + let result = MLX.atanh(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.08335167169570923, + minimum: -0.8625828623771667, + maximum: 1.104719638824463, + absoluteSum: 5.097476005554199, + positionChecksum: 2.553055922190348, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.558688223361969, -0.10614049434661865, -0.8625828623771667, + 0.07864232361316681, -0.030338648706674576, 0.7645785808563232, + ]), + tolerance: .float32) + } + } + + @Test("erf") + func test_erf() throws { + try withIntegrationState(seed: 19849) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.erf(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.1678692102432251, + minimum: -0.9197378754615784, + maximum: 0.9991385340690613, + absoluteSum: 6.304012298583984, + positionChecksum: 3.7706054051717124, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.7421068549156189, 0.12956653535366058, 0.9657787680625916, + 0.40262550115585327, 0.33637773990631104, 0.9991385340690613, + ]), + tolerance: .float32) + } + } + + @Test("erfInverse") + func test_erfInverse() throws { + try withIntegrationState(seed: 83926) { + let a = MLXRandom.uniform(low: -0.9, high: 0.9, [4, 3], dtype: .float32) + let result = MLX.erfInverse(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.016824109479784966, + minimum: -0.750474750995636, + maximum: 1.074080467224121, + absoluteSum: 5.0765275955200195, + positionChecksum: 2.2263415654500327, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.750474750995636, 0.003604303812608123, -0.47791653871536255, + -0.009403870441019535, -0.17815247178077698, 0.3282304108142853, + ]), + tolerance: .float32) + } + } + + @Test("sigmoid") + func test_sigmoid() throws { + try withIntegrationState(seed: 83793) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.sigmoid(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.5078827142715454, + minimum: 0.2062499076128006, + maximum: 0.8442580103874207, + absoluteSum: 6.094592094421387, + positionChecksum: 3.197172164916992, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.5074858069419861, 0.6258241534233093, 0.20864970982074738, + 0.3561042547225952, 0.2062499076128006, 0.8442580103874207, + ]), + tolerance: .float32) + } + } + + @Test("degrees") + func test_degrees() throws { + try withIntegrationState(seed: 24094) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.degrees(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -2.4433441162109375, + minimum: -85.32889556884766, + maximum: 128.0925750732422, + absoluteSum: 596.0240478515625, + positionChecksum: 281.8431396484375, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 128.0925750732422, 36.44943618774414, 30.74724006652832, -70.09713745117188, + 55.25877380371094, -10.830322265625, + ]), + tolerance: .float32) + } + } + + @Test("radians") + func test_radians() throws { + try withIntegrationState(seed: 7077) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.radians(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.0018685739487409592, + minimum: -0.031736548990011215, + maximum: 0.03308562934398651, + absoluteSum: 0.18592886626720428, + positionChecksum: 0.08721178770065308, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.017862409353256226, -0.009540251456201077, 0.03308562934398651, + -0.013581288047134876, -0.011545875109732151, -0.007317093200981617, + ]), + tolerance: .float32) + } + } + + @Test("stopGradient") + func test_stopGradient() throws { + try withIntegrationState(seed: 91151) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.stopGradient(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.26480633020401, + minimum: -0.7478859424591064, + maximum: 1.6692116260528564, + absoluteSum: 6.601988792419434, + positionChecksum: 2.663097381591797, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.8156958222389221, 1.6692116260528564, 0.24522577226161957, + -0.047961488366127014, 0.6883101463317871, -0.17975953221321106, + ]), + tolerance: .float32) + } + } + + @Test("round") + func test_round() throws { + try withIntegrationState(seed: 63999) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 10.0) + let result = MLX.round(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.5, + minimum: -17.0, + maximum: 23.0, + absoluteSum: 98.0, + positionChecksum: 39.166666666666664, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [-14.0, 23.0, -1.0, 2.0, 7.0, 0.0]), + tolerance: .float32) + } + } + + @Test("round/decimals") + func test_round_decimals() throws { + try withIntegrationState(seed: 2520) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 10.0) + let result = MLX.round(a, decimals: 2) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.6791664361953735, + minimum: -16.440000534057617, + maximum: 15.34999942779541, + absoluteSum: 110.8699951171875, + positionChecksum: 59.361663818359375, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -5.789999961853027, -5.799999713897705, -10.619999885559082, + 10.90999984741211, -11.369999885559082, 5.349999904632568, + ]), + tolerance: .float32) + } + } + + @Test("clip") + func test_clip() throws { + try withIntegrationState(seed: 84323) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.clip(a, min: -0.5, max: 0.5) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.20147940516471863, + minimum: -0.5, + maximum: 0.5, + absoluteSum: 4.310338497161865, + positionChecksum: 2.2437593142191568, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.5, -0.5, -0.15267881751060486, -0.22263969480991364, 0.5, + -0.2244129180908203, + ]), + tolerance: .float32) + } + } + + @Test("logicalNot") + func test_logicalNot() throws { + try withIntegrationState(seed: 65221) { + let a = MLXRandom.bernoulli(0.5, [4, 3]) + let result = MLX.logicalNot(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .bool, + mean: 0.0833333358168602, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 1.0, + positionChecksum: 0.5, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [0.0, 0.0, 0.0, 0.0, 0.0, 0.0]), + tolerance: .exact) + } + } + + @Test("isNaN") + func test_isNaN() throws { + try withIntegrationState(seed: 94051) { + let a = MLXArray([1.0, -2.5, Float.infinity, -Float.infinity, Float.nan, 0.0]) + let result = MLX.isNaN(a) + expectSummary( + result, + ArraySummary( + shape: [6], + dtype: .bool, + mean: 0.1666666716337204, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 1.0, + positionChecksum: 0.8333333333333334, + sampleIndices: [0, 1, 2, 3, 4, 5], + samples: [0.0, 0.0, 0.0, 0.0, 1.0, 0.0]), + tolerance: .exact) + } + } + + @Test("isInf") + func test_isInf() throws { + try withIntegrationState(seed: 17171) { + let a = MLXArray([1.0, -2.5, Float.infinity, -Float.infinity, Float.nan, 0.0]) + let result = MLX.isInf(a) + expectSummary( + result, + ArraySummary( + shape: [6], + dtype: .bool, + mean: 0.3333333432674408, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 2.0, + positionChecksum: 1.1666666666666667, + sampleIndices: [0, 1, 2, 3, 4, 5], + samples: [0.0, 0.0, 1.0, 1.0, 0.0, 0.0]), + tolerance: .exact) + } + } + + @Test("isFinite") + func test_isFinite() throws { + try withIntegrationState(seed: 93517) { + let a = MLXArray([1.0, -2.5, Float.infinity, -Float.infinity, Float.nan, 0.0]) + let result = MLX.isFinite(a) + expectSummary( + result, + ArraySummary( + shape: [6], + dtype: .bool, + mean: 0.5, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 3.0, + positionChecksum: 1.5, + sampleIndices: [0, 1, 2, 3, 4, 5], + samples: [1.0, 1.0, 0.0, 0.0, 0.0, 1.0]), + tolerance: .exact) + } + } + + @Test("isPosInf") + func test_isPosInf() throws { + try withIntegrationState(seed: 67816) { + let a = MLXArray([1.0, -2.5, Float.infinity, -Float.infinity, Float.nan, 0.0]) + let result = MLX.isPosInf(a) + expectSummary( + result, + ArraySummary( + shape: [6], + dtype: .bool, + mean: 0.1666666716337204, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 1.0, + positionChecksum: 0.5, + sampleIndices: [0, 1, 2, 3, 4, 5], + samples: [0.0, 0.0, 1.0, 0.0, 0.0, 0.0]), + tolerance: .exact) + } + } + + @Test("isNegInf") + func test_isNegInf() throws { + try withIntegrationState(seed: 31786) { + let a = MLXArray([1.0, -2.5, Float.infinity, -Float.infinity, Float.nan, 0.0]) + let result = MLX.isNegInf(a) + expectSummary( + result, + ArraySummary( + shape: [6], + dtype: .bool, + mean: 0.1666666716337204, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 1.0, + positionChecksum: 0.6666666666666666, + sampleIndices: [0, 1, 2, 3, 4, 5], + samples: [0.0, 0.0, 0.0, 1.0, 0.0, 0.0]), + tolerance: .exact) + } + } + + @Test("nanToNum") + func test_nanToNum() throws { + // arguments are explicit: python defaults posinf/neginf to the dtype max/min while Swift defaults them to 0 + try withIntegrationState(seed: 94828) { + let a = MLXArray([1.0, -2.5, Float.infinity, -Float.infinity, Float.nan, 0.0]) + let result = MLX.nanToNum(a, nan: 1.0, posInf: 100.0, negInf: -100.0) + expectSummary( + result, + ArraySummary( + shape: [6], + dtype: .float32, + mean: -0.0833333358168602, + minimum: -100.0, + maximum: 100.0, + absoluteSum: 204.5, + positionChecksum: 118.5, + sampleIndices: [0, 1, 2, 3, 4, 5], + samples: [1.0, -2.5, 100.0, -100.0, 1.0, 0.0]), + tolerance: .float32) + } + } + + @Test("abs/method") + func test_abs_method() throws { + try withIntegrationState(seed: 59334) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = a.abs() + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 1.0287009477615356, + minimum: 0.007374108768999577, + maximum: 2.067824363708496, + absoluteSum: 12.34441089630127, + positionChecksum: 7.004905700683594, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.6036763191223145, 1.1487308740615845, 0.4937532842159271, + 2.067824363708496, 1.0368303060531616, 2.0537874698638916, + ]), + tolerance: .float32) + } + } + + @Test("exp/method") + func test_exp_method() throws { + try withIntegrationState(seed: 65744) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = a.exp() + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 1.2003825902938843, + minimum: 0.19767798483371735, + maximum: 3.021728992462158, + absoluteSum: 14.404590606689453, + positionChecksum: 7.480013529459636, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.19767798483371735, 1.8948596715927124, 1.4384976625442505, + 3.021728992462158, 0.4999668002128601, 0.38851505517959595, + ]), + tolerance: .float32) + } + } + + @Test("sqrt/method") + func test_sqrt_method() throws { + try withIntegrationState(seed: 3785) { + let a = MLXRandom.uniform(low: 0.1, high: 2.0, [4, 3], dtype: .float32) + let result = a.sqrt() + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 1.1024643182754517, + minimum: 0.7357668280601501, + maximum: 1.4061843156814575, + absoluteSum: 13.229571342468262, + positionChecksum: 7.240408579508464, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 1.4061843156814575, 1.0316053628921509, 1.038780689239502, + 1.0228774547576904, 1.3849389553070068, 0.9444799423217773, + ]), + tolerance: .float32) + } + } + + @Test("round/method") + func test_round_method() throws { + try withIntegrationState(seed: 25776) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 10.0) + let result = a.round() + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.0, + minimum: -20.0, + maximum: 12.0, + absoluteSum: 74.0, + positionChecksum: 39.5, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [-3.0, -4.0, 4.0, 12.0, -0.0, -2.0]), + tolerance: .float32) + } + } + + @Test("asType/float16") + func test_asType_float16() throws { + try withIntegrationState(seed: 396) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = a.asType(.float16) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float16, + mean: 0.3122355341911316, + minimum: -0.763671875, + maximum: 1.630859375, + absoluteSum: 7.469482421875, + positionChecksum: 4.177174886067708, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.335693359375, 1.630859375, 0.72021484375, 0.63427734375, -0.72607421875, + 1.4150390625, + ]), + tolerance: .float16) + } + } + + @Test("asType/int32") + func test_asType_int32() throws { + try withIntegrationState(seed: 68400) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 10.0) + let result = a.asType(.int32) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .int32, + mean: -4.6666669845581055, + minimum: -23.0, + maximum: 11.0, + absoluteSum: 106.0, + positionChecksum: 49.083333333333336, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [11.0, -10.0, -5.0, -3.0, -11.0, 6.0]), + tolerance: .exact) + } + } + + @Test("exp/float16") + func test_exp_float16() throws { + try withIntegrationState(seed: 43848) { + let a = MLXRandom.normal([4, 3], dtype: .float16, loc: 0.0, scale: 1.0) + let result = MLX.exp(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float16, + mean: 1.9106242656707764, + minimum: 0.174072265625, + maximum: 6.24609375, + absoluteSum: 22.927490234375, + positionChecksum: 12.168294270833334, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 1.255859375, 2.857421875, 1.10546875, 0.388671875, 1.515625, 1.5947265625, + ]), + tolerance: .float16) + } + } + + @Test("exp/bfloat16") + func test_exp_bfloat16() throws { + try withIntegrationState(seed: 82777) { + let a = MLXRandom.normal([4, 3], dtype: .bfloat16, loc: 0.0, scale: 1.0) + let result = MLX.exp(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .bfloat16, + mean: 2.4685873985290527, + minimum: 0.21484375, + maximum: 9.625, + absoluteSum: 29.623046875, + positionChecksum: 22.324381510416668, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [0.494140625, 2.140625, 0.6796875, 0.84375, 0.7890625, 8.9375]), + tolerance: .float16) + } + } + + @Test("positive") + func test_positive() throws { + try withIntegrationState(seed: 39984) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.positive(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.24678385257720947, + minimum: -1.4747556447982788, + maximum: 1.6195077896118164, + absoluteSum: 8.60529613494873, + positionChecksum: 4.061902681986491, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.7380889058113098, -1.2995355129241943, 0.5283272862434387, + 0.3148530423641205, 0.2088206708431244, -0.6653445363044739, + ]), + tolerance: .float32) + } + } + + @Test("view/int16") + func test_view_int16() throws { + try withIntegrationState(seed: 63133) { + let a = MLXRandom.randInt(low: 0, high: 1024, [4, 3], type: Int32.self) + let result = MLX.view(a, dtype: .int16) + expectSummary( + result, + ArraySummary( + shape: [4, 6], + dtype: .int16, + mean: 213.625, + minimum: 0.0, + maximum: 919.0, + absoluteSum: 5127.0, + positionChecksum: 2602.625, + sampleIndices: [0, 5, 9, 14, 18, 23], + samples: [22.0, 0.0, 0.0, 863.0, 255.0, 0.0]), + tolerance: .exact) + } + } + + @Test("contiguous") + func test_contiguous() throws { + try withIntegrationState(seed: 69937) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.contiguous(a.T) + expectSummary( + result, + ArraySummary( + shape: [3, 4], + dtype: .float32, + mean: -0.5083404779434204, + minimum: -2.389263391494751, + maximum: 0.8913607597351074, + absoluteSum: 9.336750984191895, + positionChecksum: 4.510031382242839, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.004110712558031082, -2.389263391494751, -0.3058277368545532, + -0.40270015597343445, -0.49763554334640503, 0.01963740773499012, + ]), + tolerance: .float32) + } + } + + @Test("isClose") + func test_isClose() throws { + try withIntegrationState(seed: 90080) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.isClose(a, b, rtol: 0.5, atol: 0.1) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .bool, + mean: 0.25, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 3.0, + positionChecksum: 1.6666666666666667, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [0.0, 1.0, 0.0, 0.0, 1.0, 0.0]), + tolerance: .exact) + } + } + + @Test("allClose") + func test_allClose() throws { + try withIntegrationState(seed: 65428) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.allClose(a, a + 0.001, rtol: 0.01) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .bool, + mean: 0.0, + minimum: 0.0, + maximum: 0.0, + absoluteSum: 0.0, + positionChecksum: 0.0, + sampleIndices: [0], + samples: [0.0]), + tolerance: .exact) + } + } + + @Test("arrayEqual") + func test_arrayEqual() throws { + try withIntegrationState(seed: 84233) { + let a = MLXRandom.randInt(low: 0, high: 4, [4, 3], type: Int32.self) + let result = MLX.arrayEqual(a, a) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .bool, + mean: 1.0, + minimum: 1.0, + maximum: 1.0, + absoluteSum: 1.0, + positionChecksum: 1.0, + sampleIndices: [0], + samples: [1.0]), + tolerance: .exact) + } + } + + @Test("arrayEqual/false") + func test_arrayEqual_false() throws { + try withIntegrationState(seed: 90639) { + let a = MLXRandom.randInt(low: 0, high: 4, [4, 3], type: Int32.self) + let b = MLXRandom.randInt(low: 0, high: 4, [4, 3], type: Int32.self) + let result = MLX.arrayEqual(a, b) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .bool, + mean: 0.0, + minimum: 0.0, + maximum: 0.0, + absoluteSum: 0.0, + positionChecksum: 0.0, + sampleIndices: [0], + samples: [0.0]), + tolerance: .exact) + } + } + + @Test("realPart") + func test_realPart() throws { + // Swift exposes realPart() as a method, not a free function + try withIntegrationState(seed: 28164) { + let r = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let i = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let c = r + i.asImaginary() + let result = c.realPart() + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.27432000637054443, + minimum: -2.3572521209716797, + maximum: 0.864957869052887, + absoluteSum: 8.50927734375, + positionChecksum: 3.2805039087931314, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -2.3572521209716797, -0.278932124376297, 0.36380791664123535, + 0.6124622225761414, 0.09972750395536423, -0.07274171710014343, + ]), + tolerance: .float32) + } + } + + @Test("imaginaryPart") + func test_imaginaryPart() throws { + try withIntegrationState(seed: 58944) { + let r = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let i = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let c = r + i.asImaginary() + let result = c.imaginaryPart() + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.4124636650085449, + minimum: -1.5858224630355835, + maximum: 1.1684366464614868, + absoluteSum: 9.83724308013916, + positionChecksum: 6.4064591725667315, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.04229888319969177, 0.016773255541920662, -0.5229106545448303, + -1.2519052028656006, -1.4781920909881592, 0.7881496548652649, + ]), + tolerance: .float32) + } + } + + @Test("conjugate") + func test_conjugate() throws { + try withIntegrationState(seed: 29664) { + let r = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let i = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let c = r + i.asImaginary() + let result = MLX.conjugate(c) + expectSummary( + result.realPart(), + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.23909756541252136, + minimum: -0.8671482801437378, + maximum: 1.3631097078323364, + absoluteSum: 9.518288612365723, + positionChecksum: 5.316549301147461, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.7143032550811768, -0.7656220197677612, 0.36717352271080017, + 1.3309167623519897, -0.14781005680561066, -0.8671482801437378, + ]), + tolerance: .float32) + expectSummary( + result.imaginaryPart(), + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.04297981783747673, + minimum: -1.5997480154037476, + maximum: 1.4610912799835205, + absoluteSum: 10.585673332214355, + positionChecksum: 5.1464080810546875, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -1.5997480154037476, -0.7909235954284668, -1.143326997756958, + 1.0923233032226562, 1.4610912799835205, 1.02470064163208, + ]), + tolerance: .float32) + } + } + + @Test("asImaginary") + func test_asImaginary() throws { + // verifies the complex value the other complex cases are built from + try withIntegrationState(seed: 52131) { + let r = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let i = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let c = r + i.asImaginary() + let result = c + expectSummary( + result.realPart(), + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.4877490997314453, + minimum: -1.8236697912216187, + maximum: 0.596900999546051, + absoluteSum: 7.777519226074219, + positionChecksum: 3.4768377939860025, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.25812309980392456, -1.8236697912216187, -1.0013736486434937, + -0.42360997200012207, 0.10724081844091415, -0.371466726064682, + ]), + tolerance: .float32) + expectSummary( + result.imaginaryPart(), + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.037692904472351074, + minimum: -1.3910119533538818, + maximum: 1.7863659858703613, + absoluteSum: 9.36073112487793, + positionChecksum: 4.2958984375, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 1.7863659858703613, 0.4858953356742859, 1.0185210704803467, + 1.220772624015808, 0.1476796418428421, 0.20525570213794708, + ]), + tolerance: .float32) + } + } +} diff --git a/Tests/MLXIntegrationTests/Generated/GeneratedFFTTests.swift b/Tests/MLXIntegrationTests/Generated/GeneratedFFTTests.swift new file mode 100644 index 000000000..32f6eafeb --- /dev/null +++ b/Tests/MLXIntegrationTests/Generated/GeneratedFFTTests.swift @@ -0,0 +1,1513 @@ +// Copyright © 2026 Apple Inc. +// +// GENERATED by tools/integration_tests -- DO NOT EDIT. +// +// Values in this file were produced by python mlx and are compared against +// the Swift API. Regenerate with: +// +// python3 tools/integration_tests/generate.py +// +// python mlx: 0.32.2.dev20260910+1f8e74e3f +// vendored mlx (Cmlx): 0.32.2 +// generator revision: 2 +// cases: 41 + +import Foundation +import MLX +import Testing + +@Suite("generated: FFT") +struct GeneratedFFTTests { + + @Test("fft") + func test_fft() throws { + try withIntegrationState(seed: 3161) { + let r = MLXRandom.normal([100], dtype: .float32, loc: 0.0, scale: 1.0) + let i = MLXRandom.normal([100], dtype: .float32, loc: 0.0, scale: 1.0) + let c = r + i.asImaginary() + let result = MLX.fft(c) + expectSummary( + result.realPart(), + ArraySummary( + shape: [100], + dtype: .float32, + mean: -0.30193811655044556, + minimum: -26.834150314331055, + maximum: 25.471298217773438, + absoluteSum: 748.5565795898438, + positionChecksum: 364.7257421875, + sampleIndices: [0, 20, 40, 59, 79, 99], + samples: [ + -17.864389419555664, 0.47785067558288574, -7.920133590698242, + 1.171248197555542, 5.3682661056518555, -11.829191207885742, + ]), + tolerance: .float32) + expectSummary( + result.imaginaryPart(), + ArraySummary( + shape: [100], + dtype: .float32, + mean: -2.2110860347747803, + minimum: -33.87453842163086, + maximum: 27.859336853027344, + absoluteSum: 759.38427734375, + positionChecksum: 398.85734375, + sampleIndices: [0, 20, 40, 59, 79, 99], + samples: [ + -10.048587799072266, 0.7990264892578125, 2.4799821376800537, + 8.100057601928711, -5.986302375793457, -4.987302303314209, + ]), + tolerance: .float32) + } + } + + @Test("fft/nShort") + func test_fft_nShort() throws { + try withIntegrationState(seed: 2147) { + let r = MLXRandom.normal([100], dtype: .float32, loc: 0.0, scale: 1.0) + let i = MLXRandom.normal([100], dtype: .float32, loc: 0.0, scale: 1.0) + let c = r + i.asImaginary() + let result = MLX.fft(c, n: 80) + expectSummary( + result.realPart(), + ArraySummary( + shape: [80], + dtype: .float32, + mean: -0.30215585231781006, + minimum: -19.026569366455078, + maximum: 23.793720245361328, + absoluteSum: 603.0272216796875, + positionChecksum: 314.046435546875, + sampleIndices: [0, 16, 32, 47, 63, 79], + samples: [ + 8.233772277832031, -4.0249857902526855, 2.5544846057891846, + -9.478172302246094, 3.6788711547851562, 5.409932613372803, + ]), + tolerance: .float32) + expectSummary( + result.imaginaryPart(), + ArraySummary( + shape: [80], + dtype: .float32, + mean: -0.330033540725708, + minimum: -17.620588302612305, + maximum: 18.184288024902344, + absoluteSum: 535.7166748046875, + positionChecksum: 275.080224609375, + sampleIndices: [0, 16, 32, 47, 63, 79], + samples: [ + -4.2741594314575195, -5.628664493560791, -7.816094875335693, + 7.278494358062744, -11.146451950073242, -5.636445999145508, + ]), + tolerance: .float32) + } + } + + @Test("fft/nLong") + func test_fft_nLong() throws { + try withIntegrationState(seed: 27280) { + let r = MLXRandom.normal([100], dtype: .float32, loc: 0.0, scale: 1.0) + let i = MLXRandom.normal([100], dtype: .float32, loc: 0.0, scale: 1.0) + let c = r + i.asImaginary() + let result = MLX.fft(c, n: 120) + expectSummary( + result.realPart(), + ArraySummary( + shape: [120], + dtype: .float32, + mean: 1.245316982269287, + minimum: -24.74895668029785, + maximum: 26.73784065246582, + absoluteSum: 952.5072021484375, + positionChecksum: 484.21611328125, + sampleIndices: [0, 24, 48, 71, 95, 119], + samples: [ + 17.85637664794922, 7.068179607391357, 3.323228120803833, 18.420766830444336, + 10.24211311340332, -7.67423152923584, + ]), + tolerance: .float32) + expectSummary( + result.imaginaryPart(), + ArraySummary( + shape: [120], + dtype: .float32, + mean: 0.865074634552002, + minimum: -23.82689666748047, + maximum: 21.837039947509766, + absoluteSum: 956.6898803710938, + positionChecksum: 484.05237630208336, + sampleIndices: [0, 24, 48, 71, 95, 119], + samples: [ + 9.100713729858398, -9.670844078063965, -3.221973419189453, + -2.571913719177246, -14.796401023864746, 15.20323371887207, + ]), + tolerance: .float32) + } + } + + @Test("fft/axis") + func test_fft_axis() throws { + try withIntegrationState(seed: 67284) { + let r = MLXRandom.normal([10, 10], dtype: .float32, loc: 0.0, scale: 1.0) + let i = MLXRandom.normal([10, 10], dtype: .float32, loc: 0.0, scale: 1.0) + let c = r + i.asImaginary() + let result = MLX.fft(c, axis: 0) + expectSummary( + result.realPart(), + ArraySummary( + shape: [10, 10], + dtype: .float32, + mean: -0.465768426656723, + minimum: -9.097143173217773, + maximum: 7.025123119354248, + absoluteSum: 236.0032501220703, + positionChecksum: 112.51224609375, + sampleIndices: [0, 20, 40, 59, 79, 99], + samples: [ + -5.653952598571777, -1.792625904083252, -3.0560266971588135, + -1.7548167705535889, 3.503523111343384, -1.5613975524902344, + ]), + tolerance: .float32) + expectSummary( + result.imaginaryPart(), + ArraySummary( + shape: [10, 10], + dtype: .float32, + mean: -0.3379652798175812, + minimum: -7.193725109100342, + maximum: 7.352320671081543, + absoluteSum: 243.55081176757812, + positionChecksum: 124.103173828125, + sampleIndices: [0, 20, 40, 59, 79, 99], + samples: [ + -6.474360466003418, -1.6081452369689941, 3.647308588027954, + -0.7049122452735901, 0.582835853099823, 1.1408429145812988, + ]), + tolerance: .float32) + } + } + + @Test("fft/ortho") + func test_fft_ortho() throws { + try withIntegrationState(seed: 41079) { + let r = MLXRandom.normal([100], dtype: .float32, loc: 0.0, scale: 1.0) + let i = MLXRandom.normal([100], dtype: .float32, loc: 0.0, scale: 1.0) + let c = r + i.asImaginary() + let result = MLX.fft(c, norm: .ortho) + expectSummary( + result.realPart(), + ArraySummary( + shape: [100], + dtype: .float32, + mean: -0.1377011239528656, + minimum: -2.420893907546997, + maximum: 2.4660956859588623, + absoluteSum: 89.8805160522461, + positionChecksum: 44.410625, + sampleIndices: [0, 20, 40, 59, 79, 99], + samples: [ + 0.9301566481590271, 1.4401391744613647, -1.4074681997299194, + -2.0820672512054443, -0.47943899035453796, -0.2989818751811981, + ]), + tolerance: .float32) + expectSummary( + result.imaginaryPart(), + ArraySummary( + shape: [100], + dtype: .float32, + mean: 0.0205818060785532, + minimum: -2.0081610679626465, + maximum: 2.822983503341675, + absoluteSum: 78.08838653564453, + positionChecksum: 37.138740234375, + sampleIndices: [0, 20, 40, 59, 79, 99], + samples: [ + -1.1429274082183838, 1.4222431182861328, 0.8448073267936707, + 0.45329347252845764, 1.7841476202011108, 0.44647422432899475, + ]), + tolerance: .float32) + } + } + + @Test("ifft") + func test_ifft() throws { + try withIntegrationState(seed: 92555) { + let r = MLXRandom.normal([100], dtype: .float32, loc: 0.0, scale: 1.0) + let i = MLXRandom.normal([100], dtype: .float32, loc: 0.0, scale: 1.0) + let c = r + i.asImaginary() + let result = MLX.ifft(c) + expectSummary( + result.realPart(), + ArraySummary( + shape: [100], + dtype: .float32, + mean: 0.016596531495451927, + minimum: -0.23650917410850525, + maximum: 0.2672997713088989, + absoluteSum: 8.624792098999023, + positionChecksum: 4.234874267578125, + sampleIndices: [0, 20, 40, 59, 79, 99], + samples: [ + -0.010019684210419655, 0.14668859541416168, 0.11429940909147263, + -0.045612435787916183, 0.14252117276191711, 0.036860380321741104, + ]), + tolerance: .float32) + expectSummary( + result.imaginaryPart(), + ArraySummary( + shape: [100], + dtype: .float32, + mean: 0.018229888752102852, + minimum: -0.23968389630317688, + maximum: 0.2401280254125595, + absoluteSum: 8.7319917678833, + positionChecksum: 4.291281127929688, + sampleIndices: [0, 20, 40, 59, 79, 99], + samples: [ + -0.0015521335881203413, -0.17396560311317444, 0.0605366975069046, + -0.03737438842654228, 0.04746906831860542, -0.12103308737277985, + ]), + tolerance: .float32) + } + } + + @Test("ifft/nShort") + func test_ifft_nShort() throws { + try withIntegrationState(seed: 12331) { + let r = MLXRandom.normal([100], dtype: .float32, loc: 0.0, scale: 1.0) + let i = MLXRandom.normal([100], dtype: .float32, loc: 0.0, scale: 1.0) + let c = r + i.asImaginary() + let result = MLX.ifft(c, n: 80) + expectSummary( + result.realPart(), + ArraySummary( + shape: [80], + dtype: .float32, + mean: -0.019246213138103485, + minimum: -0.3084016740322113, + maximum: 0.20837044715881348, + absoluteSum: 7.586906909942627, + positionChecksum: 3.7272315979003907, + sampleIndices: [0, 16, 32, 47, 63, 79], + samples: [ + -0.06569941341876984, 0.08424773812294006, 0.13474620878696442, + 0.1767660677433014, -0.20869246125221252, 0.027556825429201126, + ]), + tolerance: .float32) + expectSummary( + result.imaginaryPart(), + ArraySummary( + shape: [80], + dtype: .float32, + mean: 0.016126060858368874, + minimum: -0.20641577243804932, + maximum: 0.277831107378006, + absoluteSum: 7.127358436584473, + positionChecksum: 3.7531387329101564, + sampleIndices: [0, 16, 32, 47, 63, 79], + samples: [ + -0.05102470517158508, -0.11271238327026367, 0.026768634095788002, + 0.09892851114273071, -0.12387256324291229, 0.01915312185883522, + ]), + tolerance: .float32) + } + } + + @Test("ifft/nLong") + func test_ifft_nLong() throws { + try withIntegrationState(seed: 38177) { + let r = MLXRandom.normal([100], dtype: .float32, loc: 0.0, scale: 1.0) + let i = MLXRandom.normal([100], dtype: .float32, loc: 0.0, scale: 1.0) + let c = r + i.asImaginary() + let result = MLX.ifft(c, n: 120) + expectSummary( + result.realPart(), + ArraySummary( + shape: [120], + dtype: .float32, + mean: -0.004438468720763922, + minimum: -0.2090398520231247, + maximum: 0.2671051621437073, + absoluteSum: 8.520634651184082, + positionChecksum: 4.005149841308594, + sampleIndices: [0, 24, 48, 71, 95, 119], + samples: [ + 0.07100244611501694, 0.07003705203533173, 0.2671051621437073, + -0.2090398520231247, -0.09596047550439835, 0.02532544918358326, + ]), + tolerance: .float32) + expectSummary( + result.imaginaryPart(), + ArraySummary( + shape: [120], + dtype: .float32, + mean: -0.0033197151497006416, + minimum: -0.2966475784778595, + maximum: 0.22516706585884094, + absoluteSum: 8.153952598571777, + positionChecksum: 3.686529286702474, + sampleIndices: [0, 24, 48, 71, 95, 119], + samples: [ + -0.12483212351799011, 0.09503603726625443, 0.016529258340597153, + 0.10363555699586868, 0.10037048161029816, -0.031290605664253235, + ]), + tolerance: .float32) + } + } + + @Test("ifft/axis") + func test_ifft_axis() throws { + try withIntegrationState(seed: 92391) { + let r = MLXRandom.normal([10, 10], dtype: .float32, loc: 0.0, scale: 1.0) + let i = MLXRandom.normal([10, 10], dtype: .float32, loc: 0.0, scale: 1.0) + let c = r + i.asImaginary() + let result = MLX.ifft(c, axis: 0) + expectSummary( + result.realPart(), + ArraySummary( + shape: [10, 10], + dtype: .float32, + mean: -0.03038243018090725, + minimum: -0.9769951701164246, + maximum: 0.7989615797996521, + absoluteSum: 27.44314193725586, + positionChecksum: 14.012554931640626, + sampleIndices: [0, 20, 40, 59, 79, 99], + samples: [ + 0.37563377618789673, 0.14058302342891693, 0.17051009833812714, + 0.10858341306447983, 0.29239600896835327, -0.06135018542408943, + ]), + tolerance: .float32) + expectSummary( + result.imaginaryPart(), + ArraySummary( + shape: [10, 10], + dtype: .float32, + mean: -0.07071053236722946, + minimum: -0.8794716596603394, + maximum: 0.5949147343635559, + absoluteSum: 24.313308715820312, + positionChecksum: 13.0312890625, + sampleIndices: [0, 20, 40, 59, 79, 99], + samples: [ + -0.008079779334366322, 0.014237630181014538, -0.20407943427562714, + 0.25475987792015076, -0.5532649755477905, -0.37195757031440735, + ]), + tolerance: .float32) + } + } + + @Test("ifft/ortho") + func test_ifft_ortho() throws { + try withIntegrationState(seed: 8262) { + let r = MLXRandom.normal([100], dtype: .float32, loc: 0.0, scale: 1.0) + let i = MLXRandom.normal([100], dtype: .float32, loc: 0.0, scale: 1.0) + let c = r + i.asImaginary() + let result = MLX.ifft(c, norm: .ortho) + expectSummary( + result.realPart(), + ArraySummary( + shape: [100], + dtype: .float32, + mean: -0.14974907040596008, + minimum: -2.6609697341918945, + maximum: 2.1586873531341553, + absoluteSum: 76.76493835449219, + positionChecksum: 35.6903564453125, + sampleIndices: [0, 20, 40, 59, 79, 99], + samples: [ + -0.7524006366729736, -0.7424914240837097, -1.2931209802627563, + -0.3934279680252075, 0.7683846354484558, 0.5575725436210632, + ]), + tolerance: .float32) + expectSummary( + result.imaginaryPart(), + ArraySummary( + shape: [100], + dtype: .float32, + mean: -0.1610265076160431, + minimum: -2.565251588821411, + maximum: 2.1245665550231934, + absoluteSum: 91.36912536621094, + positionChecksum: 48.544853515625, + sampleIndices: [0, 20, 40, 59, 79, 99], + samples: [ + 0.05022697523236275, -0.2160770297050476, -2.0412232875823975, + -0.921349287033081, -1.5685405731201172, -2.192484140396118, + ]), + tolerance: .float32) + } + } + + @Test("rfft") + func test_rfft() throws { + try withIntegrationState(seed: 9717) { + let c = MLXRandom.normal([100], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.rfft(c) + expectSummary( + result.realPart(), + ArraySummary( + shape: [51], + dtype: .float32, + mean: 0.08805207908153534, + minimum: -14.910188674926758, + maximum: 16.072185516357422, + absoluteSum: 306.84942626953125, + positionChecksum: 164.0440793504902, + sampleIndices: [0, 10, 20, 30, 40, 50], + samples: [ + -3.9297866821289062, -7.619717121124268, -0.037673234939575195, + 4.216407775878906, 1.6698544025421143, 9.09442138671875, + ]), + tolerance: .float32) + expectSummary( + result.imaginaryPart(), + ArraySummary( + shape: [51], + dtype: .float32, + mean: 0.4946885406970978, + minimum: -14.891788482666016, + maximum: 12.950984954833984, + absoluteSum: 284.84588623046875, + positionChecksum: 144.68315333946077, + sampleIndices: [0, 10, 20, 30, 40, 50], + samples: [ + 0.0, 3.381943702697754, -3.114729404449463, 2.588118076324463, + -1.9571096897125244, 0.0, + ]), + tolerance: .float32) + } + } + + @Test("rfft/n") + func test_rfft_n() throws { + try withIntegrationState(seed: 39379) { + let c = MLXRandom.normal([100], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.rfft(c, n: 80) + expectSummary( + result.realPart(), + ArraySummary( + shape: [41], + dtype: .float32, + mean: 1.0768029689788818, + minimum: -9.284590721130371, + maximum: 11.904441833496094, + absoluteSum: 186.73614501953125, + positionChecksum: 107.89291158536585, + sampleIndices: [0, 8, 16, 24, 32, 40], + samples: [ + 5.147792816162109, -0.7290818691253662, 2.711458206176758, + -9.033034324645996, -4.209567070007324, 8.454961776733398, + ]), + tolerance: .float32) + expectSummary( + result.imaginaryPart(), + ArraySummary( + shape: [41], + dtype: .float32, + mean: -0.010239856317639351, + minimum: -14.581356048583984, + maximum: 16.252479553222656, + absoluteSum: 198.701171875, + positionChecksum: 104.59896627286585, + sampleIndices: [0, 8, 16, 24, 32, 40], + samples: [ + 0.0, 1.6902306079864502, 3.0469067096710205, 4.965663433074951, + -2.1530351638793945, 0.0, + ]), + tolerance: .float32) + } + } + + @Test("rfft/forward") + func test_rfft_forward() throws { + try withIntegrationState(seed: 29470) { + let c = MLXRandom.normal([100], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.rfft(c, norm: .forward) + expectSummary( + result.realPart(), + ArraySummary( + shape: [51], + dtype: .float32, + mean: -0.008272930979728699, + minimum: -0.18479548394680023, + maximum: 0.1636444628238678, + absoluteSum: 2.7358570098876953, + positionChecksum: 1.4796739466050093, + sampleIndices: [0, 10, 20, 30, 40, 50], + samples: [ + -0.01737872324883938, -0.0004667544271796942, -0.0013617968652397394, + -0.09146632254123688, -0.08097280561923981, 0.1343451887369156, + ]), + tolerance: .float32) + expectSummary( + result.imaginaryPart(), + ArraySummary( + shape: [51], + dtype: .float32, + mean: 0.022197062149643898, + minimum: -0.13841067254543304, + maximum: 0.19061492383480072, + absoluteSum: 2.6014790534973145, + positionChecksum: 1.4728336708218444, + sampleIndices: [0, 10, 20, 30, 40, 50], + samples: [ + 0.0, 0.011355984024703503, 0.03543727844953537, -0.13841067254543304, + 0.0947970375418663, 0.0, + ]), + tolerance: .float32) + } + } + + @Test("irfft") + func test_irfft() throws { + try withIntegrationState(seed: 45651) { + let r = MLXRandom.normal([100], dtype: .float32, loc: 0.0, scale: 1.0) + let i = MLXRandom.normal([100], dtype: .float32, loc: 0.0, scale: 1.0) + let c = r + i.asImaginary() + let result = MLX.irfft(c) + expectSummary( + result, + ArraySummary( + shape: [198], + dtype: .float32, + mean: -0.009573065675795078, + minimum: -0.2509799599647522, + maximum: 0.19757866859436035, + absoluteSum: 14.547061920166016, + positionChecksum: 7.517489346590909, + sampleIndices: [0, 39, 79, 118, 158, 197], + samples: [ + -0.11132393777370453, 0.011129573918879032, 0.02290239743888378, + -0.10837771743535995, -0.03477989509701729, -0.05752992257475853, + ]), + tolerance: .float32) + } + } + + @Test("irfft/n") + func test_irfft_n() throws { + try withIntegrationState(seed: 16439) { + let r = MLXRandom.normal([100], dtype: .float32, loc: 0.0, scale: 1.0) + let i = MLXRandom.normal([100], dtype: .float32, loc: 0.0, scale: 1.0) + let c = r + i.asImaginary() + let result = MLX.irfft(c, n: 120) + expectSummary( + result, + ArraySummary( + shape: [120], + dtype: .float32, + mean: 0.014314486645162106, + minimum: -0.28694266080856323, + maximum: 0.346214234828949, + absoluteSum: 13.746326446533203, + positionChecksum: 6.66098378499349, + sampleIndices: [0, 24, 48, 71, 95, 119], + samples: [ + 0.09791383892297745, -0.020166436210274696, -0.1698228269815445, + 0.21948356926441193, -0.07427943497896194, -0.10716888308525085, + ]), + tolerance: .float32) + } + } + + @Test("fft2") + func test_fft2() throws { + try withIntegrationState(seed: 35659) { + let r = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let i = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let c = r + i.asImaginary() + let result = MLX.fft2(c) + expectSummary( + result.realPart(), + ArraySummary( + shape: [8, 8, 8], + dtype: .float32, + mean: -0.7178468704223633, + minimum: -30.726062774658203, + maximum: 22.191179275512695, + absoluteSum: 3245.3369140625, + positionChecksum: 1631.6517333984375, + sampleIndices: [0, 102, 204, 307, 409, 511], + samples: [ + 3.754521369934082, -6.490096569061279, 4.651754856109619, + 14.567571640014648, 8.419650077819824, -0.5989406108856201, + ]), + tolerance: .float32) + expectSummary( + result.imaginaryPart(), + ArraySummary( + shape: [8, 8, 8], + dtype: .float32, + mean: 0.3027850389480591, + minimum: -23.579471588134766, + maximum: 25.952545166015625, + absoluteSum: 3262.20458984375, + positionChecksum: 1579.3980712890625, + sampleIndices: [0, 102, 204, 307, 409, 511], + samples: [ + 0.2110457420349121, 5.290388107299805, -2.667041301727295, + 1.3738980293273926, 12.329010963439941, -5.085733413696289, + ]), + tolerance: .float32) + } + } + + @Test("fft2/s") + func test_fft2_s() throws { + try withIntegrationState(seed: 49089) { + let r = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let i = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let c = r + i.asImaginary() + let result = MLX.fft2(c, s: [3, 4]) + expectSummary( + result.realPart(), + ArraySummary( + shape: [8, 3, 4], + dtype: .float32, + mean: -0.44498777389526367, + minimum: -7.310737609863281, + maximum: 7.252446174621582, + absoluteSum: 287.83111572265625, + positionChecksum: 149.67510986328125, + sampleIndices: [0, 19, 38, 57, 76, 95], + samples: [ + -2.878303050994873, -1.4361096620559692, 0.4035429358482361, + 6.749013900756836, 1.0702213048934937, -3.8870134353637695, + ]), + tolerance: .float32) + expectSummary( + result.imaginaryPart(), + ArraySummary( + shape: [8, 3, 4], + dtype: .float32, + mean: 0.6117333173751831, + minimum: -7.615391254425049, + maximum: 8.767561912536621, + absoluteSum: 268.90863037109375, + positionChecksum: 147.26704915364584, + sampleIndices: [0, 19, 38, 57, 76, 95], + samples: [ + -1.313608169555664, 3.5226364135742188, -4.22772216796875, + -0.37668323516845703, 8.767561912536621, -5.13897180557251, + ]), + tolerance: .float32) + } + } + + @Test("fft2/axes") + func test_fft2_axes() throws { + try withIntegrationState(seed: 55820) { + let r = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let i = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let c = r + i.asImaginary() + let result = MLX.fft2(c, axes: [0, 2]) + expectSummary( + result.realPart(), + ArraySummary( + shape: [8, 8, 8], + dtype: .float32, + mean: -0.14448444545269012, + minimum: -32.98771667480469, + maximum: 22.09111785888672, + absoluteSum: 3448.20849609375, + positionChecksum: 1765.2010498046875, + sampleIndices: [0, 102, 204, 307, 409, 511], + samples: [ + -0.04199075698852539, 2.22198486328125, -0.7248520851135254, + -4.36610221862793, -14.677776336669922, -9.698610305786133, + ]), + tolerance: .float32) + expectSummary( + result.imaginaryPart(), + ArraySummary( + shape: [8, 8, 8], + dtype: .float32, + mean: 0.6306827068328857, + minimum: -23.172414779663086, + maximum: 25.345382690429688, + absoluteSum: 3380.444580078125, + positionChecksum: 1648.940185546875, + sampleIndices: [0, 102, 204, 307, 409, 511], + samples: [ + 10.400022506713867, -4.030159950256348, 11.986205101013184, + 3.302541732788086, -2.0049829483032227, 1.248794436454773, + ]), + tolerance: .float32) + } + } + + @Test("ifft2") + func test_ifft2() throws { + try withIntegrationState(seed: 56653) { + let r = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let i = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let c = r + i.asImaginary() + let result = MLX.ifft2(c) + expectSummary( + result.realPart(), + ArraySummary( + shape: [8, 8, 8], + dtype: .float32, + mean: -0.0019193659536540508, + minimum: -0.36529168486595154, + maximum: 0.41034168004989624, + absoluteSum: 50.309425354003906, + positionChecksum: 25.473011016845703, + sampleIndices: [0, 102, 204, 307, 409, 511], + samples: [ + -0.004118995741009712, -0.17179054021835327, -0.09692619740962982, + -0.009604329243302345, -0.15581375360488892, 0.12122435122728348, + ]), + tolerance: .float32) + expectSummary( + result.imaginaryPart(), + ArraySummary( + shape: [8, 8, 8], + dtype: .float32, + mean: 0.004310143645852804, + minimum: -0.3949809670448303, + maximum: 0.33320337533950806, + absoluteSum: 51.721160888671875, + positionChecksum: 26.324140548706055, + sampleIndices: [0, 102, 204, 307, 409, 511], + samples: [ + 0.010328251868486404, 0.20638486742973328, -0.15939217805862427, + 0.16267183423042297, 0.06081516668200493, 0.11304222047328949, + ]), + tolerance: .float32) + } + } + + @Test("ifft2/s") + func test_ifft2_s() throws { + try withIntegrationState(seed: 62821) { + let r = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let i = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let c = r + i.asImaginary() + let result = MLX.ifft2(c, s: [3, 4]) + expectSummary( + result.realPart(), + ArraySummary( + shape: [8, 3, 4], + dtype: .float32, + mean: 0.011248579248785973, + minimum: -0.7063571214675903, + maximum: 0.7444247007369995, + absoluteSum: 23.008529663085938, + positionChecksum: 11.561017354329428, + sampleIndices: [0, 19, 38, 57, 76, 95], + samples: [ + 0.21029286086559296, -0.2403487265110016, -0.21712619066238403, + -0.17577317357063293, -0.4854602813720703, -0.17871509492397308, + ]), + tolerance: .float32) + expectSummary( + result.imaginaryPart(), + ArraySummary( + shape: [8, 3, 4], + dtype: .float32, + mean: 0.046339940279722214, + minimum: -0.6162391304969788, + maximum: 0.8353030681610107, + absoluteSum: 18.841552734375, + positionChecksum: 9.221977233886719, + sampleIndices: [0, 19, 38, 57, 76, 95], + samples: [ + 0.17300370335578918, 0.613670825958252, 0.1690458208322525, + -0.17819438874721527, -0.15327095985412598, -0.05432131886482239, + ]), + tolerance: .float32) + } + } + + @Test("ifft2/axes") + func test_ifft2_axes() throws { + try withIntegrationState(seed: 15965) { + let r = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let i = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let c = r + i.asImaginary() + let result = MLX.ifft2(c, axes: [0, 2]) + expectSummary( + result.realPart(), + ArraySummary( + shape: [8, 8, 8], + dtype: .float32, + mean: -0.001953849568963051, + minimum: -0.4720636308193207, + maximum: 0.5106112360954285, + absoluteSum: 53.05870819091797, + positionChecksum: 27.151283264160156, + sampleIndices: [0, 102, 204, 307, 409, 511], + samples: [ + -0.0033838478848338127, -0.10563762485980988, -0.022997520864009857, + 0.1938212811946869, -0.20600074529647827, -0.08471836894750595, + ]), + tolerance: .float32) + expectSummary( + result.imaginaryPart(), + ArraySummary( + shape: [8, 8, 8], + dtype: .float32, + mean: -0.00011229666415601969, + minimum: -0.2937026619911194, + maximum: 0.40563082695007324, + absoluteSum: 47.61387634277344, + positionChecksum: 24.217378616333008, + sampleIndices: [0, 102, 204, 307, 409, 511], + samples: [ + 0.09508480131626129, -0.14579325914382935, -0.18434379994869232, + 0.11046306043863297, 0.020605478435754776, 0.0009845122694969177, + ]), + tolerance: .float32) + } + } + + @Test("fftn") + func test_fftn() throws { + try withIntegrationState(seed: 20276) { + let r = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let i = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let c = r + i.asImaginary() + let result = MLX.fftn(c) + expectSummary( + result.realPart(), + ArraySummary( + shape: [8, 8, 8], + dtype: .float32, + mean: -1.9362115859985352, + minimum: -71.32183837890625, + maximum: 69.47334289550781, + absoluteSum: 9444.30859375, + positionChecksum: 4618.19921875, + sampleIndices: [0, 102, 204, 307, 409, 511], + samples: [ + -40.14142608642578, -46.13766860961914, -42.42741775512695, + -28.554052352905273, 51.425296783447266, -16.503108978271484, + ]), + tolerance: .float32) + expectSummary( + result.imaginaryPart(), + ArraySummary( + shape: [8, 8, 8], + dtype: .float32, + mean: -0.8104066252708435, + minimum: -74.49031066894531, + maximum: 64.66445922851562, + absoluteSum: 8921.556640625, + positionChecksum: 4586.01806640625, + sampleIndices: [0, 102, 204, 307, 409, 511], + samples: [ + 26.956165313720703, 5.077007293701172, -26.426931381225586, + -0.8038463592529297, -10.05626392364502, 9.05034351348877, + ]), + tolerance: .float32) + } + } + + @Test("fftn/s") + func test_fftn_s() throws { + try withIntegrationState(seed: 56309) { + let r = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let i = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let c = r + i.asImaginary() + let result = MLX.fftn(c, s: [3, 4], axes: [0, 1]) + expectSummary( + result.realPart(), + ArraySummary( + shape: [3, 4, 8], + dtype: .float32, + mean: -0.4631645083427429, + minimum: -9.234552383422852, + maximum: 7.494760513305664, + absoluteSum: 269.39691162109375, + positionChecksum: 135.80084228515625, + sampleIndices: [0, 19, 38, 57, 76, 95], + samples: [ + 2.372249126434326, -0.43843603134155273, 2.477461814880371, + 0.17259609699249268, 5.412355422973633, 1.3921151161193848, + ]), + tolerance: .float32) + expectSummary( + result.imaginaryPart(), + ArraySummary( + shape: [3, 4, 8], + dtype: .float32, + mean: 0.6277634501457214, + minimum: -6.913886070251465, + maximum: 11.678964614868164, + absoluteSum: 276.4273986816406, + positionChecksum: 139.78584798177084, + sampleIndices: [0, 19, 38, 57, 76, 95], + samples: [ + 1.672297716140747, 3.777714252471924, 3.101365566253662, + -1.5718486309051514, 1.980665922164917, -5.670261859893799, + ]), + tolerance: .float32) + } + } + + @Test("fftn/axes") + func test_fftn_axes() throws { + try withIntegrationState(seed: 33152) { + let r = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let i = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let c = r + i.asImaginary() + let result = MLX.fftn(c, axes: [0, 2]) + expectSummary( + result.realPart(), + ArraySummary( + shape: [8, 8, 8], + dtype: .float32, + mean: 0.069861501455307, + minimum: -26.00143814086914, + maximum: 24.77533531188965, + absoluteSum: 3271.591796875, + positionChecksum: 1632.01904296875, + sampleIndices: [0, 102, 204, 307, 409, 511], + samples: [ + 3.444851875305176, 15.690192222595215, -6.4840569496154785, + 2.230067729949951, 1.1066653728485107, -12.76463794708252, + ]), + tolerance: .float32) + expectSummary( + result.imaginaryPart(), + ArraySummary( + shape: [8, 8, 8], + dtype: .float32, + mean: -0.1090853214263916, + minimum: -22.648542404174805, + maximum: 21.17329978942871, + absoluteSum: 3264.510986328125, + positionChecksum: 1672.77587890625, + sampleIndices: [0, 102, 204, 307, 409, 511], + samples: [ + -0.9768218994140625, 5.668068885803223, -2.5782880783081055, + 5.189621925354004, 7.344736099243164, -5.029444694519043, + ]), + tolerance: .float32) + } + } + + @Test("ifftn") + func test_ifftn() throws { + try withIntegrationState(seed: 41746) { + let r = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let i = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let c = r + i.asImaginary() + let result = MLX.ifftn(c) + expectSummary( + result.realPart(), + ArraySummary( + shape: [8, 8, 8], + dtype: .float32, + mean: 0.0007427402306348085, + minimum: -0.15388083457946777, + maximum: 0.1925613135099411, + absoluteSum: 17.775423049926758, + positionChecksum: 9.001312255859375, + sampleIndices: [0, 102, 204, 307, 409, 511], + samples: [ + 0.049047283828258514, -0.009180162101984024, 0.03828829899430275, + 0.038404885679483414, 0.07229334115982056, 0.006673771888017654, + ]), + tolerance: .float32) + expectSummary( + result.imaginaryPart(), + ArraySummary( + shape: [8, 8, 8], + dtype: .float32, + mean: 0.0006254386971704662, + minimum: -0.18057334423065186, + maximum: 0.12072519958019257, + absoluteSum: 18.66337776184082, + positionChecksum: 9.401142120361328, + sampleIndices: [0, 102, 204, 307, 409, 511], + samples: [ + 0.050497107207775116, -0.03860977664589882, -0.06357460469007492, + -0.02733970433473587, 0.05312860757112503, -0.020095955580472946, + ]), + tolerance: .float32) + } + } + + @Test("ifftn/s") + func test_ifftn_s() throws { + try withIntegrationState(seed: 2929) { + let r = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let i = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let c = r + i.asImaginary() + let result = MLX.ifftn(c, s: [3, 4], axes: [0, 1]) + expectSummary( + result.realPart(), + ArraySummary( + shape: [3, 4, 8], + dtype: .float32, + mean: -0.03799237683415413, + minimum: -0.7618478536605835, + maximum: 0.8563539981842041, + absoluteSum: 22.170364379882812, + positionChecksum: 10.169722239176432, + sampleIndices: [0, 19, 38, 57, 76, 95], + samples: [ + 0.49846258759498596, -0.5838869214057922, -0.09089171886444092, + 0.23879596590995789, 0.07031029462814331, -0.007234610617160797, + ]), + tolerance: .float32) + expectSummary( + result.imaginaryPart(), + ArraySummary( + shape: [3, 4, 8], + dtype: .float32, + mean: 0.03390979766845703, + minimum: -0.6404374241828918, + maximum: 0.5621622800827026, + absoluteSum: 21.58087730407715, + positionChecksum: 11.274425506591797, + sampleIndices: [0, 19, 38, 57, 76, 95], + samples: [ + 0.42583590745925903, 0.26459410786628723, -0.007343888282775879, + 0.465649276971817, -0.23035968840122223, 0.15064139664173126, + ]), + tolerance: .float32) + } + } + + @Test("ifftn/axes") + func test_ifftn_axes() throws { + try withIntegrationState(seed: 23441) { + let r = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let i = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let c = r + i.asImaginary() + let result = MLX.ifftn(c, axes: [0, 2]) + expectSummary( + result.realPart(), + ArraySummary( + shape: [8, 8, 8], + dtype: .float32, + mean: -0.008047770708799362, + minimum: -0.4037286043167114, + maximum: 0.4154580235481262, + absoluteSum: 49.243099212646484, + positionChecksum: 24.31148910522461, + sampleIndices: [0, 102, 204, 307, 409, 511], + samples: [ + 0.015621490776538849, -0.12264824658632278, -0.04659213125705719, + 0.0034754127264022827, -0.09425322711467743, 0.11336925625801086, + ]), + tolerance: .float32) + expectSummary( + result.imaginaryPart(), + ArraySummary( + shape: [8, 8, 8], + dtype: .float32, + mean: -0.002788180485367775, + minimum: -0.3810046315193176, + maximum: 0.4174731373786926, + absoluteSum: 50.17457962036133, + positionChecksum: 25.259380340576172, + sampleIndices: [0, 102, 204, 307, 409, 511], + samples: [ + 0.035300590097904205, 0.08818861842155457, -0.04737217724323273, + -0.02006063610315323, 0.08536103367805481, -0.05704180896282196, + ]), + tolerance: .float32) + } + } + + @Test("rfft2") + func test_rfft2() throws { + try withIntegrationState(seed: 10302) { + let c = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.rfft2(c) + expectSummary( + result.realPart(), + ArraySummary( + shape: [8, 8, 5], + dtype: .float32, + mean: 0.06260216236114502, + minimum: -15.983535766601562, + maximum: 15.228923797607422, + absoluteSum: 1476.531494140625, + positionChecksum: 682.576953125, + sampleIndices: [0, 64, 128, 191, 255, 319], + samples: [ + -15.983535766601562, 2.8445022106170654, -1.0613014698028564, + -3.684018135070801, 2.3670763969421387, 0.32685017585754395, + ]), + tolerance: .float32) + expectSummary( + result.imaginaryPart(), + ArraySummary( + shape: [8, 8, 5], + dtype: .float32, + mean: 0.055629778653383255, + minimum: -17.28950309753418, + maximum: 16.831018447875977, + absoluteSum: 1299.924560546875, + positionChecksum: 631.76669921875, + sampleIndices: [0, 64, 128, 191, 255, 319], + samples: [ + 0.0, 0.0, -0.3817894458770752, -2.158597230911255, -4.935375690460205, + -2.6903560161590576, + ]), + tolerance: .float32) + } + } + + @Test("rfft2/axes") + func test_rfft2_axes() throws { + try withIntegrationState(seed: 47042) { + let c = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.rfft2(c, axes: [0, 2]) + expectSummary( + result.realPart(), + ArraySummary( + shape: [8, 8, 5], + dtype: .float32, + mean: 0.7185362577438354, + minimum: -14.52419662475586, + maximum: 16.742055892944336, + absoluteSum: 1495.4569091796875, + positionChecksum: 726.011865234375, + sampleIndices: [0, 64, 128, 191, 255, 319], + samples: [ + 5.749485015869141, -7.658758640289307, -0.8341364860534668, + 2.9679250717163086, 6.413942813873291, 7.591013431549072, + ]), + tolerance: .float32) + expectSummary( + result.imaginaryPart(), + ArraySummary( + shape: [8, 8, 5], + dtype: .float32, + mean: 0.23355670273303986, + minimum: -15.147994995117188, + maximum: 14.86296272277832, + absoluteSum: 1337.6533203125, + positionChecksum: 668.728515625, + sampleIndices: [0, 64, 128, 191, 255, 319], + samples: [ + 0.0, -5.069750785827637, -0.016709327697753906, 2.846343517303467, + 13.277387619018555, 4.648657321929932, + ]), + tolerance: .float32) + } + } + + @Test("rfftn") + func test_rfftn() throws { + try withIntegrationState(seed: 87041) { + let c = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.rfftn(c) + expectSummary( + result.realPart(), + ArraySummary( + shape: [8, 8, 5], + dtype: .float32, + mean: -1.1790698766708374, + minimum: -43.51289367675781, + maximum: 40.62009048461914, + absoluteSum: 4045.7822265625, + positionChecksum: 2089.2171875, + sampleIndices: [0, 64, 128, 191, 255, 319], + samples: [ + 2.826408863067627, 18.93130874633789, -16.50474739074707, 5.28663969039917, + -1.7076082229614258, -12.364112854003906, + ]), + tolerance: .float32) + expectSummary( + result.imaginaryPart(), + ArraySummary( + shape: [8, 8, 5], + dtype: .float32, + mean: -0.06564798206090927, + minimum: -53.53523254394531, + maximum: 53.53523254394531, + absoluteSum: 4218.8720703125, + positionChecksum: 2227.225, + sampleIndices: [0, 64, 128, 191, 255, 319], + samples: [ + 0.0, -14.105966567993164, -23.85370635986328, -16.85041046142578, + -1.7126522064208984, 8.338759422302246, + ]), + tolerance: .float32) + } + } + + @Test("rfftn/axes") + func test_rfftn_axes() throws { + try withIntegrationState(seed: 46670) { + let c = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.rfftn(c, axes: [0, 2]) + expectSummary( + result.realPart(), + ArraySummary( + shape: [8, 8, 5], + dtype: .float32, + mean: 0.07727640122175217, + minimum: -16.674985885620117, + maximum: 17.34377098083496, + absoluteSum: 1540.9918212890625, + positionChecksum: 760.8494140625, + sampleIndices: [0, 64, 128, 191, 255, 319], + samples: [ + -15.779212951660156, -1.9917529821395874, -7.319093704223633, + 1.1237940788269043, 17.34377098083496, 6.796459197998047, + ]), + tolerance: .float32) + expectSummary( + result.imaginaryPart(), + ArraySummary( + shape: [8, 8, 5], + dtype: .float32, + mean: 0.20542661845684052, + minimum: -16.39087677001953, + maximum: 16.39087677001953, + absoluteSum: 1281.357666015625, + positionChecksum: 672.893212890625, + sampleIndices: [0, 64, 128, 191, 255, 319], + samples: [ + 0.0, -0.4302929639816284, 3.292222499847412, -1.093297004699707, + 16.39087677001953, 9.523144721984863, + ]), + tolerance: .float32) + } + } + + @Test("irfft2") + func test_irfft2() throws { + try withIntegrationState(seed: 34359) { + let r = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let i = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let c = r + i.asImaginary() + let result = MLX.irfft2(c) + expectSummary( + result, + ArraySummary( + shape: [8, 8, 14], + dtype: .float32, + mean: -0.004153141751885414, + minimum: -0.4831690192222595, + maximum: 0.39062052965164185, + absoluteSum: 90.77721405029297, + positionChecksum: 45.66217476981027, + sampleIndices: [0, 179, 358, 537, 716, 895], + samples: [ + -0.04847263544797897, 0.20513692498207092, -0.229575976729393, + -0.044226326048374176, -0.07635796815156937, -0.04189996048808098, + ]), + tolerance: .float32) + } + } + + @Test("irfft2/axes") + func test_irfft2_axes() throws { + try withIntegrationState(seed: 42730) { + let r = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let i = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let c = r + i.asImaginary() + let result = MLX.irfft2(c, axes: [0, 2]) + expectSummary( + result, + ArraySummary( + shape: [8, 8, 14], + dtype: .float32, + mean: -0.0023053386248648167, + minimum: -0.3896021842956543, + maximum: 0.4151984751224518, + absoluteSum: 91.33279418945312, + positionChecksum: 45.873779296875, + sampleIndices: [0, 179, 358, 537, 716, 895], + samples: [ + -0.11438766866922379, 0.0355401448905468, 0.13291791081428528, + -0.1825760304927826, 0.04413542523980141, -0.03667726740241051, + ]), + tolerance: .float32) + } + } + + @Test("irfftn") + func test_irfftn() throws { + try withIntegrationState(seed: 58696) { + let r = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let i = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let c = r + i.asImaginary() + let result = MLX.irfftn(c) + expectSummary( + result, + ArraySummary( + shape: [8, 8, 14], + dtype: .float32, + mean: 0.00038374235737137496, + minimum: -0.13770824670791626, + maximum: 0.16200695931911469, + absoluteSum: 32.69968795776367, + positionChecksum: 16.683267865862167, + sampleIndices: [0, 179, 358, 537, 716, 895], + samples: [ + 0.05379920452833176, 0.045248210430145264, 0.008416865952312946, + 0.0006176760070957243, 0.025639304891228676, -0.06305667012929916, + ]), + tolerance: .float32) + } + } + + @Test("irfftn/axes") + func test_irfftn_axes() throws { + try withIntegrationState(seed: 9062) { + let r = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let i = MLXRandom.normal([8, 8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let c = r + i.asImaginary() + let result = MLX.irfftn(c, axes: [0, 2]) + expectSummary( + result, + ArraySummary( + shape: [8, 8, 14], + dtype: .float32, + mean: 0.0023820779751986265, + minimum: -0.34605321288108826, + maximum: 0.40482693910598755, + absoluteSum: 86.85626220703125, + positionChecksum: 44.459620884486604, + sampleIndices: [0, 179, 358, 537, 716, 895], + samples: [ + -0.09112237393856049, 0.04324769228696823, 0.007158036343753338, + 0.15657760202884674, 0.08600138127803802, -0.22598956525325775, + ]), + tolerance: .float32) + } + } + + @Test("fftfreq") + func test_fftfreq() throws { + try withIntegrationState(seed: 29017) { + let result = MLX.fftfreq(16) + expectSummary( + result, + ArraySummary( + shape: [16], + dtype: .float32, + mean: -0.03125, + minimum: -0.5, + maximum: 0.4375, + absoluteSum: 4.0, + positionChecksum: 2.25, + sampleIndices: [0, 3, 6, 9, 12, 15], + samples: [0.0, 0.1875, 0.375, -0.4375, -0.25, -0.0625]), + tolerance: .float32) + } + } + + @Test("fftfreq/d") + func test_fftfreq_d() throws { + try withIntegrationState(seed: 91367) { + let result = MLX.fftfreq(16, d: 0.25) + expectSummary( + result, + ArraySummary( + shape: [16], + dtype: .float32, + mean: -0.125, + minimum: -2.0, + maximum: 1.75, + absoluteSum: 16.0, + positionChecksum: 9.0, + sampleIndices: [0, 3, 6, 9, 12, 15], + samples: [0.0, 0.75, 1.5, -1.75, -1.0, -0.25]), + tolerance: .float32) + } + } + + @Test("rfftfreq") + func test_rfftfreq() throws { + try withIntegrationState(seed: 94336) { + let result = MLX.rfftfreq(16) + expectSummary( + result, + ArraySummary( + shape: [9], + dtype: .float32, + mean: 0.25, + minimum: 0.0, + maximum: 0.5, + absoluteSum: 2.25, + positionChecksum: 1.6666666666666667, + sampleIndices: [0, 2, 3, 5, 6, 8], + samples: [0.0, 0.125, 0.1875, 0.3125, 0.375, 0.5]), + tolerance: .float32) + } + } + + @Test("fftshift") + func test_fftshift() throws { + try withIntegrationState(seed: 60895) { + let a = MLXRandom.normal([8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.fftshift(a) + expectSummary( + result, + ArraySummary( + shape: [8, 8], + dtype: .float32, + mean: 0.16273652017116547, + minimum: -2.3805737495422363, + maximum: 2.3832385540008545, + absoluteSum: 52.84351348876953, + positionChecksum: 24.834022521972656, + sampleIndices: [0, 13, 25, 38, 50, 63], + samples: [ + 1.3552110195159912, 0.08066458255052567, 1.1147830486297607, + 0.0029193360824137926, 0.69813072681427, -0.2748071253299713, + ]), + tolerance: .float32) + } + } + + @Test("fftshift/axes") + func test_fftshift_axes() throws { + try withIntegrationState(seed: 3962) { + let a = MLXRandom.normal([8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.fftshift(a, axes: [0]) + expectSummary( + result, + ArraySummary( + shape: [8, 8], + dtype: .float32, + mean: -0.012623829767107964, + minimum: -1.9678735733032227, + maximum: 1.761767864227295, + absoluteSum: 51.23778533935547, + positionChecksum: 25.559467315673828, + sampleIndices: [0, 13, 25, 38, 50, 63], + samples: [ + -0.9273226857185364, 1.414719820022583, 0.45456352829933167, + 1.4807580709457397, -0.10517635941505432, 0.7928916811943054, + ]), + tolerance: .float32) + } + } + + @Test("ifftshift") + func test_ifftshift() throws { + try withIntegrationState(seed: 43884) { + let a = MLXRandom.normal([8, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.ifftshift(a) + expectSummary( + result, + ArraySummary( + shape: [8, 8], + dtype: .float32, + mean: 0.12466083467006683, + minimum: -1.6268000602722168, + maximum: 2.804321765899658, + absoluteSum: 48.001792907714844, + positionChecksum: 23.223478317260742, + sampleIndices: [0, 13, 25, 38, 50, 63], + samples: [ + -0.6544477939605713, 0.8364695906639099, 1.0058847665786743, + 2.804321765899658, -0.050703976303339005, -0.33208170533180237, + ]), + tolerance: .float32) + } + } +} diff --git a/Tests/MLXIntegrationTests/Generated/GeneratedFactoryTests.swift b/Tests/MLXIntegrationTests/Generated/GeneratedFactoryTests.swift new file mode 100644 index 000000000..27bb07d78 --- /dev/null +++ b/Tests/MLXIntegrationTests/Generated/GeneratedFactoryTests.swift @@ -0,0 +1,500 @@ +// Copyright © 2026 Apple Inc. +// +// GENERATED by tools/integration_tests -- DO NOT EDIT. +// +// Values in this file were produced by python mlx and are compared against +// the Swift API. Regenerate with: +// +// python3 tools/integration_tests/generate.py +// +// python mlx: 0.32.2.dev20260910+1f8e74e3f +// vendored mlx (Cmlx): 0.32.2 +// generator revision: 2 +// cases: 23 + +import Foundation +import MLX +import Testing + +@Suite("generated: Factory") +struct GeneratedFactoryTests { + + @Test("arange/stop") + func test_arange_stop() throws { + try withIntegrationState(seed: 76748) { + let result = MLX.arange(10) + expectSummary( + result, + ArraySummary( + shape: [10], + dtype: .int32, + mean: 4.5, + minimum: 0.0, + maximum: 9.0, + absoluteSum: 45.0, + positionChecksum: 33.0, + sampleIndices: [0, 2, 4, 5, 7, 9], + samples: [0.0, 2.0, 4.0, 5.0, 7.0, 9.0]), + tolerance: .exact) + } + } + + @Test("arange/range") + func test_arange_range() throws { + try withIntegrationState(seed: 98523) { + let result = MLX.arange(2, 12) + expectSummary( + result, + ArraySummary( + shape: [10], + dtype: .int32, + mean: 6.5, + minimum: 2.0, + maximum: 11.0, + absoluteSum: 65.0, + positionChecksum: 44.0, + sampleIndices: [0, 2, 4, 5, 7, 9], + samples: [2.0, 4.0, 6.0, 7.0, 9.0, 11.0]), + tolerance: .exact) + } + } + + @Test("arange/step") + func test_arange_step() throws { + try withIntegrationState(seed: 7526) { + let result = MLX.arange(0.0, 2.0, step: 0.25) + expectSummary( + result, + ArraySummary( + shape: [8], + dtype: .float32, + mean: 0.875, + minimum: 0.0, + maximum: 1.75, + absoluteSum: 7.0, + positionChecksum: 5.25, + sampleIndices: [0, 1, 3, 4, 6, 7], + samples: [0.0, 0.25, 0.75, 1.0, 1.5, 1.75]), + tolerance: .float32) + } + } + + @Test("arange/dtype") + func test_arange_dtype() throws { + try withIntegrationState(seed: 62487) { + let result = MLX.arange(0, 10, step: 2, dtype: .int32) + expectSummary( + result, + ArraySummary( + shape: [5], + dtype: .int32, + mean: 4.0, + minimum: 0.0, + maximum: 8.0, + absoluteSum: 20.0, + positionChecksum: 16.0, + sampleIndices: [0, 1, 2, 3, 4], + samples: [0.0, 2.0, 4.0, 6.0, 8.0]), + tolerance: .exact) + } + } + + @Test("linspace") + func test_linspace() throws { + try withIntegrationState(seed: 87908) { + let result = MLX.linspace(0.0, 1.0, count: 9) + expectSummary( + result, + ArraySummary( + shape: [9], + dtype: .float32, + mean: 0.5, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 4.5, + positionChecksum: 3.3333333333333335, + sampleIndices: [0, 2, 3, 5, 6, 8], + samples: [0.0, 0.25, 0.375, 0.625, 0.75, 1.0]), + tolerance: .float32) + } + } + + @Test("linspace/int") + func test_linspace_int() throws { + // integer bounds still produce float32, matching python + try withIntegrationState(seed: 96709) { + let result = MLX.linspace(0, 10, count: 6) + expectSummary( + result, + ArraySummary( + shape: [6], + dtype: .float32, + mean: 5.0, + minimum: 0.0, + maximum: 10.0, + absoluteSum: 30.0, + positionChecksum: 23.333333333333332, + sampleIndices: [0, 1, 2, 3, 4, 5], + samples: [0.0, 2.0, 4.0, 6.0, 8.0, 10.0]), + tolerance: .float32) + } + } + + @Test("linspace/dtype") + func test_linspace_dtype() throws { + try withIntegrationState(seed: 61299) { + let result = MLX.linspace(0, 10, count: 6, dtype: .int32) + expectSummary( + result, + ArraySummary( + shape: [6], + dtype: .int32, + mean: 5.0, + minimum: 0.0, + maximum: 10.0, + absoluteSum: 30.0, + positionChecksum: 23.333333333333332, + sampleIndices: [0, 1, 2, 3, 4, 5], + samples: [0.0, 2.0, 4.0, 6.0, 8.0, 10.0]), + tolerance: .exact) + } + } + + @Test("linspace/endpoint") + func test_linspace_endpoint() throws { + // python has no endpoint parameter: the half-open interval is the inclusive one with the last sample dropped + try withIntegrationState(seed: 29128) { + let result = MLX.linspace(0.0, 1.0, count: 5, endpoint: false) + expectSummary( + result, + ArraySummary( + shape: [5], + dtype: .float32, + mean: 0.4000000059604645, + minimum: 0.0, + maximum: 0.800000011920929, + absoluteSum: 2.0, + positionChecksum: 1.6, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + 0.0, 0.20000000298023224, 0.4000000059604645, 0.6000000238418579, + 0.800000011920929, + ]), + tolerance: .float32) + } + } + + @Test("zeros") + func test_zeros() throws { + try withIntegrationState(seed: 87556) { + let result = MLX.zeros([3, 4]) + expectSummary( + result, + ArraySummary( + shape: [3, 4], + dtype: .float32, + mean: 0.0, + minimum: 0.0, + maximum: 0.0, + absoluteSum: 0.0, + positionChecksum: 0.0, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [0.0, 0.0, 0.0, 0.0, 0.0, 0.0]), + tolerance: .float32) + } + } + + @Test("zeros/dtype") + func test_zeros_dtype() throws { + try withIntegrationState(seed: 92392) { + let result = MLX.zeros([3, 4], dtype: .int32) + expectSummary( + result, + ArraySummary( + shape: [3, 4], + dtype: .int32, + mean: 0.0, + minimum: 0.0, + maximum: 0.0, + absoluteSum: 0.0, + positionChecksum: 0.0, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [0.0, 0.0, 0.0, 0.0, 0.0, 0.0]), + tolerance: .exact) + } + } + + @Test("ones") + func test_ones() throws { + try withIntegrationState(seed: 95935) { + let result = MLX.ones([3, 4]) + expectSummary( + result, + ArraySummary( + shape: [3, 4], + dtype: .float32, + mean: 1.0, + minimum: 1.0, + maximum: 1.0, + absoluteSum: 12.0, + positionChecksum: 6.5, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [1.0, 1.0, 1.0, 1.0, 1.0, 1.0]), + tolerance: .float32) + } + } + + @Test("ones/dtype") + func test_ones_dtype() throws { + try withIntegrationState(seed: 51426) { + let result = MLX.ones([3, 4], dtype: .int16) + expectSummary( + result, + ArraySummary( + shape: [3, 4], + dtype: .int16, + mean: 1.0, + minimum: 1.0, + maximum: 1.0, + absoluteSum: 12.0, + positionChecksum: 6.5, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [1.0, 1.0, 1.0, 1.0, 1.0, 1.0]), + tolerance: .exact) + } + } + + @Test("zeros/like") + func test_zeros_like() throws { + try withIntegrationState(seed: 79637) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.zeros(like: a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.0, + minimum: 0.0, + maximum: 0.0, + absoluteSum: 0.0, + positionChecksum: 0.0, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [0.0, 0.0, 0.0, 0.0, 0.0, 0.0]), + tolerance: .float32) + } + } + + @Test("ones/like") + func test_ones_like() throws { + try withIntegrationState(seed: 45715) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.ones(like: a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 1.0, + minimum: 1.0, + maximum: 1.0, + absoluteSum: 12.0, + positionChecksum: 6.5, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [1.0, 1.0, 1.0, 1.0, 1.0, 1.0]), + tolerance: .float32) + } + } + + @Test("full") + func test_full() throws { + try withIntegrationState(seed: 96888) { + let result = MLX.full([2, 3], values: Float(2.5)) + expectSummary( + result, + ArraySummary( + shape: [2, 3], + dtype: .float32, + mean: 2.5, + minimum: 2.5, + maximum: 2.5, + absoluteSum: 15.0, + positionChecksum: 8.75, + sampleIndices: [0, 1, 2, 3, 4, 5], + samples: [2.5, 2.5, 2.5, 2.5, 2.5, 2.5]), + tolerance: .float32) + } + } + + @Test("full/dtype") + func test_full_dtype() throws { + try withIntegrationState(seed: 47656) { + let result = MLX.full([2, 3], values: MLXArray(7), dtype: .int32) + expectSummary( + result, + ArraySummary( + shape: [2, 3], + dtype: .int32, + mean: 7.0, + minimum: 7.0, + maximum: 7.0, + absoluteSum: 42.0, + positionChecksum: 24.5, + sampleIndices: [0, 1, 2, 3, 4, 5], + samples: [7.0, 7.0, 7.0, 7.0, 7.0, 7.0]), + tolerance: .exact) + } + } + + @Test("eye") + func test_eye() throws { + try withIntegrationState(seed: 78454) { + let result = MLX.eye(4) + expectSummary( + result, + ArraySummary( + shape: [4, 4], + dtype: .float32, + mean: 0.25, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 4.0, + positionChecksum: 2.125, + sampleIndices: [0, 3, 6, 9, 12, 15], + samples: [1.0, 0.0, 0.0, 0.0, 0.0, 1.0]), + tolerance: .float32) + } + } + + @Test("eye/rectangular") + func test_eye_rectangular() throws { + try withIntegrationState(seed: 51823) { + let result = MLX.eye(4, m: 6, k: 1) + expectSummary( + result, + ArraySummary( + shape: [4, 6], + dtype: .float32, + mean: 0.1666666716337204, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 4.0, + positionChecksum: 2.0833333333333335, + sampleIndices: [0, 5, 9, 14, 18, 23], + samples: [0.0, 0.0, 0.0, 0.0, 0.0, 0.0]), + tolerance: .float32) + } + } + + @Test("identity") + func test_identity() throws { + try withIntegrationState(seed: 81094) { + let result = MLX.identity(5) + expectSummary( + result, + ArraySummary( + shape: [5, 5], + dtype: .float32, + mean: 0.19999998807907104, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 5.0, + positionChecksum: 2.6, + sampleIndices: [0, 5, 10, 14, 19, 24], + samples: [1.0, 0.0, 0.0, 0.0, 0.0, 1.0]), + tolerance: .float32) + } + } + + @Test("bartlett") + func test_bartlett() throws { + try withIntegrationState(seed: 66814) { + let result = MLX.bartlett(16) + expectSummary( + result, + ArraySummary( + shape: [16], + dtype: .float32, + mean: 0.46666663885116577, + minimum: 0.0, + maximum: 0.9333333969116211, + absoluteSum: 7.466666221618652, + positionChecksum: 3.9666662216186523, + sampleIndices: [0, 3, 6, 9, 12, 15], + samples: [ + 0.0, 0.40000003576278687, 0.8000000715255737, 0.7999999523162842, + 0.39999985694885254, 0.0, + ]), + tolerance: .float32) + } + } + + @Test("hanning") + func test_hanning() throws { + try withIntegrationState(seed: 34335) { + let result = MLX.hanning(16) + expectSummary( + result, + ArraySummary( + shape: [16], + dtype: .float32, + mean: 0.4687499701976776, + minimum: 0.0, + maximum: 0.9890736937522888, + absoluteSum: 7.499999523162842, + positionChecksum: 3.984375, + sampleIndices: [0, 3, 6, 9, 12, 15], + samples: [ + 0.0, 0.34549155831336975, 0.904508650302887, 0.9045085310935974, + 0.3454914093017578, 7.64274186065882e-15, + ]), + tolerance: .float32) + } + } + + @Test("hamming") + func test_hamming() throws { + try withIntegrationState(seed: 86817) { + let result = MLX.hamming(16) + expectSummary( + result, + ArraySummary( + shape: [16], + dtype: .float32, + mean: 0.5112500190734863, + minimum: 0.08000001311302185, + maximum: 0.9899479746818542, + absoluteSum: 8.180000305175781, + positionChecksum: 4.345624923706055, + sampleIndices: [0, 3, 6, 9, 12, 15], + samples: [ + 0.08000001311302185, 0.39785224199295044, 0.9121478796005249, + 0.9121477603912354, 0.3978521227836609, 0.08000001311302185, + ]), + tolerance: .float32) + } + } + + @Test("blackman") + func test_blackman() throws { + try withIntegrationState(seed: 98518) { + let result = MLX.blackman(16) + expectSummary( + result, + ArraySummary( + shape: [16], + dtype: .float32, + mean: 0.39375001192092896, + minimum: 0.0, + maximum: 0.9821575284004211, + absoluteSum: 6.300000190734863, + positionChecksum: 3.346874713897705, + sampleIndices: [0, 3, 6, 9, 12, 15], + samples: [ + 0.0, 0.20077016949653625, 0.8492298722267151, 0.8492297530174255, + 0.20077009499073029, 0.0, + ]), + tolerance: .float32) + } + } +} diff --git a/Tests/MLXIntegrationTests/Generated/GeneratedFastTests.swift b/Tests/MLXIntegrationTests/Generated/GeneratedFastTests.swift new file mode 100644 index 000000000..6a0102a4a --- /dev/null +++ b/Tests/MLXIntegrationTests/Generated/GeneratedFastTests.swift @@ -0,0 +1,201 @@ +// Copyright © 2026 Apple Inc. +// +// GENERATED by tools/integration_tests -- DO NOT EDIT. +// +// Values in this file were produced by python mlx and are compared against +// the Swift API. Regenerate with: +// +// python3 tools/integration_tests/generate.py +// +// python mlx: 0.32.2.dev20260910+1f8e74e3f +// vendored mlx (Cmlx): 0.32.2 +// generator revision: 2 +// cases: 7 + +import Foundation +import MLX +import Testing + +@Suite("generated: Fast") +struct GeneratedFastTests { + + @Test("rmsNorm") + func test_rmsNorm() throws { + try withIntegrationState(seed: 66852) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let weight = MLXRandom.normal([16], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.rmsNorm(x, weight: weight, eps: 1e-5) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: -0.06485586613416672, + minimum: -3.9406793117523193, + maximum: 3.7056970596313477, + absoluteSum: 191.23782348632812, + positionChecksum: 95.85537719726562, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + -1.0069217681884766, -0.8062147498130798, 1.2844102382659912, + 0.10886163264513016, 1.0201539993286133, -0.6644749045372009, + ]), + tolerance: .float32) + } + } + + @Test("layerNorm") + func test_layerNorm() throws { + try withIntegrationState(seed: 45114) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let weight = MLXRandom.normal([16], dtype: .float32, loc: 0.0, scale: 1.0) + let bias = MLXRandom.normal([16], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.layerNorm(x, weight: weight, bias: bias, eps: 1e-5) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: 0.23579367995262146, + minimum: -3.412724733352661, + maximum: 3.5934901237487793, + absoluteSum: 225.9371337890625, + positionChecksum: 108.25701904296875, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + -0.18614619970321655, -0.702523946762085, -0.5655195713043213, + 1.331803560256958, -1.6021459102630615, -0.23616120219230652, + ]), + tolerance: .float32) + } + } + + @Test("layerNorm/noAffine") + func test_layerNorm_noAffine() throws { + try withIntegrationState(seed: 56298) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.layerNorm(x, eps: 1e-5) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: -6.51925802230835e-09, + minimum: -2.750777006149292, + maximum: 2.616267681121826, + absoluteSum: 200.528564453125, + positionChecksum: 99.69488525390625, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + -0.5727211833000183, 0.3872288167476654, 2.390326976776123, + 1.077345848083496, -1.7498226165771484, -0.0715562179684639, + ]), + tolerance: .float32) + } + } + + @Test("RoPE") + func test_RoPE() throws { + try withIntegrationState(seed: 66733) { + let x = MLXRandom.normal([2, 4, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.RoPE( + x, dimensions: 16, traditional: false, base: 10000.0, scale: 1.0, offset: 0) + expectSummary( + result, + ArraySummary( + shape: [2, 4, 8, 16], + dtype: .float32, + mean: -0.00884411484003067, + minimum: -3.2420690059661865, + maximum: 3.197585344314575, + absoluteSum: 821.5716552734375, + positionChecksum: 414.88031005859375, + sampleIndices: [0, 205, 409, 614, 818, 1023], + samples: [ + -0.6869831681251526, 0.25631067156791687, 0.8028337955474854, + -0.9851393699645996, -0.39124467968940735, 2.371519088745117, + ]), + tolerance: .float32) + } + } + + @Test("RoPE/traditional") + func test_RoPE_traditional() throws { + try withIntegrationState(seed: 74377) { + let x = MLXRandom.normal([2, 4, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.RoPE( + x, dimensions: 8, traditional: true, base: 10000.0, scale: 0.5, offset: 2) + expectSummary( + result, + ArraySummary( + shape: [2, 4, 8, 16], + dtype: .float32, + mean: -0.028798004612326622, + minimum: -2.746978521347046, + maximum: 3.72855281829834, + absoluteSum: 807.062744140625, + positionChecksum: 392.5193786621094, + sampleIndices: [0, 205, 409, 614, 818, 1023], + samples: [ + -1.4806770086288452, -0.9948027729988098, 1.3070459365844727, + -0.7166029214859009, 0.09865694493055344, -2.09395694732666, + ]), + tolerance: .float32) + } + } + + @Test("scaledDotProductAttention") + func test_scaledDotProductAttention() throws { + try withIntegrationState(seed: 82359) { + let queries = MLXRandom.normal([1, 2, 4, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let keys = MLXRandom.normal([1, 2, 4, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let values = MLXRandom.normal([1, 2, 4, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.scaledDotProductAttention( + queries: queries, keys: keys, values: values, scale: 0.35, mask: nil) + expectSummary( + result, + ArraySummary( + shape: [1, 2, 4, 8], + dtype: .float32, + mean: 0.197175532579422, + minimum: -1.386329174041748, + maximum: 2.5051798820495605, + absoluteSum: 37.204132080078125, + positionChecksum: 20.866201400756836, + sampleIndices: [0, 13, 25, 38, 50, 63], + samples: [ + 0.5489547848701477, 0.5017918944358826, 0.6151727437973022, + -0.010526720434427261, 2.5051798820495605, 0.639729380607605, + ]), + tolerance: .float32) + } + } + + @Test("scaledDotProductAttention/mask") + func test_scaledDotProductAttention_mask() throws { + try withIntegrationState(seed: 63535) { + let queries = MLXRandom.normal([1, 2, 4, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let keys = MLXRandom.normal([1, 2, 4, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let values = MLXRandom.normal([1, 2, 4, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let mask = MLX.tril(MLX.ones([4, 4])) .== 1 + let result = MLX.scaledDotProductAttention( + queries: queries, keys: keys, values: values, scale: 0.35, mask: mask) + expectSummary( + result, + ArraySummary( + shape: [1, 2, 4, 8], + dtype: .float32, + mean: 0.0719122439622879, + minimum: -1.5812186002731323, + maximum: 1.1895414590835571, + absoluteSum: 33.8242301940918, + positionChecksum: 17.089391708374023, + sampleIndices: [0, 13, 25, 38, 50, 63], + samples: [ + -0.4037952125072479, 0.0268552303314209, -0.7593691945075989, + -0.8534163236618042, -0.503488302230835, 0.6569595336914062, + ]), + tolerance: .float32) + } + } +} diff --git a/Tests/MLXIntegrationTests/Generated/GeneratedIndexingTests.swift b/Tests/MLXIntegrationTests/Generated/GeneratedIndexingTests.swift new file mode 100644 index 000000000..e5c2b106c --- /dev/null +++ b/Tests/MLXIntegrationTests/Generated/GeneratedIndexingTests.swift @@ -0,0 +1,166 @@ +// Copyright © 2026 Apple Inc. +// +// GENERATED by tools/integration_tests -- DO NOT EDIT. +// +// Values in this file were produced by python mlx and are compared against +// the Swift API. Regenerate with: +// +// python3 tools/integration_tests/generate.py +// +// python mlx: 0.32.2.dev20260910+1f8e74e3f +// vendored mlx (Cmlx): 0.32.2 +// generator revision: 2 +// cases: 6 + +import Foundation +import MLX +import Testing + +@Suite("generated: Indexing") +struct GeneratedIndexingTests { + + @Test("index/range") + func test_index_range() throws { + try withIntegrationState(seed: 56799) { + let a = MLXRandom.normal([6, 5], dtype: .float32, loc: 0.0, scale: 1.0) + let result = a[1 ..< 4, 0 ..< 3] + expectSummary( + result, + ArraySummary( + shape: [3, 3], + dtype: .float32, + mean: 0.20556111633777618, + minimum: -1.9838948249816895, + maximum: 2.7164838314056396, + absoluteSum: 10.062643051147461, + positionChecksum: 5.407042609320746, + sampleIndices: [0, 2, 3, 5, 6, 8], + samples: [ + -0.7313419580459595, 1.4610944986343384, -1.9838948249816895, + 2.7164838314056396, -1.3910598754882812, 0.4201461672782898, + ]), + tolerance: .float32) + } + } + + @Test("index/lastColumn") + func test_index_lastColumn() throws { + try withIntegrationState(seed: 67368) { + let a = MLXRandom.normal([6, 5], dtype: .float32, loc: 0.0, scale: 1.0) + let result = a[0..., -1] + expectSummary( + result, + ArraySummary( + shape: [6], + dtype: .float32, + mean: -0.07870198786258698, + minimum: -1.1848795413970947, + maximum: 1.5605968236923218, + absoluteSum: 3.762835741043091, + positionChecksum: 1.838436762491862, + sampleIndices: [0, 1, 2, 3, 4, 5], + samples: [ + -0.4085389971733093, -1.1848795413970947, -0.5241052508354187, + 1.5605968236923218, 0.07067099958658218, 0.014044167473912239, + ]), + tolerance: .float32) + } + } + + @Test("index/newAxis") + func test_index_newAxis() throws { + try withIntegrationState(seed: 7679) { + let a = MLXRandom.normal([6, 5], dtype: .float32, loc: 0.0, scale: 1.0) + let result = a[.ellipsis, .newAxis] + expectSummary( + result, + ArraySummary( + shape: [6, 5, 1], + dtype: .float32, + mean: -0.09582757204771042, + minimum: -1.8003718852996826, + maximum: 1.4870085716247559, + absoluteSum: 22.161052703857422, + positionChecksum: 11.798158772786458, + sampleIndices: [0, 6, 12, 17, 23, 29], + samples: [ + 0.08150071650743484, 0.26980727910995483, 0.2373223602771759, + -0.4347092807292938, 1.3970370292663574, -0.6056225299835205, + ]), + tolerance: .float32) + } + } + + @Test("index/stride") + func test_index_stride() throws { + try withIntegrationState(seed: 33648) { + let a = MLXRandom.normal([6, 5], dtype: .float32, loc: 0.0, scale: 1.0) + let result = a[.stride(by: 2)] + expectSummary( + result, + ArraySummary( + shape: [3, 5], + dtype: .float32, + mean: 0.07586007565259933, + minimum: -1.3099104166030884, + maximum: 2.5639853477478027, + absoluteSum: 12.466233253479004, + positionChecksum: 6.5820460001627605, + sampleIndices: [0, 3, 6, 8, 11, 14], + samples: [ + -0.0001792133116396144, -0.5972116589546204, -0.8127403259277344, + 2.5639853477478027, 0.8785732388496399, -0.3988332748413086, + ]), + tolerance: .float32) + } + } + + @Test("index/reverse") + func test_index_reverse() throws { + try withIntegrationState(seed: 42592) { + let a = MLXRandom.normal([6, 5], dtype: .float32, loc: 0.0, scale: 1.0) + let result = a[.ellipsis, .stride(by: -1)] + expectSummary( + result, + ArraySummary( + shape: [6, 5], + dtype: .float32, + mean: -0.13014544546604156, + minimum: -2.513247489929199, + maximum: 1.7819768190383911, + absoluteSum: 19.7058162689209, + positionChecksum: 8.921414184570313, + sampleIndices: [0, 6, 12, 17, 23, 29], + samples: [ + 1.695084571838379, -0.2929709255695343, 0.3612818419933319, + 0.3690113425254822, -0.16253980994224548, 0.26952359080314636, + ]), + tolerance: .float32) + } + } + + @Test("index/array") + func test_index_array() throws { + try withIntegrationState(seed: 61633) { + let a = MLXRandom.normal([6, 5], dtype: .float32, loc: 0.0, scale: 1.0) + let indices = MLXRandom.randInt(low: 0, high: 6, [4], type: Int32.self) + let result = a[indices] + expectSummary( + result, + ArraySummary( + shape: [4, 5], + dtype: .float32, + mean: 0.04731393977999687, + minimum: -1.259320855140686, + maximum: 2.2643637657165527, + absoluteSum: 13.387282371520996, + positionChecksum: 7.656300354003906, + sampleIndices: [0, 4, 8, 11, 15, 19], + samples: [ + 0.24186666309833527, -0.5770583152770996, 0.4117497503757477, + -1.259320855140686, -0.39659786224365234, 2.2643637657165527, + ]), + tolerance: .float32) + } + } +} diff --git a/Tests/MLXIntegrationTests/Generated/GeneratedLinalgTests.swift b/Tests/MLXIntegrationTests/Generated/GeneratedLinalgTests.swift new file mode 100644 index 000000000..2be56a9e6 --- /dev/null +++ b/Tests/MLXIntegrationTests/Generated/GeneratedLinalgTests.swift @@ -0,0 +1,405 @@ +// Copyright © 2026 Apple Inc. +// +// GENERATED by tools/integration_tests -- DO NOT EDIT. +// +// Values in this file were produced by python mlx and are compared against +// the Swift API. Regenerate with: +// +// python3 tools/integration_tests/generate.py +// +// python mlx: 0.32.2.dev20260910+1f8e74e3f +// vendored mlx (Cmlx): 0.32.2 +// generator revision: 2 +// cases: 16 + +import Foundation +import MLX +import Testing + +@Suite("generated: Linalg") +struct GeneratedLinalgTests { + + @Test("norm") + func test_norm() throws { + try withIntegrationState(seed: 61303) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.norm(a, stream: .cpu) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 2.4654452800750732, + minimum: 2.4654452800750732, + maximum: 2.4654452800750732, + absoluteSum: 2.4654452800750732, + positionChecksum: 2.4654452800750732, + sampleIndices: [0], + samples: [2.4654452800750732]), + tolerance: .float32) + } + } + + @Test("norm/fro") + func test_norm_fro() throws { + try withIntegrationState(seed: 13358) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.norm(a, ord: .fro, stream: .cpu) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 3.750117301940918, + minimum: 3.750117301940918, + maximum: 3.750117301940918, + absoluteSum: 3.750117301940918, + positionChecksum: 3.750117301940918, + sampleIndices: [0], + samples: [3.750117301940918]), + tolerance: .float32) + } + } + + @Test("norm/ord1/axis") + func test_norm_ord1_axis() throws { + try withIntegrationState(seed: 50165) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.norm(a, ord: 1.0, axis: 0, stream: .cpu) + expectSummary( + result, + ArraySummary( + shape: [3], + dtype: .float32, + mean: 2.603952407836914, + minimum: 1.1081544160842896, + maximum: 4.008610248565674, + absoluteSum: 7.811856746673584, + positionChecksum: 4.678925196329753, + sampleIndices: [0, 1, 2], + samples: [2.69509220123291, 4.008610248565674, 1.1081544160842896]), + tolerance: .float32) + } + } + + @Test("norm/axes") + func test_norm_axes() throws { + try withIntegrationState(seed: 97742) { + let a = MLXRandom.normal([2, 4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.norm(a, axes: [1, 2], stream: .cpu) + expectSummary( + result, + ArraySummary( + shape: [2], + dtype: .float32, + mean: 3.1824984550476074, + minimum: 2.9801127910614014, + maximum: 3.3848843574523926, + absoluteSum: 6.364996910095215, + positionChecksum: 4.672554969787598, + sampleIndices: [0, 1], + samples: [3.3848843574523926, 2.9801127910614014]), + tolerance: .float32) + } + } + + @Test("norm/keepDims") + func test_norm_keepDims() throws { + try withIntegrationState(seed: 52042) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.norm(a, axis: -1, keepDims: true, stream: .cpu) + expectSummary( + result, + ArraySummary( + shape: [4, 1], + dtype: .float32, + mean: 1.7273744344711304, + minimum: 0.8218560218811035, + maximum: 2.5514132976531982, + absoluteSum: 6.9094977378845215, + positionChecksum: 3.605181932449341, + sampleIndices: [0, 1, 2, 3], + samples: [ + 2.5514132976531982, 2.026794910430908, 1.5094335079193115, + 0.8218560218811035, + ]), + tolerance: .float32) + } + } + + @Test("inv") + func test_inv() throws { + try withIntegrationState(seed: 45668) { + let a = MLXRandom.normal([4, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let spd = MLX.matmul(a, a.T) + 4.0 * MLX.eye(4) + let result = MLX.inv(spd, stream: .cpu) + expectSummary( + result, + ArraySummary( + shape: [4, 4], + dtype: .float32, + mean: 0.03864220157265663, + minimum: -0.040189146995544434, + maximum: 0.21930460631847382, + absoluteSum: 0.9170589447021484, + positionChecksum: 0.4542868137359619, + sampleIndices: [0, 3, 6, 9, 12, 15], + samples: [ + 0.18467867374420166, -0.01917562633752823, 0.0063242255710065365, + 0.006324225105345249, -0.01917562447488308, 0.15314576029777527, + ]), + tolerance: .float32) + } + } + + @Test("triInv") + func test_triInv() throws { + try withIntegrationState(seed: 99938) { + let a = MLXRandom.normal([4, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let spd = MLX.matmul(a, a.T) + 4.0 * MLX.eye(4) + let lower = MLX.tril(spd) + let result = MLX.triInv(lower, stream: .cpu) + expectSummary( + result, + ArraySummary( + shape: [4, 4], + dtype: .float32, + mean: 0.04744876176118851, + minimum: -0.017073122784495354, + maximum: 0.20120765268802643, + absoluteSum: 0.8115707635879517, + positionChecksum: 0.4696817994117737, + sampleIndices: [0, 3, 6, 9, 12, 15], + samples: [ + 0.16164816915988922, 0.0, 0.0, 0.03999091684818268, 0.006804905831813812, + 0.20120765268802643, + ]), + tolerance: .float32) + } + } + + @Test("cholesky") + func test_cholesky() throws { + try withIntegrationState(seed: 62767) { + let a = MLXRandom.normal([4, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let spd = MLX.matmul(a, a.T) + 4.0 * MLX.eye(4) + let result = MLX.cholesky(spd, stream: .cpu) + expectSummary( + result, + ArraySummary( + shape: [4, 4], + dtype: .float32, + mean: 0.8446391820907593, + minimum: -0.9616007208824158, + maximum: 3.0558056831359863, + absoluteSum: 17.262596130371094, + positionChecksum: 10.090898513793945, + sampleIndices: [0, 3, 6, 9, 12, 15], + samples: [ + 3.0558056831359863, 0.0, 0.0, -0.9125845432281494, 0.9799381494522095, + 2.471771001815796, + ]), + tolerance: .float32) + } + } + + @Test("cholesky/upper") + func test_cholesky_upper() throws { + try withIntegrationState(seed: 38847) { + let a = MLXRandom.normal([4, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let spd = MLX.matmul(a, a.T) + 4.0 * MLX.eye(4) + let result = MLX.cholesky(spd, upper: true, stream: .cpu) + expectSummary( + result, + ArraySummary( + shape: [4, 4], + dtype: .float32, + mean: 0.7331788539886475, + minimum: -0.7065603137016296, + maximum: 3.136165142059326, + absoluteSum: 13.893322944641113, + positionChecksum: 6.400030612945557, + sampleIndices: [0, 3, 6, 9, 12, 15], + samples: [ + 3.0352115631103516, -0.1763664186000824, -0.03923870995640755, 0.0, 0.0, + 2.291856050491333, + ]), + tolerance: .float32) + } + } + + @Test("choleskyInv") + func test_choleskyInv() throws { + try withIntegrationState(seed: 92684) { + let a = MLXRandom.normal([4, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let spd = MLX.matmul(a, a.T) + 4.0 * MLX.eye(4) + let l = MLX.cholesky(spd, stream: .cpu) + let result = MLX.choleskyInv(l, stream: .cpu) + expectSummary( + result, + ArraySummary( + shape: [4, 4], + dtype: .float32, + mean: 0.037270255386829376, + minimum: -0.026440046727657318, + maximum: 0.20521625876426697, + absoluteSum: 0.795662522315979, + positionChecksum: 0.3411255478858948, + sampleIndices: [0, 3, 6, 9, 12, 15], + samples: [ + 0.20521625876426697, -0.01557123102247715, -0.026440046727657318, + -0.026440046727657318, -0.01557123102247715, 0.07566997408866882, + ]), + tolerance: .float32) + } + } + + @Test("pinv") + func test_pinv() throws { + try withIntegrationState(seed: 49272) { + let a = MLXRandom.normal([5, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.pinv(a, stream: .cpu) + expectSummary( + result, + ArraySummary( + shape: [3, 5], + dtype: .float32, + mean: 0.14109282195568085, + minimum: -0.7837748527526855, + maximum: 0.786260187625885, + absoluteSum: 5.669179916381836, + positionChecksum: 2.6510897318522137, + sampleIndices: [0, 3, 6, 8, 11, 14], + samples: [ + -0.7837748527526855, -0.39348694682121277, 0.19790665805339813, + 0.3927387297153473, 0.528360903263092, 0.377410352230072, + ]), + tolerance: .float32) + } + } + + @Test("cross") + func test_cross() throws { + try withIntegrationState(seed: 83359) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.cross(a, b) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.19259056448936462, + minimum: -0.550257682800293, + maximum: 3.921427011489868, + absoluteSum: 6.5804123878479, + positionChecksum: 4.06615416208903, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.028709368780255318, -0.5001630187034607, -0.550257682800293, + 3.921427011489868, 0.14121165871620178, 0.2556290924549103, + ]), + tolerance: .float32) + } + } + + @Test("solve") + func test_solve() throws { + try withIntegrationState(seed: 29554) { + let a = MLXRandom.normal([4, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let spd = MLX.matmul(a, a.T) + 4.0 * MLX.eye(4) + let b = MLXRandom.normal([4, 2], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.solve(spd, b, stream: .cpu) + expectSummary( + result, + ArraySummary( + shape: [4, 2], + dtype: .float32, + mean: -0.011842835694551468, + minimum: -0.191814124584198, + maximum: 0.21963128447532654, + absoluteSum: 1.128840684890747, + positionChecksum: 0.6843239068984985, + sampleIndices: [0, 1, 3, 4, 6, 7], + samples: [ + -0.18137195706367493, -0.09362111240625381, -0.10529350489377975, + -0.191814124584198, 0.09047023206949234, 0.2069474756717682, + ]), + tolerance: .float32) + } + } + + @Test("solveTriangular") + func test_solveTriangular() throws { + try withIntegrationState(seed: 13290) { + let a = MLXRandom.normal([4, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let spd = MLX.matmul(a, a.T) + 4.0 * MLX.eye(4) + let lower = MLX.tril(spd) + let b = MLXRandom.normal([4, 2], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.solveTriangular(lower, b, stream: .cpu) + expectSummary( + result, + ArraySummary( + shape: [4, 2], + dtype: .float32, + mean: 0.061786651611328125, + minimum: -0.15068136155605316, + maximum: 0.28889700770378113, + absoluteSum: 0.938929557800293, + positionChecksum: 0.47943881154060364, + sampleIndices: [0, 1, 3, 4, 6, 7], + samples: [ + 0.06748572736978531, 0.28889700770378113, -0.03723353147506714, + 0.017042802646756172, -0.034403301775455475, 0.1020500585436821, + ]), + tolerance: .float32) + } + } + + @Test("det") + func test_det() throws { + try withIntegrationState(seed: 74576) { + let a = MLXRandom.normal([4, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let spd = MLX.matmul(a, a.T) + 4.0 * MLX.eye(4) + let result = MLX.det(spd, stream: .cpu) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 2897.23046875, + minimum: 2897.23046875, + maximum: 2897.23046875, + absoluteSum: 2897.23046875, + positionChecksum: 2897.23046875, + sampleIndices: [0], + samples: [2897.23046875]), + tolerance: .float32) + } + } + + @Test("eigvalsh") + func test_eigvalsh() throws { + // symmetric input: eigenvalues are returned in ascending order + try withIntegrationState(seed: 56419) { + let a = MLXRandom.normal([4, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let spd = MLX.matmul(a, a.T) + 4.0 * MLX.eye(4) + let result = MLX.eigvalsh(spd, stream: .cpu) + expectSummary( + result, + ArraySummary( + shape: [4], + dtype: .float32, + mean: 8.057674407958984, + minimum: 4.122847080230713, + maximum: 15.702493667602539, + absoluteSum: 32.23069763183594, + positionChecksum: 24.50469970703125, + sampleIndices: [0, 1, 2, 3], + samples: [ + 4.122847080230713, 6.13009786605835, 6.275260925292969, 15.702493667602539, + ]), + tolerance: .float32) + } + } +} diff --git a/Tests/MLXIntegrationTests/Generated/GeneratedLossesTests.swift b/Tests/MLXIntegrationTests/Generated/GeneratedLossesTests.swift new file mode 100644 index 000000000..8555b75f2 --- /dev/null +++ b/Tests/MLXIntegrationTests/Generated/GeneratedLossesTests.swift @@ -0,0 +1,555 @@ +// Copyright © 2026 Apple Inc. +// +// GENERATED by tools/integration_tests -- DO NOT EDIT. +// +// Values in this file were produced by python mlx and are compared against +// the Swift API. Regenerate with: +// +// python3 tools/integration_tests/generate.py +// +// python mlx: 0.32.2.dev20260910+1f8e74e3f +// vendored mlx (Cmlx): 0.32.2 +// generator revision: 2 +// cases: 23 + +import Foundation +import MLX +import MLXNN +import Testing + +@Suite("generated: Losses") +struct GeneratedLossesTests { + + @Test("crossEntropy/none") + func test_crossEntropy_none() throws { + try withIntegrationState(seed: 30964) { + let logits = MLXRandom.normal([4, 5], dtype: .float32, loc: 0.0, scale: 1.0) + let targets = MLXRandom.randInt(low: 0, high: 5, [4], type: Int32.self) + let result = MLXNN.crossEntropy(logits: logits, targets: targets, reduction: .none) + expectSummary( + result, + ArraySummary( + shape: [4], + dtype: .float32, + mean: 1.681290864944458, + minimum: 1.0067064762115479, + maximum: 2.6064746379852295, + absoluteSum: 6.725163459777832, + positionChecksum: 3.676652669906616, + sampleIndices: [0, 1, 2, 3], + samples: [ + 2.6064746379852295, 1.2626385688781738, 1.8493441343307495, + 1.0067064762115479, + ]), + tolerance: .float32) + } + } + + @Test("mseLoss/none") + func test_mseLoss_none() throws { + try withIntegrationState(seed: 14505) { + let predictions = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let targets = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.mseLoss(predictions: predictions, targets: targets, reduction: .none) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 2.298888921737671, + minimum: 0.011387436650693417, + maximum: 6.817906379699707, + absoluteSum: 27.586666107177734, + positionChecksum: 15.784774780273438, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.08654762804508209, 4.789333820343018, 0.49890097975730896, + 1.4658656120300293, 4.84807825088501, 1.1505988836288452, + ]), + tolerance: .float32) + } + } + + @Test("l1Loss/none") + func test_l1Loss_none() throws { + try withIntegrationState(seed: 96391) { + let predictions = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let targets = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.l1Loss(predictions: predictions, targets: targets, reduction: .none) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 1.1529018878936768, + minimum: 0.06354738771915436, + maximum: 2.881420135498047, + absoluteSum: 13.834822654724121, + positionChecksum: 6.185560862223308, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 1.8485625982284546, 1.470560073852539, 0.10926105082035065, + 0.631935179233551, 0.40655890107154846, 1.2638256549835205, + ]), + tolerance: .float32) + } + } + + @Test("crossEntropy/mean") + func test_crossEntropy_mean() throws { + try withIntegrationState(seed: 20203) { + let logits = MLXRandom.normal([4, 5], dtype: .float32, loc: 0.0, scale: 1.0) + let targets = MLXRandom.randInt(low: 0, high: 5, [4], type: Int32.self) + let result = MLXNN.crossEntropy(logits: logits, targets: targets, reduction: .mean) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 1.336897850036621, + minimum: 1.336897850036621, + maximum: 1.336897850036621, + absoluteSum: 1.336897850036621, + positionChecksum: 1.336897850036621, + sampleIndices: [0], + samples: [1.336897850036621]), + tolerance: .float32) + } + } + + @Test("mseLoss/mean") + func test_mseLoss_mean() throws { + try withIntegrationState(seed: 21174) { + let predictions = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let targets = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.mseLoss(predictions: predictions, targets: targets, reduction: .mean) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 1.5020225048065186, + minimum: 1.5020225048065186, + maximum: 1.5020225048065186, + absoluteSum: 1.5020225048065186, + positionChecksum: 1.5020225048065186, + sampleIndices: [0], + samples: [1.5020225048065186]), + tolerance: .float32) + } + } + + @Test("l1Loss/mean") + func test_l1Loss_mean() throws { + try withIntegrationState(seed: 72088) { + let predictions = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let targets = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.l1Loss(predictions: predictions, targets: targets, reduction: .mean) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 1.33334481716156, + minimum: 1.33334481716156, + maximum: 1.33334481716156, + absoluteSum: 1.33334481716156, + positionChecksum: 1.33334481716156, + sampleIndices: [0], + samples: [1.33334481716156]), + tolerance: .float32) + } + } + + @Test("crossEntropy/sum") + func test_crossEntropy_sum() throws { + try withIntegrationState(seed: 79221) { + let logits = MLXRandom.normal([4, 5], dtype: .float32, loc: 0.0, scale: 1.0) + let targets = MLXRandom.randInt(low: 0, high: 5, [4], type: Int32.self) + let result = MLXNN.crossEntropy(logits: logits, targets: targets, reduction: .sum) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 7.240591049194336, + minimum: 7.240591049194336, + maximum: 7.240591049194336, + absoluteSum: 7.240591049194336, + positionChecksum: 7.240591049194336, + sampleIndices: [0], + samples: [7.240591049194336]), + tolerance: .float32) + } + } + + @Test("mseLoss/sum") + func test_mseLoss_sum() throws { + try withIntegrationState(seed: 59747) { + let predictions = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let targets = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.mseLoss(predictions: predictions, targets: targets, reduction: .sum) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 35.32448959350586, + minimum: 35.32448959350586, + maximum: 35.32448959350586, + absoluteSum: 35.32448959350586, + positionChecksum: 35.32448959350586, + sampleIndices: [0], + samples: [35.32448959350586]), + tolerance: .float32) + } + } + + @Test("l1Loss/sum") + func test_l1Loss_sum() throws { + try withIntegrationState(seed: 61040) { + let predictions = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let targets = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.l1Loss(predictions: predictions, targets: targets, reduction: .sum) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 11.155969619750977, + minimum: 11.155969619750977, + maximum: 11.155969619750977, + absoluteSum: 11.155969619750977, + positionChecksum: 11.155969619750977, + sampleIndices: [0], + samples: [11.155969619750977]), + tolerance: .float32) + } + } + + @Test("crossEntropy/weights") + func test_crossEntropy_weights() throws { + try withIntegrationState(seed: 88185) { + let logits = MLXRandom.normal([4, 5], dtype: .float32, loc: 0.0, scale: 1.0) + let targets = MLXRandom.randInt(low: 0, high: 5, [4], type: Int32.self) + let weights = MLXRandom.uniform(low: 0.5, high: 1.5, [4], dtype: .float32) + let result = MLXNN.crossEntropy( + logits: logits, targets: targets, weights: weights, reduction: .mean) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 2.4085261821746826, + minimum: 2.4085261821746826, + maximum: 2.4085261821746826, + absoluteSum: 2.4085261821746826, + positionChecksum: 2.4085261821746826, + sampleIndices: [0], + samples: [2.4085261821746826]), + tolerance: .float32) + } + } + + @Test("crossEntropy/labelSmoothing") + func test_crossEntropy_labelSmoothing() throws { + try withIntegrationState(seed: 27380) { + let logits = MLXRandom.normal([4, 5], dtype: .float32, loc: 0.0, scale: 1.0) + let targets = MLXRandom.randInt(low: 0, high: 5, [4], type: Int32.self) + let result = MLXNN.crossEntropy( + logits: logits, targets: targets, labelSmoothing: 0.1, reduction: .mean) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 1.5042340755462646, + minimum: 1.5042340755462646, + maximum: 1.5042340755462646, + absoluteSum: 1.5042340755462646, + positionChecksum: 1.5042340755462646, + sampleIndices: [0], + samples: [1.5042340755462646]), + tolerance: .float32) + } + } + + @Test("crossEntropy/probabilities") + func test_crossEntropy_probabilities() throws { + try withIntegrationState(seed: 36) { + let logits = MLXRandom.normal([4, 5], dtype: .float32, loc: 0.0, scale: 1.0) + let t = MLXRandom.normal([4, 5], dtype: .float32, loc: 0.0, scale: 1.0) + let targets = MLX.softmax(t, axis: -1) + let result = MLXNN.crossEntropy(logits: logits, targets: targets, reduction: .mean) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 2.351641893386841, + minimum: 2.351641893386841, + maximum: 2.351641893386841, + absoluteSum: 2.351641893386841, + positionChecksum: 2.351641893386841, + sampleIndices: [0], + samples: [2.351641893386841]), + tolerance: .float32) + } + } + + @Test("binaryCrossEntropy/logits") + func test_binaryCrossEntropy_logits() throws { + try withIntegrationState(seed: 63990) { + let logits = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let targets = MLXRandom.bernoulli(0.5, [4, 3]) + let result = MLXNN.binaryCrossEntropy( + logits: logits, targets: targets.asType(.float32), reduction: .mean) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 0.7426967620849609, + minimum: 0.7426967620849609, + maximum: 0.7426967620849609, + absoluteSum: 0.7426967620849609, + positionChecksum: 0.7426967620849609, + sampleIndices: [0], + samples: [0.7426967620849609]), + tolerance: .float32) + } + } + + @Test("binaryCrossEntropy/probabilities") + func test_binaryCrossEntropy_probabilities() throws { + try withIntegrationState(seed: 37781) { + let p = MLXRandom.uniform(low: 0.1, high: 0.9, [4, 3], dtype: .float32) + let targets = MLXRandom.bernoulli(0.5, [4, 3]) + let result = MLXNN.binaryCrossEntropy( + logits: p, targets: targets.asType(.float32), withLogits: false, reduction: .mean) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 0.8758541941642761, + minimum: 0.8758541941642761, + maximum: 0.8758541941642761, + absoluteSum: 0.8758541941642761, + positionChecksum: 0.8758541941642761, + sampleIndices: [0], + samples: [0.8758541941642761]), + tolerance: .float32) + } + } + + @Test("nllLoss") + func test_nllLoss() throws { + try withIntegrationState(seed: 70502) { + let logits = MLXRandom.normal([4, 5], dtype: .float32, loc: 0.0, scale: 1.0) + let inputs = MLX.log(MLX.softmax(logits, axis: -1)) + let targets = MLXRandom.randInt(low: 0, high: 5, [4], type: Int32.self) + let result = MLXNN.nllLoss(inputs: inputs, targets: targets, reduction: .mean) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 1.6425096988677979, + minimum: 1.6425096988677979, + maximum: 1.6425096988677979, + absoluteSum: 1.6425096988677979, + positionChecksum: 1.6425096988677979, + sampleIndices: [0], + samples: [1.6425096988677979]), + tolerance: .float32) + } + } + + @Test("klDivLoss") + func test_klDivLoss() throws { + try withIntegrationState(seed: 90379) { + let a = MLXRandom.normal([4, 5], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([4, 5], dtype: .float32, loc: 0.0, scale: 1.0) + let inputs = MLX.log(MLX.softmax(a, axis: -1)) + let targets = MLX.log(MLX.softmax(b, axis: -1)) + let result = MLXNN.klDivLoss(inputs: inputs, targets: targets, reduction: .mean) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 0.4696533679962158, + minimum: 0.4696533679962158, + maximum: 0.4696533679962158, + absoluteSum: 0.4696533679962158, + positionChecksum: 0.4696533679962158, + sampleIndices: [0], + samples: [0.4696533679962158]), + tolerance: .float32) + } + } + + @Test("smoothL1Loss") + func test_smoothL1Loss() throws { + try withIntegrationState(seed: 16579) { + let predictions = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let targets = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.smoothL1Loss( + predictions: predictions, targets: targets, reduction: .mean) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 1.2783452272415161, + minimum: 1.2783452272415161, + maximum: 1.2783452272415161, + absoluteSum: 1.2783452272415161, + positionChecksum: 1.2783452272415161, + sampleIndices: [0], + samples: [1.2783452272415161]), + tolerance: .float32) + } + } + + @Test("smoothL1Loss/beta") + func test_smoothL1Loss_beta() throws { + try withIntegrationState(seed: 63527) { + let predictions = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let targets = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.smoothL1Loss( + predictions: predictions, targets: targets, beta: 0.5, reduction: .none) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 1.0599005222320557, + minimum: 0.011740408837795258, + maximum: 3.1852192878723145, + absoluteSum: 12.718805313110352, + positionChecksum: 6.934684753417969, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 1.0957175493240356, 0.011740408837795258, 1.24065363407135, + 0.014730868861079216, 1.9761757850646973, 0.6690434217453003, + ]), + tolerance: .float32) + } + } + + @Test("tripletLoss") + func test_tripletLoss() throws { + try withIntegrationState(seed: 98850) { + let anchors = MLXRandom.normal([4, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let positives = MLXRandom.normal([4, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let negatives = MLXRandom.normal([4, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.tripletLoss( + anchors: anchors, positives: positives, negatives: negatives, reduction: .mean) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 1.7783023118972778, + minimum: 1.7783023118972778, + maximum: 1.7783023118972778, + absoluteSum: 1.7783023118972778, + positionChecksum: 1.7783023118972778, + sampleIndices: [0], + samples: [1.7783023118972778]), + tolerance: .float32) + } + } + + @Test("hingeLoss") + func test_hingeLoss() throws { + try withIntegrationState(seed: 90714) { + let inputs = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let signs = MLXRandom.bernoulli(0.5, [4, 3]) + let targets = MLX.which(signs, 1.0, -1.0) + let result = MLXNN.hingeLoss(inputs: inputs, targets: targets, reduction: .mean) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 0.9594718813896179, + minimum: 0.9594718813896179, + maximum: 0.9594718813896179, + absoluteSum: 0.9594718813896179, + positionChecksum: 0.9594718813896179, + sampleIndices: [0], + samples: [0.9594718813896179]), + tolerance: .float32) + } + } + + @Test("huberLoss") + func test_huberLoss() throws { + try withIntegrationState(seed: 48851) { + let inputs = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let targets = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.huberLoss( + inputs: inputs, targets: targets, delta: 0.5, reduction: .mean) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 0.3303348124027252, + minimum: 0.3303348124027252, + maximum: 0.3303348124027252, + absoluteSum: 0.3303348124027252, + positionChecksum: 0.3303348124027252, + sampleIndices: [0], + samples: [0.3303348124027252]), + tolerance: .float32) + } + } + + @Test("logCoshLoss") + func test_logCoshLoss() throws { + try withIntegrationState(seed: 71520) { + let inputs = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let targets = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.logCoshLoss(inputs: inputs, targets: targets, reduction: .mean) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 0.8429071307182312, + minimum: 0.8429071307182312, + maximum: 0.8429071307182312, + absoluteSum: 0.8429071307182312, + positionChecksum: 0.8429071307182312, + sampleIndices: [0], + samples: [0.8429071307182312]), + tolerance: .float32) + } + } + + @Test("cosineSimilarityLoss") + func test_cosineSimilarityLoss() throws { + try withIntegrationState(seed: 8460) { + let x1 = MLXRandom.normal([4, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let x2 = MLXRandom.normal([4, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXNN.cosineSimilarityLoss(x1: x1, x2: x2, reduction: .mean) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: -0.043844811618328094, + minimum: -0.043844811618328094, + maximum: -0.043844811618328094, + absoluteSum: 0.043844811618328094, + positionChecksum: 0.043844811618328094, + sampleIndices: [0], + samples: [-0.043844811618328094]), + tolerance: .float32) + } + } +} diff --git a/Tests/MLXIntegrationTests/Generated/GeneratedModuleActivationsTests.swift b/Tests/MLXIntegrationTests/Generated/GeneratedModuleActivationsTests.swift new file mode 100644 index 000000000..ad9dbb228 --- /dev/null +++ b/Tests/MLXIntegrationTests/Generated/GeneratedModuleActivationsTests.swift @@ -0,0 +1,965 @@ +// Copyright © 2026 Apple Inc. +// +// GENERATED by tools/integration_tests -- DO NOT EDIT. +// +// Values in this file were produced by python mlx and are compared against +// the Swift API. Regenerate with: +// +// python3 tools/integration_tests/generate.py +// +// python mlx: 0.32.2.dev20260910+1f8e74e3f +// vendored mlx (Cmlx): 0.32.2 +// generator revision: 2 +// cases: 34 + +import Foundation +import MLX +import MLXNN +import Testing + +@Suite("generated: ModuleActivations") +struct GeneratedModuleActivationsTests { + + @Test("Identity") + func test_Identity() throws { + try withIntegrationState(seed: 34070) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = Identity() + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: -0.017332343384623528, + minimum: -3.227935791015625, + maximum: 2.948580741882324, + absoluteSum: 204.18060302734375, + positionChecksum: 105.40927124023438, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + -0.3334098160266876, 0.5547546148300171, -0.570229172706604, + -0.4087222218513489, -1.4258067607879639, -0.2681485116481781, + ]), + tolerance: .float32) + } + } + + @Test("Sigmoid") + func test_Sigmoid() throws { + try withIntegrationState(seed: 9913) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = Sigmoid() + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: 0.5089752078056335, + minimum: 0.04935688152909279, + maximum: 0.9119895696640015, + absoluteSum: 130.2976531982422, + positionChecksum: 67.42926788330078, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + 0.18784688413143158, 0.2168743908405304, 0.7701733708381653, + 0.1441628485918045, 0.6048126220703125, 0.3422277867794037, + ]), + tolerance: .float32) + } + } + + @Test("ReLU") + func test_ReLU() throws { + try withIntegrationState(seed: 16541) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = ReLU() + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: 0.3870805501937866, + minimum: 0.0, + maximum: 2.8602685928344727, + absoluteSum: 99.09262084960938, + positionChecksum: 47.08649444580078, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + 0.8868525624275208, 1.1041529178619385, 0.0, 1.1260745525360107, 0.0, 0.0, + ]), + tolerance: .float32) + } + } + + @Test("ReLU6") + func test_ReLU6() throws { + try withIntegrationState(seed: 54790) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = ReLU6() + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: 0.38837987184524536, + minimum: 0.0, + maximum: 2.9087584018707275, + absoluteSum: 99.42524719238281, + positionChecksum: 49.03082275390625, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + 0.0, 0.0, 0.3142980933189392, 0.0, 0.22747474908828735, 0.24886935949325562, + ]), + tolerance: .float32) + } + } + + @Test("ReLUSquared") + func test_ReLUSquared() throws { + try withIntegrationState(seed: 86178) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = ReLUSquared() + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: 0.5489815473556519, + minimum: 0.0, + maximum: 6.659243106842041, + absoluteSum: 140.53927612304688, + positionChecksum: 72.461181640625, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + 0.0, 0.13198506832122803, 2.2052338123321533, 0.0, 0.0, 0.19578227400779724, + ]), + tolerance: .float32) + } + } + + @Test("LeakyReLU") + func test_LeakyReLU() throws { + try withIntegrationState(seed: 42095) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = LeakyReLU() + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: 0.4470902681350708, + minimum: -0.03602779656648636, + maximum: 2.8485965728759766, + absoluteSum: 116.6517333984375, + positionChecksum: 62.598758697509766, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + -0.01762707717716694, 0.8074039816856384, 0.26561081409454346, + 0.8008984327316284, -0.017988381907343864, 1.6079342365264893, + ]), + tolerance: .float32) + } + } + + @Test("LeakyReLU/slope") + func test_LeakyReLU_slope() throws { + try withIntegrationState(seed: 76511) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = LeakyReLU(negativeSlope: 0.2) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: 0.3090943694114685, + minimum: -0.5242452025413513, + maximum: 3.7356669902801514, + absoluteSum: 117.27056884765625, + positionChecksum: 58.45794677734375, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + -0.061162084341049194, -0.01622505858540535, 0.19621963798999786, + 0.9918569326400757, -0.08134179562330246, 2.142888307571411, + ]), + tolerance: .float32) + } + } + + @Test("ELU") + func test_ELU() throws { + try withIntegrationState(seed: 85369) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = ELU() + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: 0.16009706258773804, + minimum: -0.9679020643234253, + maximum: 2.54449462890625, + absoluteSum: 162.6907958984375, + positionChecksum: 79.58113098144531, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + -0.8480066657066345, -0.45555174350738525, 0.7897208333015442, + 0.7451042532920837, -0.7147349119186401, 0.16544604301452637, + ]), + tolerance: .float32) + } + } + + @Test("ELU/alpha") + func test_ELU_alpha() throws { + try withIntegrationState(seed: 48075) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = ELU(alpha: 0.5) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: 0.2684254050254822, + minimum: -0.46672412753105164, + maximum: 3.140700578689575, + absoluteSum: 134.91976928710938, + positionChecksum: 70.10220336914062, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + 0.214599609375, 0.44013354182243347, 1.0033553838729858, 1.466883897781372, + 0.8385943174362183, 0.34053468704223633, + ]), + tolerance: .float32) + } + } + + @Test("CELU") + func test_CELU() throws { + try withIntegrationState(seed: 40999) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = CELU() + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: 0.27527904510498047, + minimum: -0.9589483737945557, + maximum: 2.492225408554077, + absoluteSum: 168.98414611816406, + positionChecksum: 87.66098022460938, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + -0.8132992386817932, 0.05264721438288689, 0.7257114052772522, + -0.36488646268844604, -0.6876139640808105, -0.2170543074607849, + ]), + tolerance: .float32) + } + } + + @Test("SiLU") + func test_SiLU() throws { + try withIntegrationState(seed: 94108) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = SiLU() + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: 0.29252636432647705, + minimum: -0.2782213091850281, + maximum: 3.081599473953247, + absoluteSum: 117.98585510253906, + positionChecksum: 61.460487365722656, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + 1.104568600654602, -0.07563341408967972, -0.049893204122781754, + -0.24418365955352783, -0.2484523504972458, 0.3073877692222595, + ]), + tolerance: .float32) + } + } + + @Test("SELU") + func test_SELU() throws { + try withIntegrationState(seed: 97176) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = SELU() + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: 0.12662386894226074, + minimum: -1.6931291818618774, + maximum: 2.9934651851654053, + absoluteSum: 210.03921508789062, + positionChecksum: 101.77021026611328, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + 0.8677928447723389, 0.6797854900360107, 0.36840230226516724, + -0.4589764475822449, -0.5336998105049133, 0.24484601616859436, + ]), + tolerance: .float32) + } + } + + @Test("Mish") + func test_Mish() throws { + try withIntegrationState(seed: 86013) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = Mish() + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: 0.1439376026391983, + minimum: -0.3088423013687134, + maximum: 2.430266857147217, + absoluteSum: 103.19464111328125, + positionChecksum: 48.830989837646484, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + 1.398360013961792, -0.014234477654099464, 0.8261752128601074, + -0.3047178089618683, -0.29602304100990295, -0.1958984136581421, + ]), + tolerance: .float32) + } + } + + @Test("Tanh") + func test_Tanh() throws { + try withIntegrationState(seed: 65516) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = Tanh() + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: -0.049693141132593155, + minimum: -0.9932478070259094, + maximum: 0.9843816161155701, + absoluteSum: 147.73788452148438, + positionChecksum: 76.96502685546875, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + 0.2730482816696167, 0.02616955153644085, -0.3087116479873657, + 0.5383175611495972, 0.22366760671138763, -0.2514197528362274, + ]), + tolerance: .float32) + } + } + + @Test("GELU") + func test_GELU() throws { + try withIntegrationState(seed: 56464) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = GELU() + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: 0.279835045337677, + minimum: -0.1699712574481964, + maximum: 2.9361231327056885, + absoluteSum: 100.90007781982422, + positionChecksum: 48.360328674316406, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + 0.3581196069717407, 0.6479476094245911, 2.9361231327056885, + 0.5618335008621216, 0.5602741837501526, 1.0405502319335938, + ]), + tolerance: .float32) + } + } + + @Test("GELU/precise") + func test_GELU_precise() throws { + try withIntegrationState(seed: 5645) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = GELU(approximation: .precise) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: 0.26011642813682556, + minimum: -0.17004051804542542, + maximum: 2.3098645210266113, + absoluteSum: 98.03465270996094, + positionChecksum: 50.40411376953125, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + 1.5097497701644897, -0.05096469447016716, -0.16896995902061462, + 1.1236003637313843, -0.16408808529376984, 0.2334529012441635, + ]), + tolerance: .float32) + } + } + + @Test("GELU/fast") + func test_GELU_fast() throws { + try withIntegrationState(seed: 11942) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = GELU(approximation: .fast) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: 0.2423984706401825, + minimum: -0.16360986232757568, + maximum: 2.6349453926086426, + absoluteSum: 93.47364807128906, + positionChecksum: 45.76264953613281, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + 1.0383901596069336, -0.09115244448184967, 1.3228586912155151, + -0.1617099493741989, -0.16003209352493286, 0.08449243754148483, + ]), + tolerance: .float32) + } + } + + @Test("HardSwish") + func test_HardSwish() throws { + try withIntegrationState(seed: 53183) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = HardSwish() + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: 0.1723780632019043, + minimum: -0.3749813735485077, + maximum: 3.0541574954986572, + absoluteSum: 100.03382873535156, + positionChecksum: 49.33647155761719, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + -0.3716834783554077, 0.10859835892915726, 0.44419172406196594, + 0.8135273456573486, -0.215162992477417, -0.1706112176179886, + ]), + tolerance: .float32) + } + } + + @Test("HardTanh") + func test_HardTanh() throws { + try withIntegrationState(seed: 35106) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = HardTanh() + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: -0.05235737934708595, + minimum: -1.0, + maximum: 1.0, + absoluteSum: 155.99127197265625, + positionChecksum: 75.52446746826172, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + 0.07812279462814331, 0.43642038106918335, -0.6854920983314514, -1.0, + -0.8840487003326416, -1.0, + ]), + tolerance: .float32) + } + } + + @Test("HardShrink") + func test_HardShrink() throws { + try withIntegrationState(seed: 10671) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = HardShrink() + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: -0.1441853642463684, + minimum: -2.248732566833496, + maximum: 2.803218364715576, + absoluteSum: 192.95077514648438, + positionChecksum: 102.37713623046875, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + -0.5038650631904602, 0.6754856109619141, 0.0, 0.0, -0.7869282364845276, + -1.0194823741912842, + ]), + tolerance: .float32) + } + } + + @Test("HardShrink/lambda") + func test_HardShrink_lambda() throws { + try withIntegrationState(seed: 29385) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = HardShrink(lambda: 0.2) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: 0.06390944868326187, + minimum: -2.980595827102661, + maximum: 2.9201619625091553, + absoluteSum: 205.49563598632812, + positionChecksum: 104.30592346191406, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + -0.8073967695236206, 2.180393934249878, 0.3150389492511749, + -0.29229095578193665, -1.2221306562423706, 0.0, + ]), + tolerance: .float32) + } + } + + @Test("Softplus") + func test_Softplus() throws { + try withIntegrationState(seed: 89825) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = Softplus() + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: 0.7824997305870056, + minimum: 0.0757046714425087, + maximum: 3.211472511291504, + absoluteSum: 200.31993103027344, + positionChecksum: 98.73367309570312, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + 1.0939661264419556, 0.4173877239227295, 0.9904651641845703, + 0.8673735857009888, 1.2018206119537354, 0.3981037139892578, + ]), + tolerance: .float32) + } + } + + @Test("Softsign") + func test_Softsign() throws { + try withIntegrationState(seed: 69518) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = Softsign() + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: -0.007850009948015213, + minimum: -0.686312198638916, + maximum: 0.7605715990066528, + absoluteSum: 97.87689208984375, + positionChecksum: 49.9131965637207, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + -0.2170916199684143, -0.6419945359230042, -0.6414955258369446, + -0.34234535694122314, 0.4860936999320984, 0.4500749111175537, + ]), + tolerance: .float32) + } + } + + @Test("Softshrink") + func test_Softshrink() throws { + try withIntegrationState(seed: 1067) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = Softshrink() + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: -0.00816175527870655, + minimum: -2.2436203956604004, + maximum: 2.848755121231079, + absoluteSum: 98.35267639160156, + positionChecksum: 51.70696258544922, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + 0.0, 0.14205098152160645, 0.0, 0.6208071708679199, 0.1955741047859192, + 0.4562690854072571, + ]), + tolerance: .float32) + } + } + + @Test("Softshrink/lambda") + func test_Softshrink_lambda() throws { + try withIntegrationState(seed: 65904) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = Softshrink(lambda: 0.2) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: -0.07054832577705383, + minimum: -3.6172547340393066, + maximum: 2.4144792556762695, + absoluteSum: 156.82997131347656, + positionChecksum: 75.9085693359375, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + -0.08811737596988678, 0.0, 0.861285388469696, -1.2284928560256958, 0.0, + -0.251159131526947, + ]), + tolerance: .float32) + } + } + + @Test("Softmax") + func test_Softmax() throws { + try withIntegrationState(seed: 76552) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = Softmax() + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: 0.0625, + minimum: 0.002436474896967411, + maximum: 0.45796260237693787, + absoluteSum: 16.0, + positionChecksum: 8.038972854614258, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + 0.011925941333174706, 0.015899190679192543, 0.038694579154253006, + 0.03524196892976761, 0.015219993889331818, 0.12464060634374619, + ]), + tolerance: .float32) + } + } + + @Test("Softmin") + func test_Softmin() throws { + try withIntegrationState(seed: 7761) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = Softmin() + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: 0.0625, + minimum: 0.0016566670965403318, + maximum: 0.6610931754112244, + absoluteSum: 16.0, + positionChecksum: 8.019933700561523, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + 0.2532690465450287, 0.0048106457106769085, 0.008263001218438148, + 0.09788679331541061, 0.13839659094810486, 0.05248511955142021, + ]), + tolerance: .float32) + } + } + + @Test("LogSoftmax") + func test_LogSoftmax() throws { + try withIntegrationState(seed: 82322) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = LogSoftmax() + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: -3.238546848297119, + minimum: -5.579071521759033, + maximum: -0.37829434871673584, + absoluteSum: 829.0679931640625, + positionChecksum: 420.2214050292969, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + -2.0407981872558594, -2.1493773460388184, -4.613822937011719, + -2.9762957096099854, -4.447968482971191, -1.2499594688415527, + ]), + tolerance: .float32) + } + } + + @Test("LogSigmoid") + func test_LogSigmoid() throws { + try withIntegrationState(seed: 79811) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = LogSigmoid() + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: -0.8042079210281372, + minimum: -2.9104275703430176, + maximum: -0.05430770292878151, + absoluteSum: 205.87722778320312, + positionChecksum: 103.44607543945312, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + -0.646495521068573, -0.8605636358261108, -0.6894464492797852, + -1.5663222074508667, -0.37833893299102783, -0.36881721019744873, + ]), + tolerance: .float32) + } + } + + @Test("Step") + func test_Step() throws { + try withIntegrationState(seed: 68019) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = Step() + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .int32, + mean: 0.484375, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 124.0, + positionChecksum: 63.3359375, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [1.0, 0.0, 1.0, 1.0, 0.0, 1.0]), + tolerance: .exact) + } + } + + @Test("Step/threshold") + func test_Step_threshold() throws { + try withIntegrationState(seed: 39538) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = Step(threshold: 0.5) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .int32, + mean: 0.2890625, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 74.0, + positionChecksum: 35.3515625, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [0.0, 1.0, 0.0, 0.0, 0.0, 1.0]), + tolerance: .exact) + } + } + + @Test("GLU") + func test_GLU() throws { + try withIntegrationState(seed: 70327) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = GLU() + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 8], + dtype: .float32, + mean: -0.019226614385843277, + minimum: -2.286162853240967, + maximum: 1.3630620241165161, + absoluteSum: 49.290138244628906, + positionChecksum: 22.783916473388672, + sampleIndices: [0, 25, 51, 76, 102, 127], + samples: [ + 0.8543118834495544, -0.3675357699394226, -0.8153498768806458, + 0.032235924154520035, -0.4751351773738861, 0.2292122095823288, + ]), + tolerance: .float32) + } + } + + @Test("PReLU") + func test_PReLU() throws { + try withIntegrationState(seed: 7115) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = PReLU() + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, [("weight", [1])]) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: 0.6037850379943848, + minimum: 0.00014645590272266418, + maximum: 3.3934574127197266, + absoluteSum: 154.5689697265625, + positionChecksum: 72.80538940429688, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + 0.044510796666145325, 0.5477976202964783, 0.10613550990819931, + 0.43876010179519653, 0.5276330709457397, 0.3869277536869049, + ]), + tolerance: .float32) + } + } + + @Test("PReLU/count") + func test_PReLU_count() throws { + try withIntegrationState(seed: 70876) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = PReLU(count: 16, value: 0.3) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, [("weight", [16])]) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: 0.4174256920814514, + minimum: -0.8470731973648071, + maximum: 2.8875372409820557, + absoluteSum: 128.29067993164062, + positionChecksum: 61.88407897949219, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + 1.7127238512039185, 0.9431889653205872, 0.14541825652122498, + -0.006025301292538643, 0.03635843098163605, 0.8640216588973999, + ]), + tolerance: .float32) + } + } +} diff --git a/Tests/MLXIntegrationTests/Generated/GeneratedModuleAttentionTests.swift b/Tests/MLXIntegrationTests/Generated/GeneratedModuleAttentionTests.swift new file mode 100644 index 000000000..b6e4522b9 --- /dev/null +++ b/Tests/MLXIntegrationTests/Generated/GeneratedModuleAttentionTests.swift @@ -0,0 +1,124 @@ +// Copyright © 2026 Apple Inc. +// +// GENERATED by tools/integration_tests -- DO NOT EDIT. +// +// Values in this file were produced by python mlx and are compared against +// the Swift API. Regenerate with: +// +// python3 tools/integration_tests/generate.py +// +// python mlx: 0.32.2.dev20260910+1f8e74e3f +// vendored mlx (Cmlx): 0.32.2 +// generator revision: 2 +// cases: 3 + +import Foundation +import MLX +import MLXNN +import Testing + +@Suite("generated: ModuleAttention") +struct GeneratedModuleAttentionTests { + + @Test("MultiHeadAttention") + func test_MultiHeadAttention() throws { + try withIntegrationState(seed: 24603) { + let x = MLXRandom.normal([2, 6, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = MultiHeadAttention(dimensions: 16, numHeads: 4) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters( + module, + [ + ("key_proj.weight", [16, 16]), ("out_proj.weight", [16, 16]), + ("query_proj.weight", [16, 16]), ("value_proj.weight", [16, 16]), + ]) + let result = module(x, keys: x, values: x) + expectSummary( + result, + ArraySummary( + shape: [2, 6, 16], + dtype: .float32, + mean: 0.05091768503189087, + minimum: -0.7064140439033508, + maximum: 1.1835263967514038, + absoluteSum: 40.15238952636719, + positionChecksum: 21.73735809326172, + sampleIndices: [0, 38, 76, 115, 153, 191], + samples: [ + 0.27396339178085327, 0.2314339280128479, 0.05882430821657181, + -0.01647859998047352, -0.12509478628635406, -0.09450595080852509, + ]), + tolerance: .float32) + } + } + + @Test("MultiHeadAttention/bias") + func test_MultiHeadAttention_bias() throws { + try withIntegrationState(seed: 27712) { + let x = MLXRandom.normal([2, 6, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = MultiHeadAttention(dimensions: 16, numHeads: 4, bias: true) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters( + module, + [ + ("key_proj.bias", [16]), ("key_proj.weight", [16, 16]), ("out_proj.bias", [16]), + ("out_proj.weight", [16, 16]), ("query_proj.bias", [16]), + ("query_proj.weight", [16, 16]), ("value_proj.bias", [16]), + ("value_proj.weight", [16, 16]), + ]) + let result = module(x, keys: x, values: x) + expectSummary( + result, + ArraySummary( + shape: [2, 6, 16], + dtype: .float32, + mean: 0.26262208819389343, + minimum: -1.393288493156433, + maximum: 2.2614102363586426, + absoluteSum: 78.51240539550781, + positionChecksum: 30.15442403157552, + sampleIndices: [0, 38, 76, 115, 153, 191], + samples: [ + 2.2614102363586426, 0.005398362874984741, 0.10388334095478058, + 0.21886420249938965, 0.3588539659976959, 0.5437532663345337, + ]), + tolerance: .float32) + } + } + + @Test("MultiHeadAttention/mask") + func test_MultiHeadAttention_mask() throws { + try withIntegrationState(seed: 65643) { + let x = MLXRandom.normal([2, 6, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let mask = MultiHeadAttention.createAdditiveCausalMask(6) + let module = MultiHeadAttention(dimensions: 16, numHeads: 4) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters( + module, + [ + ("key_proj.weight", [16, 16]), ("out_proj.weight", [16, 16]), + ("query_proj.weight", [16, 16]), ("value_proj.weight", [16, 16]), + ]) + let result = module(x, keys: x, values: x, mask: mask) + expectSummary( + result, + ArraySummary( + shape: [2, 6, 16], + dtype: .float32, + mean: 0.024126103147864342, + minimum: -1.0875264406204224, + maximum: 0.8838620781898499, + absoluteSum: 60.221214294433594, + positionChecksum: 30.494303385416668, + sampleIndices: [0, 38, 76, 115, 153, 191], + samples: [ + 0.47034403681755066, -0.05138539895415306, 0.000970873050391674, + -0.5258967280387878, -0.12062773108482361, -1.0875264406204224, + ]), + tolerance: .float32) + } + } +} diff --git a/Tests/MLXIntegrationTests/Generated/GeneratedModuleConvolutionTests.swift b/Tests/MLXIntegrationTests/Generated/GeneratedModuleConvolutionTests.swift new file mode 100644 index 000000000..72d10f2bb --- /dev/null +++ b/Tests/MLXIntegrationTests/Generated/GeneratedModuleConvolutionTests.swift @@ -0,0 +1,306 @@ +// Copyright © 2026 Apple Inc. +// +// GENERATED by tools/integration_tests -- DO NOT EDIT. +// +// Values in this file were produced by python mlx and are compared against +// the Swift API. Regenerate with: +// +// python3 tools/integration_tests/generate.py +// +// python mlx: 0.32.2.dev20260910+1f8e74e3f +// vendored mlx (Cmlx): 0.32.2 +// generator revision: 2 +// cases: 10 + +import Foundation +import MLX +import MLXNN +import Testing + +@Suite("generated: ModuleConvolution") +struct GeneratedModuleConvolutionTests { + + @Test("Conv1d") + func test_Conv1d() throws { + try withIntegrationState(seed: 80399) { + let x = MLXRandom.normal([2, 10, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let module = Conv1d(inputChannels: 4, outputChannels: 3, kernelSize: 3) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, [("bias", [3]), ("weight", [3, 3, 4])]) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 3], + dtype: .float32, + mean: -0.11053398251533508, + minimum: -2.7783169746398926, + maximum: 3.0954372882843018, + absoluteSum: 45.54235076904297, + positionChecksum: 22.108502705891926, + sampleIndices: [0, 9, 19, 28, 38, 47], + samples: [ + 3.0954372882843018, 0.5464149713516235, 0.0334894061088562, + -0.09258026629686356, 1.676621913909912, -0.2100052833557129, + ]), + tolerance: .float32) + } + } + + @Test("Conv1d/stridePadding") + func test_Conv1d_stridePadding() throws { + try withIntegrationState(seed: 36805) { + let x = MLXRandom.normal([2, 10, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let module = Conv1d( + inputChannels: 4, outputChannels: 3, kernelSize: 3, stride: 2, padding: 1) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, [("bias", [3]), ("weight", [3, 3, 4])]) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 5, 3], + dtype: .float32, + mean: -0.17636795341968536, + minimum: -2.4992713928222656, + maximum: 2.399353504180908, + absoluteSum: 23.765769958496094, + positionChecksum: 10.79921162923177, + sampleIndices: [0, 6, 12, 17, 23, 29], + samples: [ + -1.8267059326171875, 2.399353504180908, -1.4456043243408203, + -0.010623306035995483, -1.917959451675415, 0.3439697027206421, + ]), + tolerance: .float32) + } + } + + @Test("Conv1d/noBias") + func test_Conv1d_noBias() throws { + try withIntegrationState(seed: 91755) { + let x = MLXRandom.normal([2, 10, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let module = Conv1d(inputChannels: 4, outputChannels: 3, kernelSize: 3, bias: false) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, [("weight", [3, 3, 4])]) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 3], + dtype: .float32, + mean: 0.02788696438074112, + minimum: -1.910292148590088, + maximum: 1.8673605918884277, + absoluteSum: 34.24660110473633, + positionChecksum: 15.799858093261719, + sampleIndices: [0, 9, 19, 28, 38, 47], + samples: [ + -0.6793568730354309, -1.4656273126602173, -0.2600403130054474, + 0.4456599950790405, -0.5511552691459656, -0.8705923557281494, + ]), + tolerance: .float32) + } + } + + @Test("Conv2d") + func test_Conv2d() throws { + try withIntegrationState(seed: 77676) { + let x = MLXRandom.normal([2, 8, 8, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let module = Conv2d(inputChannels: 3, outputChannels: 4, kernelSize: 3) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, [("bias", [4]), ("weight", [4, 3, 3, 3])]) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 6, 6, 4], + dtype: .float32, + mean: -0.19274936616420746, + minimum: -5.951562881469727, + maximum: 5.163274765014648, + absoluteSum: 359.1477355957031, + positionChecksum: 194.03301323784723, + sampleIndices: [0, 57, 115, 172, 230, 287], + samples: [ + -3.717362642288208, -0.37188804149627686, 3.1527411937713623, + -1.7658523321151733, -0.7540384531021118, -1.4457536935806274, + ]), + tolerance: .float32) + } + } + + @Test("Conv2d/stridePadding") + func test_Conv2d_stridePadding() throws { + try withIntegrationState(seed: 70327) { + let x = MLXRandom.normal([2, 8, 8, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let module = Conv2d( + inputChannels: 3, outputChannels: 4, kernelSize: 3, stride: 2, padding: 1) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, [("bias", [4]), ("weight", [4, 3, 3, 3])]) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 4, 4, 4], + dtype: .float32, + mean: -0.16016966104507446, + minimum: -4.5320258140563965, + maximum: 3.6546685695648193, + absoluteSum: 128.35308837890625, + positionChecksum: 58.57609558105469, + sampleIndices: [0, 25, 51, 76, 102, 127], + samples: [ + -0.4038444757461548, -1.0283045768737793, 2.045264720916748, + 0.5583839416503906, -0.3402917981147766, -1.4707754850387573, + ]), + tolerance: .float32) + } + } + + @Test("Conv3d") + func test_Conv3d() throws { + try withIntegrationState(seed: 74637) { + let x = MLXRandom.normal([1, 4, 6, 6, 2], dtype: .float32, loc: 0.0, scale: 1.0) + let module = Conv3d(inputChannels: 2, outputChannels: 3, kernelSize: 2) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, [("bias", [3]), ("weight", [3, 2, 2, 2, 2])]) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [1, 3, 5, 5, 3], + dtype: .float32, + mean: -0.23182404041290283, + minimum: -3.550316095352173, + maximum: 3.392254590988159, + absoluteSum: 165.49330139160156, + positionChecksum: 86.74056423611111, + sampleIndices: [0, 45, 90, 134, 179, 224], + samples: [ + -0.2590126395225525, -1.4220865964889526, 1.4350488185882568, + -1.436377763748169, -0.26314473152160645, 0.30482929944992065, + ]), + tolerance: .float32) + } + } + + @Test("ConvTransposed1d") + func test_ConvTransposed1d() throws { + try withIntegrationState(seed: 94091) { + let x = MLXRandom.normal([2, 8, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let module = ConvTransposed1d(inputChannels: 4, outputChannels: 3, kernelSize: 3) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, [("bias", [3]), ("weight", [3, 3, 4])]) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 10, 3], + dtype: .float32, + mean: -0.16663511097431183, + minimum: -1.8646354675292969, + maximum: 1.3452529907226562, + absoluteSum: 25.450639724731445, + positionChecksum: 12.595225016276041, + sampleIndices: [0, 12, 24, 35, 47, 59], + samples: [ + -0.7107505798339844, 0.8472874164581299, -0.5283459424972534, + -0.22366145253181458, 0.5042819380760193, -0.5187655091285706, + ]), + tolerance: .float32) + } + } + + @Test("ConvTransposed1d/stride") + func test_ConvTransposed1d_stride() throws { + try withIntegrationState(seed: 96722) { + let x = MLXRandom.normal([2, 8, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let module = ConvTransposed1d( + inputChannels: 4, outputChannels: 3, kernelSize: 3, stride: 2, padding: 1, + outputPadding: 1) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, [("bias", [3]), ("weight", [3, 3, 4])]) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 16, 3], + dtype: .float32, + mean: -0.16872534155845642, + minimum: -2.399252414703369, + maximum: 1.8922030925750732, + absoluteSum: 51.313255310058594, + positionChecksum: 26.30189259847005, + sampleIndices: [0, 19, 38, 57, 76, 95], + samples: [ + 0.008320808410644531, -0.12555092573165894, 0.2811652421951294, + -1.5479183197021484, -0.2482323944568634, 0.49897417426109314, + ]), + tolerance: .float32) + } + } + + @Test("ConvTransposed2d") + func test_ConvTransposed2d() throws { + try withIntegrationState(seed: 17256) { + let x = MLXRandom.normal([2, 6, 6, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let module = ConvTransposed2d(inputChannels: 3, outputChannels: 4, kernelSize: 3) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, [("bias", [4]), ("weight", [4, 3, 3, 3])]) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 8, 4], + dtype: .float32, + mean: -0.13198184967041016, + minimum: -4.098501682281494, + maximum: 3.9713425636291504, + absoluteSum: 394.7406005859375, + positionChecksum: 200.25326538085938, + sampleIndices: [0, 102, 204, 307, 409, 511], + samples: [ + -0.45735645294189453, 0.35417571663856506, 0.3488615155220032, + 1.8767280578613281, -0.24553531408309937, -0.6553887128829956, + ]), + tolerance: .float32) + } + } + + @Test("ConvTransposed3d") + func test_ConvTransposed3d() throws { + try withIntegrationState(seed: 18025) { + let x = MLXRandom.normal([1, 4, 4, 4, 2], dtype: .float32, loc: 0.0, scale: 1.0) + let module = ConvTransposed3d(inputChannels: 2, outputChannels: 3, kernelSize: 2) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, [("bias", [3]), ("weight", [3, 2, 2, 2, 2])]) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [1, 5, 5, 5, 3], + dtype: .float32, + mean: -0.18090620636940002, + minimum: -4.558497428894043, + maximum: 3.8166122436523438, + absoluteSum: 291.21270751953125, + positionChecksum: 155.24589583333332, + sampleIndices: [0, 75, 150, 224, 299, 374], + samples: [ + -1.108421802520752, 0.058435797691345215, -0.2240082025527954, + 1.144876480102539, 0.4049822986125946, 0.36139535903930664, + ]), + tolerance: .float32) + } + } +} diff --git a/Tests/MLXIntegrationTests/Generated/GeneratedModuleDropoutTests.swift b/Tests/MLXIntegrationTests/Generated/GeneratedModuleDropoutTests.swift new file mode 100644 index 000000000..53ff30898 --- /dev/null +++ b/Tests/MLXIntegrationTests/Generated/GeneratedModuleDropoutTests.swift @@ -0,0 +1,109 @@ +// Copyright © 2026 Apple Inc. +// +// GENERATED by tools/integration_tests -- DO NOT EDIT. +// +// Values in this file were produced by python mlx and are compared against +// the Swift API. Regenerate with: +// +// python3 tools/integration_tests/generate.py +// +// python mlx: 0.32.2.dev20260910+1f8e74e3f +// vendored mlx (Cmlx): 0.32.2 +// generator revision: 2 +// cases: 3 + +import Foundation +import MLX +import MLXNN +import Testing + +@Suite("generated: ModuleDropout") +struct GeneratedModuleDropoutTests { + + @Test("Dropout/eval") + func test_Dropout_eval() throws { + // eval mode is the identity; training mode is stochastic and cannot be compared value-by-value + try withIntegrationState(seed: 81768) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = Dropout() + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: 0.05213969200849533, + minimum: -2.930783987045288, + maximum: 3.0198965072631836, + absoluteSum: 203.62631225585938, + positionChecksum: 104.69136810302734, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + 0.10624595731496811, -0.14009599387645721, 0.12543074786663055, + -1.2116183042526245, -0.6216195225715637, 0.6007198095321655, + ]), + tolerance: .float32) + } + } + + @Test("Dropout2d/eval") + func test_Dropout2d_eval() throws { + // eval mode is the identity; training mode is stochastic and cannot be compared value-by-value + try withIntegrationState(seed: 24737) { + let x = MLXRandom.normal([2, 8, 8, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let module = Dropout2d() + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 8, 4], + dtype: .float32, + mean: -0.05712714046239853, + minimum: -3.243155002593994, + maximum: 2.784644842147827, + absoluteSum: 398.7391357421875, + positionChecksum: 194.79966735839844, + sampleIndices: [0, 102, 204, 307, 409, 511], + samples: [ + 2.463557481765747, 0.979790449142456, -1.3742802143096924, + -1.0140464305877686, 0.9424789547920227, 0.03446979075670242, + ]), + tolerance: .float32) + } + } + + @Test("Dropout3d/eval") + func test_Dropout3d_eval() throws { + // eval mode is the identity; training mode is stochastic and cannot be compared value-by-value + try withIntegrationState(seed: 78645) { + let x = MLXRandom.normal([2, 4, 8, 8, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let module = Dropout3d() + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 4, 8, 8, 4], + dtype: .float32, + mean: 0.00032923324033617973, + minimum: -3.3692824840545654, + maximum: 3.9168732166290283, + absoluteSum: 1644.861328125, + positionChecksum: 813.3851318359375, + sampleIndices: [0, 409, 819, 1228, 1638, 2047], + samples: [ + 1.438454508781433, 0.07942906767129898, -1.1821718215942383, + -0.8269960880279541, 0.7035877704620361, 0.6526941061019897, + ]), + tolerance: .float32) + } + } +} diff --git a/Tests/MLXIntegrationTests/Generated/GeneratedModuleLinearTests.swift b/Tests/MLXIntegrationTests/Generated/GeneratedModuleLinearTests.swift new file mode 100644 index 000000000..4bd4c35d4 --- /dev/null +++ b/Tests/MLXIntegrationTests/Generated/GeneratedModuleLinearTests.swift @@ -0,0 +1,163 @@ +// Copyright © 2026 Apple Inc. +// +// GENERATED by tools/integration_tests -- DO NOT EDIT. +// +// Values in this file were produced by python mlx and are compared against +// the Swift API. Regenerate with: +// +// python3 tools/integration_tests/generate.py +// +// python mlx: 0.32.2.dev20260910+1f8e74e3f +// vendored mlx (Cmlx): 0.32.2 +// generator revision: 2 +// cases: 5 + +import Foundation +import MLX +import MLXNN +import Testing + +@Suite("generated: ModuleLinear") +struct GeneratedModuleLinearTests { + + @Test("Linear") + func test_Linear() throws { + try withIntegrationState(seed: 45463) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = Linear(16, 5) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, [("bias", [5]), ("weight", [5, 16])]) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 5], + dtype: .float32, + mean: -0.030131885781884193, + minimum: -3.8416616916656494, + maximum: 3.5228404998779297, + absoluteSum: 76.65104675292969, + positionChecksum: 48.569488525390625, + sampleIndices: [0, 16, 32, 47, 63, 79], + samples: [ + -0.2741870582103729, 0.2716590166091919, -0.3759540915489197, + -0.07041969895362854, -0.8177027106285095, 3.5228404998779297, + ]), + tolerance: .float32) + } + } + + @Test("Linear/noBias") + func test_Linear_noBias() throws { + try withIntegrationState(seed: 88724) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = Linear(16, 5, bias: false) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, [("weight", [5, 16])]) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 5], + dtype: .float32, + mean: 0.11451786011457443, + minimum: -2.7055249214172363, + maximum: 3.1095452308654785, + absoluteSum: 64.63607025146484, + positionChecksum: 34.513467407226564, + sampleIndices: [0, 16, 32, 47, 63, 79], + samples: [ + -0.16995857656002045, 1.6557778120040894, -0.1264149248600006, + 0.21584510803222656, -0.18836407363414764, -1.5945661067962646, + ]), + tolerance: .float32) + } + } + + @Test("Bilinear") + func test_Bilinear() throws { + try withIntegrationState(seed: 3925) { + let x = MLXRandom.normal([2, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let y = MLXRandom.normal([2, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let module = Bilinear(16, 8, 5) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, [("bias", [5]), ("weight", [5, 8, 16])]) + let result = module(x, y) + expectSummary( + result, + ArraySummary( + shape: [2, 5], + dtype: .float32, + mean: 0.6467210650444031, + minimum: -1.8467845916748047, + maximum: 4.493323802947998, + absoluteSum: 13.118966102600098, + positionChecksum: 4.266597747802734, + sampleIndices: [0, 2, 4, 5, 7, 9], + samples: [ + 4.493323802947998, 1.3232699632644653, -1.8467845916748047, + -0.7817299365997314, -0.029828175902366638, 0.7220736145973206, + ]), + tolerance: .float32) + } + } + + @Test("Embedding") + func test_Embedding() throws { + try withIntegrationState(seed: 44353) { + let x = MLXRandom.randInt(low: 0, high: 10, [2, 6], type: Int32.self) + let module = Embedding(embeddingCount: 10, dimensions: 8) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, [("weight", [10, 8])]) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 6, 8], + dtype: .float32, + mean: 0.07708332687616348, + minimum: -0.4000000059604645, + maximum: 0.48750001192092896, + absoluteSum: 27.900001525878906, + positionChecksum: 15.243748982747396, + sampleIndices: [0, 19, 38, 57, 76, 95], + samples: [ + -0.09999999403953552, 0.4375, -0.32499998807907104, 0.3125, + -0.3499999940395355, 0.2875000238418579, + ]), + tolerance: .float32) + } + } + + @Test("Embedding/asLinear") + func test_Embedding_asLinear() throws { + try withIntegrationState(seed: 8970) { + let x = MLXRandom.normal([2, 6, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let module = Embedding(embeddingCount: 10, dimensions: 8) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, [("weight", [10, 8])]) + let result = module.asLinear(x) + expectSummary( + result, + ArraySummary( + shape: [2, 6, 10], + dtype: .float32, + mean: -0.015905074775218964, + minimum: -2.8021512031555176, + maximum: 2.801405429840088, + absoluteSum: 62.93570327758789, + positionChecksum: 31.226318359375, + sampleIndices: [0, 24, 48, 71, 95, 119], + samples: [ + -2.8021512031555176, 0.007820483297109604, -0.06263187527656555, + 0.44406285881996155, -0.010466308332979679, -1.489911437034607, + ]), + tolerance: .float32) + } + } +} diff --git a/Tests/MLXIntegrationTests/Generated/GeneratedModuleNormalizationTests.swift b/Tests/MLXIntegrationTests/Generated/GeneratedModuleNormalizationTests.swift new file mode 100644 index 000000000..c36813f8b --- /dev/null +++ b/Tests/MLXIntegrationTests/Generated/GeneratedModuleNormalizationTests.swift @@ -0,0 +1,308 @@ +// Copyright © 2026 Apple Inc. +// +// GENERATED by tools/integration_tests -- DO NOT EDIT. +// +// Values in this file were produced by python mlx and are compared against +// the Swift API. Regenerate with: +// +// python3 tools/integration_tests/generate.py +// +// python mlx: 0.32.2.dev20260910+1f8e74e3f +// vendored mlx (Cmlx): 0.32.2 +// generator revision: 2 +// cases: 10 + +import Foundation +import MLX +import MLXNN +import Testing + +@Suite("generated: ModuleNormalization") +struct GeneratedModuleNormalizationTests { + + @Test("LayerNorm") + func test_LayerNorm() throws { + try withIntegrationState(seed: 42737) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = LayerNorm(dimensions: 16) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, [("bias", [16]), ("weight", [16])]) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: -0.0198993980884552, + minimum: -1.304665446281433, + maximum: 1.182969570159912, + absoluteSum: 74.0040054321289, + positionChecksum: 37.30052185058594, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + -0.03280201554298401, -0.5021070241928101, -0.24673715233802795, + 0.10571826994419098, 0.23079484701156616, 0.9026331305503845, + ]), + tolerance: .float32) + } + } + + @Test("LayerNorm/noAffine") + func test_LayerNorm_noAffine() throws { + try withIntegrationState(seed: 71715) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = LayerNorm(dimensions: 16, affine: false) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: 1.862645149230957e-09, + minimum: -2.7004194259643555, + maximum: 2.4202475547790527, + absoluteSum: 202.51004028320312, + positionChecksum: 101.72293090820312, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + 1.364133596420288, 0.1137239933013916, -0.1254626214504242, + 0.008553391322493553, -0.20014025270938873, -1.4844169616699219, + ]), + tolerance: .float32) + } + } + + @Test("RMSNorm") + func test_RMSNorm() throws { + try withIntegrationState(seed: 74037) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = RMSNorm(dimensions: 16) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, [("weight", [16])]) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: -0.02211989276111126, + minimum: -0.9623683094978333, + maximum: 1.1941437721252441, + absoluteSum: 51.25483322143555, + positionChecksum: 26.368240356445312, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + -0.029898809269070625, -0.015912216156721115, -0.06328903138637543, + -0.10268687456846237, -0.041741251945495605, -0.15935736894607544, + ]), + tolerance: .float32) + } + } + + @Test("GroupNorm") + func test_GroupNorm() throws { + try withIntegrationState(seed: 30539) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = GroupNorm(groupCount: 4, dimensions: 16) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, [("bias", [16]), ("weight", [16])]) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: -0.010309945791959763, + minimum: -1.3219797611236572, + maximum: 1.3990989923477173, + absoluteSum: 73.74345397949219, + positionChecksum: 36.934391021728516, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + 0.17130780220031738, -0.5502849221229553, -0.1820463389158249, + 0.13417498767375946, -0.11009815335273743, 1.3455901145935059, + ]), + tolerance: .float32) + } + } + + @Test("GroupNorm/pytorchCompatible") + func test_GroupNorm_pytorchCompatible() throws { + try withIntegrationState(seed: 49381) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = GroupNorm(groupCount: 4, dimensions: 16, pytorchCompatible: true) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, [("bias", [16]), ("weight", [16])]) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: -0.030907047912478447, + minimum: -1.2020220756530762, + maximum: 1.2239670753479004, + absoluteSum: 73.9453125, + positionChecksum: 37.40062713623047, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + -0.6553022861480713, -0.08870765566825867, -0.05848407745361328, + 0.06563569605350494, 0.47302502393722534, 0.6089176535606384, + ]), + tolerance: .float32) + } + } + + @Test("InstanceNorm") + func test_InstanceNorm() throws { + try withIntegrationState(seed: 7588) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = InstanceNorm(dimensions: 16) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: -9.313225746154785e-09, + minimum: -2.293649435043335, + maximum: 2.4967687129974365, + absoluteSum: 211.96717834472656, + positionChecksum: 104.72128295898438, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + 1.8193515539169312, -0.2942982316017151, -0.3463905453681946, + 0.1246948391199112, 2.166072130203247, -1.5744894742965698, + ]), + tolerance: .float32) + } + } + + @Test("InstanceNorm/affine") + func test_InstanceNorm_affine() throws { + try withIntegrationState(seed: 86086) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = InstanceNorm(dimensions: 16, affine: true) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, [("bias", [16]), ("weight", [16])]) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: -0.03125, + minimum: -1.5313172340393066, + maximum: 1.369132399559021, + absoluteSum: 72.45367431640625, + positionChecksum: 35.79180145263672, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + -0.10569164156913757, -0.18218281865119934, -0.2373899221420288, + 0.01955053210258484, 0.29726508259773254, 0.3344022035598755, + ]), + tolerance: .float32) + } + } + + @Test("BatchNorm/eval") + func test_BatchNorm_eval() throws { + // eval mode uses the (rewritten) running statistics + try withIntegrationState(seed: 32289) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = BatchNorm(featureCount: 16) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters( + module, + [("bias", [16]), ("running_mean", [16]), ("running_var", [16]), ("weight", [16])]) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: Double.nan, + minimum: Double.nan, + maximum: Double.nan, + absoluteSum: Double.nan, + positionChecksum: Double.nan, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + Double.nan, Double.nan, Double.nan, 0.19975754618644714, + 0.17086565494537354, 0.7938545942306519, + ]), + tolerance: .float32) + } + } + + @Test("BatchNorm/training") + func test_BatchNorm_training() throws { + // training mode normalizes with the batch statistics; deterministic, unlike Dropout + try withIntegrationState(seed: 8639) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = BatchNorm(featureCount: 16) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(true) + expectParameters( + module, + [("bias", [16]), ("running_mean", [16]), ("running_var", [16]), ("weight", [16])]) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: -0.0312499962747097, + minimum: -1.5580859184265137, + maximum: 1.4559376239776611, + absoluteSum: 74.08992004394531, + positionChecksum: 36.43730926513672, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + -0.5708073377609253, -0.2267388105392456, -0.17827975749969482, + 0.17880982160568237, 0.019031628966331482, 0.21674005687236786, + ]), + tolerance: .float32) + } + } + + @Test("BatchNorm/noTrackRunningStats") + func test_BatchNorm_noTrackRunningStats() throws { + try withIntegrationState(seed: 91221) { + let x = MLXRandom.normal([2, 8, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let module = BatchNorm(featureCount: 16, trackRunningStats: false) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, [("bias", [16]), ("weight", [16])]) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 16], + dtype: .float32, + mean: -0.03125, + minimum: -1.1940438747406006, + maximum: 1.2612149715423584, + absoluteSum: 74.57074737548828, + positionChecksum: 37.47613525390625, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + -0.3706926703453064, -0.5099658370018005, -0.023155465722084045, + 0.10782907903194427, 0.10105209052562714, 1.017714262008667, + ]), + tolerance: .float32) + } + } +} diff --git a/Tests/MLXIntegrationTests/Generated/GeneratedModulePoolingTests.swift b/Tests/MLXIntegrationTests/Generated/GeneratedModulePoolingTests.swift new file mode 100644 index 000000000..1208a34ec --- /dev/null +++ b/Tests/MLXIntegrationTests/Generated/GeneratedModulePoolingTests.swift @@ -0,0 +1,246 @@ +// Copyright © 2026 Apple Inc. +// +// GENERATED by tools/integration_tests -- DO NOT EDIT. +// +// Values in this file were produced by python mlx and are compared against +// the Swift API. Regenerate with: +// +// python3 tools/integration_tests/generate.py +// +// python mlx: 0.32.2.dev20260910+1f8e74e3f +// vendored mlx (Cmlx): 0.32.2 +// generator revision: 2 +// cases: 8 + +import Foundation +import MLX +import MLXNN +import Testing + +@Suite("generated: ModulePooling") +struct GeneratedModulePoolingTests { + + @Test("MaxPool1d") + func test_MaxPool1d() throws { + try withIntegrationState(seed: 29744) { + let x = MLXRandom.normal([2, 16, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let module = MaxPool1d(kernelSize: 2, stride: 2) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 4], + dtype: .float32, + mean: 0.5671930313110352, + minimum: -1.6935234069824219, + maximum: 2.3494832515716553, + absoluteSum: 54.285911560058594, + positionChecksum: 28.31031036376953, + sampleIndices: [0, 13, 25, 38, 50, 63], + samples: [ + -1.6935234069824219, 0.8108446002006531, 0.4986465275287628, + 0.9882473945617676, 0.5640528202056885, 0.9553364515304565, + ]), + tolerance: .float32) + } + } + + @Test("MaxPool2d") + func test_MaxPool2d() throws { + try withIntegrationState(seed: 98643) { + let x = MLXRandom.normal([2, 8, 8, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let module = MaxPool2d(kernelSize: 2, stride: 2) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 4, 4, 4], + dtype: .float32, + mean: 1.061241626739502, + minimum: -0.757491409778595, + maximum: 3.0933570861816406, + absoluteSum: 142.98167419433594, + positionChecksum: 71.34737396240234, + sampleIndices: [0, 25, 51, 76, 102, 127], + samples: [ + 1.6198300123214722, 0.06948284059762955, -0.12833164632320404, + 1.8321651220321655, 1.6542015075683594, 0.931220293045044, + ]), + tolerance: .float32) + } + } + + @Test("MaxPool3d") + func test_MaxPool3d() throws { + try withIntegrationState(seed: 32850) { + let x = MLXRandom.normal([2, 4, 8, 8, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let module = MaxPool3d(kernelSize: 2, stride: 2) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 2, 4, 4, 4], + dtype: .float32, + mean: 1.3603229522705078, + minimum: -0.38012799620628357, + maximum: 3.5154623985290527, + absoluteSum: 349.03045654296875, + positionChecksum: 171.13186645507812, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + 0.4982339143753052, 1.8628464937210083, 0.7978113293647766, + 1.2853648662567139, 1.1705279350280762, 1.8667374849319458, + ]), + tolerance: .float32) + } + } + + @Test("MaxPool2d/padding") + func test_MaxPool2d_padding() throws { + try withIntegrationState(seed: 60585) { + let x = MLXRandom.normal([2, 8, 8, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let module = MaxPool2d(kernelSize: 3, stride: 2, padding: 1) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 4, 4, 4], + dtype: .float32, + mean: 1.378739356994629, + minimum: 0.026348749175667763, + maximum: 2.6323399543762207, + absoluteSum: 176.4786376953125, + positionChecksum: 89.82207489013672, + sampleIndices: [0, 25, 51, 76, 102, 127], + samples: [ + 1.1064631938934326, 1.1356632709503174, 1.3230632543563843, + 1.220381259918213, 2.2429463863372803, 1.0948508977890015, + ]), + tolerance: .float32) + } + } + + @Test("AvgPool1d") + func test_AvgPool1d() throws { + try withIntegrationState(seed: 69169) { + let x = MLXRandom.normal([2, 16, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let module = AvgPool1d(kernelSize: 2, stride: 2) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 4], + dtype: .float32, + mean: 0.07272765040397644, + minimum: -2.300410270690918, + maximum: 1.7861367464065552, + absoluteSum: 36.99386978149414, + positionChecksum: 19.82023811340332, + sampleIndices: [0, 13, 25, 38, 50, 63], + samples: [ + 0.10242053866386414, 0.11492392420768738, 0.6828688383102417, + 0.174363374710083, -0.2848045229911804, -1.0003540515899658, + ]), + tolerance: .float32) + } + } + + @Test("AvgPool2d") + func test_AvgPool2d() throws { + try withIntegrationState(seed: 9458) { + let x = MLXRandom.normal([2, 8, 8, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let module = AvgPool2d(kernelSize: 2, stride: 2) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 4, 4, 4], + dtype: .float32, + mean: -0.08202866464853287, + minimum: -1.2255580425262451, + maximum: 1.491541862487793, + absoluteSum: 48.739749908447266, + positionChecksum: 24.986846923828125, + sampleIndices: [0, 25, 51, 76, 102, 127], + samples: [ + -0.4838985204696655, 0.0945008397102356, -0.11756622791290283, + 0.09166160225868225, -0.45479485392570496, 0.6930731534957886, + ]), + tolerance: .float32) + } + } + + @Test("AvgPool3d") + func test_AvgPool3d() throws { + try withIntegrationState(seed: 62547) { + let x = MLXRandom.normal([2, 4, 8, 8, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let module = AvgPool3d(kernelSize: 2, stride: 2) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 2, 4, 4, 4], + dtype: .float32, + mean: -0.001542090903967619, + minimum: -1.1280696392059326, + maximum: 0.8791224360466003, + absoluteSum: 73.494384765625, + positionChecksum: 36.46441650390625, + sampleIndices: [0, 51, 102, 153, 204, 255], + samples: [ + -0.4229593873023987, 0.4040684401988983, 0.1515209674835205, + 0.12673257291316986, 0.5045558214187622, 0.24828201532363892, + ]), + tolerance: .float32) + } + } + + @Test("AvgPool2d/padding") + func test_AvgPool2d_padding() throws { + try withIntegrationState(seed: 60609) { + let x = MLXRandom.normal([2, 8, 8, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let module = AvgPool2d(kernelSize: 3, stride: 2, padding: 1) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 4, 4, 4], + dtype: .float32, + mean: 0.038703806698322296, + minimum: -0.7088121771812439, + maximum: 0.7457450032234192, + absoluteSum: 24.897192001342773, + positionChecksum: 12.805156707763672, + sampleIndices: [0, 25, 51, 76, 102, 127], + samples: [ + 0.28262028098106384, -0.18140751123428345, 0.34491124749183655, + 0.2511276602745056, -0.11475422978401184, -0.3191181719303131, + ]), + tolerance: .float32) + } + } +} diff --git a/Tests/MLXIntegrationTests/Generated/GeneratedModulePositionalTests.swift b/Tests/MLXIntegrationTests/Generated/GeneratedModulePositionalTests.swift new file mode 100644 index 000000000..2ca306dee --- /dev/null +++ b/Tests/MLXIntegrationTests/Generated/GeneratedModulePositionalTests.swift @@ -0,0 +1,191 @@ +// Copyright © 2026 Apple Inc. +// +// GENERATED by tools/integration_tests -- DO NOT EDIT. +// +// Values in this file were produced by python mlx and are compared against +// the Swift API. Regenerate with: +// +// python3 tools/integration_tests/generate.py +// +// python mlx: 0.32.2.dev20260910+1f8e74e3f +// vendored mlx (Cmlx): 0.32.2 +// generator revision: 2 +// cases: 6 + +import Foundation +import MLX +import MLXNN +import Testing + +@Suite("generated: ModulePositional") +struct GeneratedModulePositionalTests { + + @Test("RoPE") + func test_RoPE() throws { + try withIntegrationState(seed: 70361) { + let x = MLXRandom.normal([2, 4, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let module = RoPE(dimensions: 8) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 4, 8], + dtype: .float32, + mean: -0.17599579691886902, + minimum: -2.1909637451171875, + maximum: 1.6637766361236572, + absoluteSum: 48.165287017822266, + positionChecksum: 21.670984268188477, + sampleIndices: [0, 13, 25, 38, 50, 63], + samples: [ + 1.1579456329345703, -1.2916109561920166, -0.03200756013393402, + 0.1322382688522339, 0.4934418797492981, 0.43062347173690796, + ]), + tolerance: .float32) + } + } + + @Test("RoPE/traditional") + func test_RoPE_traditional() throws { + try withIntegrationState(seed: 84217) { + let x = MLXRandom.normal([2, 4, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let module = RoPE(dimensions: 8, traditional: true, base: 500.0, scale: 0.5) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 4, 8], + dtype: .float32, + mean: 0.0022530164569616318, + minimum: -2.0894505977630615, + maximum: 2.9366118907928467, + absoluteSum: 46.10230255126953, + positionChecksum: 20.703702926635742, + sampleIndices: [0, 13, 25, 38, 50, 63], + samples: [ + 1.162760615348816, 0.5243428945541382, -0.13487887382507324, + -0.6033228039741516, 1.8898133039474487, -0.35043108463287354, + ]), + tolerance: .float32) + } + } + + @Test("RoPE/offset") + func test_RoPE_offset() throws { + try withIntegrationState(seed: 44446) { + let x = MLXRandom.normal([2, 4, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let module = RoPE(dimensions: 8) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x, offset: 2) + expectSummary( + result, + ArraySummary( + shape: [2, 4, 8], + dtype: .float32, + mean: -0.13316422700881958, + minimum: -2.70853853225708, + maximum: 2.8416025638580322, + absoluteSum: 53.91521453857422, + positionChecksum: 30.396228790283203, + sampleIndices: [0, 13, 25, 38, 50, 63], + samples: [ + -0.6359602212905884, -0.9670241475105286, -0.32228633761405945, + -1.538472056388855, -0.9030501842498779, 0.39927321672439575, + ]), + tolerance: .float32) + } + } + + @Test("SinusoidalPositionalEncoding") + func test_SinusoidalPositionalEncoding() throws { + try withIntegrationState(seed: 72004) { + let x = MLXRandom.normal([2, 4, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let module = SinusoidalPositionalEncoding(dimensions: 8) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 4, 8, 8], + dtype: .float32, + mean: 0.23017042875289917, + minimum: -0.4999251067638397, + maximum: 0.5, + absoluteSum: 137.57159423828125, + positionChecksum: 69.34072875976562, + sampleIndices: [0, 102, 204, 307, 409, 511], + samples: [ + 0.46070006489753723, 0.4999995231628418, 0.1861143559217453, + 8.04811788839288e-05, -0.020217712968587875, 0.5, + ]), + tolerance: .float32) + } + } + + @Test("SinusoidalPositionalEncoding/cosineFirst") + func test_SinusoidalPositionalEncoding_cosineFirst() throws { + try withIntegrationState(seed: 28369) { + let x = MLXRandom.normal([2, 4, 8], dtype: .float32, loc: 0.0, scale: 1.0) + let module = SinusoidalPositionalEncoding( + dimensions: 8, cosineFirst: true, fullTurns: true) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 4, 8, 8], + dtype: .float32, + mean: 0.17157743871212006, + minimum: -0.4993824362754822, + maximum: 0.5, + absoluteSum: 142.95547485351562, + positionChecksum: 71.41983795166016, + sampleIndices: [0, 102, 204, 307, 409, 511], + samples: [ + -0.4852396547794342, 0.007850918918848038, 0.09497486054897308, + 0.4999999701976776, 0.4368760883808136, 0.00014115065278019756, + ]), + tolerance: .float32) + } + } + + @Test("ALiBi") + func test_ALiBi() throws { + try withIntegrationState(seed: 51200) { + let scores = MLXRandom.normal([1, 4, 6, 6], dtype: .float32, loc: 0.0, scale: 1.0) + let module = ALiBi() + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(attentionScores: scores) + expectSummary( + result, + ArraySummary( + shape: [1, 4, 6, 6], + dtype: .float32, + mean: -0.1866210252046585, + minimum: -2.942293643951416, + maximum: 2.480872631072998, + absoluteSum: 124.28718566894531, + positionChecksum: 60.13202582465278, + sampleIndices: [0, 29, 57, 86, 114, 143], + samples: [ + -0.519200325012207, -1.532111644744873, 1.0923478603363037, + 0.1874494105577469, -1.2987788915634155, -1.300144910812378, + ]), + tolerance: .float32) + } + } +} diff --git a/Tests/MLXIntegrationTests/Generated/GeneratedModuleRecurrentTests.swift b/Tests/MLXIntegrationTests/Generated/GeneratedModuleRecurrentTests.swift new file mode 100644 index 000000000..8ebaa14d1 --- /dev/null +++ b/Tests/MLXIntegrationTests/Generated/GeneratedModuleRecurrentTests.swift @@ -0,0 +1,163 @@ +// Copyright © 2026 Apple Inc. +// +// GENERATED by tools/integration_tests -- DO NOT EDIT. +// +// Values in this file were produced by python mlx and are compared against +// the Swift API. Regenerate with: +// +// python3 tools/integration_tests/generate.py +// +// python mlx: 0.32.2.dev20260910+1f8e74e3f +// vendored mlx (Cmlx): 0.32.2 +// generator revision: 2 +// cases: 5 + +import Foundation +import MLX +import MLXNN +import Testing + +@Suite("generated: ModuleRecurrent") +struct GeneratedModuleRecurrentTests { + + @Test("RNN") + func test_RNN() throws { + try withIntegrationState(seed: 26873) { + let x = MLXRandom.normal([2, 5, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let module = RNN(inputSize: 4, hiddenSize: 3) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, [("Whh", [3, 3]), ("Wxh", [3, 4]), ("bias", [3])]) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 5, 3], + dtype: .float32, + mean: -0.1807376742362976, + minimum: -0.9347691535949707, + maximum: 0.7658460736274719, + absoluteSum: 11.239614486694336, + positionChecksum: 6.376473999023437, + sampleIndices: [0, 6, 12, 17, 23, 29], + samples: [ + -0.35051560401916504, -0.33143526315689087, -0.8694400787353516, + -0.299630343914032, -0.5201353430747986, 0.19474464654922485, + ]), + tolerance: .float32) + } + } + + @Test("RNN/noBias") + func test_RNN_noBias() throws { + try withIntegrationState(seed: 54659) { + let x = MLXRandom.normal([2, 5, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let module = RNN(inputSize: 4, hiddenSize: 3, bias: false) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, [("Whh", [3, 3]), ("Wxh", [3, 4])]) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 5, 3], + dtype: .float32, + mean: -0.04676102474331856, + minimum: -0.7166611552238464, + maximum: 0.7711976170539856, + absoluteSum: 7.510426044464111, + positionChecksum: 3.6819297790527346, + sampleIndices: [0, 6, 12, 17, 23, 29], + samples: [ + 0.40871885418891907, -0.7166611552238464, 0.23077711462974548, + 0.0683915838599205, 0.21350324153900146, -0.5236168503761292, + ]), + tolerance: .float32) + } + } + + @Test("RNN/hidden") + func test_RNN_hidden() throws { + try withIntegrationState(seed: 52626) { + let x = MLXRandom.normal([2, 5, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let hidden = MLXRandom.normal([2, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let module = RNN(inputSize: 4, hiddenSize: 3) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, [("Whh", [3, 3]), ("Wxh", [3, 4]), ("bias", [3])]) + let result = module(x, hidden: hidden) + expectSummary( + result, + ArraySummary( + shape: [2, 5, 3], + dtype: .float32, + mean: -0.0772470086812973, + minimum: -0.911475419998169, + maximum: 0.95676189661026, + absoluteSum: 12.752267837524414, + positionChecksum: 6.333803304036459, + sampleIndices: [0, 6, 12, 17, 23, 29], + samples: [ + 0.6317100524902344, -0.3342420160770416, -0.4475969076156616, + -0.8060013651847839, 0.5694852471351624, -0.5043428540229797, + ]), + tolerance: .float32) + } + } + + @Test("GRU") + func test_GRU() throws { + try withIntegrationState(seed: 26515) { + let x = MLXRandom.normal([2, 5, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let module = GRU(inputSize: 4, hiddenSize: 3) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, [("Wh", [9, 3]), ("Wx", [9, 4]), ("b", [9]), ("bhn", [3])]) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 5, 3], + dtype: .float32, + mean: 0.3049412667751312, + minimum: -0.1211482584476471, + maximum: 0.6732051968574524, + absoluteSum: 9.443852424621582, + positionChecksum: 5.356678263346354, + sampleIndices: [0, 6, 12, 17, 23, 29], + samples: [ + 0.00961192138493061, 0.2725975513458252, -0.1211482584476471, + 0.378728985786438, 0.33235982060432434, 0.5703566074371338, + ]), + tolerance: .float32) + } + } + + @Test("GRU/noBias") + func test_GRU_noBias() throws { + try withIntegrationState(seed: 84535) { + let x = MLXRandom.normal([2, 5, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let module = GRU(inputSize: 4, hiddenSize: 3, bias: false) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, [("Wh", [9, 3]), ("Wx", [9, 4])]) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 5, 3], + dtype: .float32, + mean: 0.028713371604681015, + minimum: -0.3385675251483917, + maximum: 0.5721112489700317, + absoluteSum: 5.291589260101318, + positionChecksum: 2.6151153564453127, + sampleIndices: [0, 6, 12, 17, 23, 29], + samples: [ + 0.03434830531477928, 0.13748443126678467, 0.4411085546016693, + -0.020659416913986206, -0.03192802146077156, -0.3385675251483917, + ]), + tolerance: .float32) + } + } +} diff --git a/Tests/MLXIntegrationTests/Generated/GeneratedModuleUpsampleTests.swift b/Tests/MLXIntegrationTests/Generated/GeneratedModuleUpsampleTests.swift new file mode 100644 index 000000000..51808a0ea --- /dev/null +++ b/Tests/MLXIntegrationTests/Generated/GeneratedModuleUpsampleTests.swift @@ -0,0 +1,162 @@ +// Copyright © 2026 Apple Inc. +// +// GENERATED by tools/integration_tests -- DO NOT EDIT. +// +// Values in this file were produced by python mlx and are compared against +// the Swift API. Regenerate with: +// +// python3 tools/integration_tests/generate.py +// +// python mlx: 0.32.2.dev20260910+1f8e74e3f +// vendored mlx (Cmlx): 0.32.2 +// generator revision: 2 +// cases: 5 + +import Foundation +import MLX +import MLXNN +import Testing + +@Suite("generated: ModuleUpsample") +struct GeneratedModuleUpsampleTests { + + @Test("Upsample/nearest") + func test_Upsample_nearest() throws { + try withIntegrationState(seed: 39945) { + let x = MLXRandom.normal([2, 4, 4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let module = Upsample(scaleFactor: 2.0) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 8, 3], + dtype: .float32, + mean: -0.08313984423875809, + minimum: -1.9472978115081787, + maximum: 1.9504168033599854, + absoluteSum: 274.71282958984375, + positionChecksum: 145.42476399739584, + sampleIndices: [0, 77, 153, 230, 306, 383], + samples: [ + -0.9178804755210876, -0.5090406537055969, 0.6381695866584778, + 0.4494686424732208, 0.5923683643341064, -1.504755973815918, + ]), + tolerance: .float32) + } + } + + @Test("Upsample/linear") + func test_Upsample_linear() throws { + try withIntegrationState(seed: 2223) { + let x = MLXRandom.normal([2, 4, 4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let module = Upsample(scaleFactor: 2.0, mode: .linear()) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 8, 3], + dtype: .float32, + mean: 0.07992034405469894, + minimum: -2.166339159011841, + maximum: 2.1690256595611572, + absoluteSum: 198.35586547851562, + positionChecksum: 104.31976318359375, + sampleIndices: [0, 77, 153, 230, 306, 383], + samples: [ + -0.6088932156562805, 0.43297699093818665, -0.9983523488044739, + 0.5581990480422974, 0.9666178822517395, 0.052967239171266556, + ]), + tolerance: .float32) + } + } + + @Test("Upsample/linear/alignCorners") + func test_Upsample_linear_alignCorners() throws { + try withIntegrationState(seed: 15253) { + let x = MLXRandom.normal([2, 4, 4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let module = Upsample(scaleFactor: 2.0, mode: .linear(alignCorners: true)) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 8, 3], + dtype: .float32, + mean: -0.15478087961673737, + minimum: -2.823728561401367, + maximum: 1.8036024570465088, + absoluteSum: 241.51821899414062, + positionChecksum: 117.62404378255208, + sampleIndices: [0, 77, 153, 230, 306, 383], + samples: [ + -2.823728561401367, -0.9696316719055176, 0.006423458456993103, + -0.19877995550632477, -0.32860827445983887, 0.10723376274108887, + ]), + tolerance: .float32) + } + } + + @Test("Upsample/cubic") + func test_Upsample_cubic() throws { + try withIntegrationState(seed: 98182) { + let x = MLXRandom.normal([2, 4, 4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let module = Upsample(scaleFactor: 2.0, mode: .cubic()) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 8, 3], + dtype: .float32, + mean: -0.05168304964900017, + minimum: -2.2958436012268066, + maximum: 2.3127639293670654, + absoluteSum: 248.95880126953125, + positionChecksum: 116.72739664713542, + sampleIndices: [0, 77, 153, 230, 306, 383], + samples: [ + -1.6836540699005127, 0.061760637909173965, -0.49417844414711, + -0.23918810486793518, 0.038516491651535034, 0.0978599488735199, + ]), + tolerance: .float32) + } + } + + @Test("Upsample/scaleFactors") + func test_Upsample_scaleFactors() throws { + try withIntegrationState(seed: 48748) { + let x = MLXRandom.normal([2, 4, 4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let module = Upsample(scaleFactor: [2.0, 3.0]) + module.update(parameters: module.mapParameters { deterministicParameter($0) }) + module.train(false) + expectParameters(module, []) + let result = module(x) + expectSummary( + result, + ArraySummary( + shape: [2, 8, 12, 3], + dtype: .float32, + mean: -0.11570308357477188, + minimum: -1.899631142616272, + maximum: 2.335793972015381, + absoluteSum: 417.2052001953125, + positionChecksum: 212.24254014756946, + sampleIndices: [0, 115, 230, 345, 460, 575], + samples: [ + -0.7018576264381409, 0.24986159801483154, -1.3673666715621948, + 0.004559833090752363, 0.339202344417572, 1.4709136486053467, + ]), + tolerance: .float32) + } + } +} diff --git a/Tests/MLXIntegrationTests/Generated/GeneratedOptimizerFunctionsTests.swift b/Tests/MLXIntegrationTests/Generated/GeneratedOptimizerFunctionsTests.swift new file mode 100644 index 000000000..ab8ee659d --- /dev/null +++ b/Tests/MLXIntegrationTests/Generated/GeneratedOptimizerFunctionsTests.swift @@ -0,0 +1,297 @@ +// Copyright © 2026 Apple Inc. +// +// GENERATED by tools/integration_tests -- DO NOT EDIT. +// +// Values in this file were produced by python mlx and are compared against +// the Swift API. Regenerate with: +// +// python3 tools/integration_tests/generate.py +// +// python mlx: 0.32.2.dev20260910+1f8e74e3f +// vendored mlx (Cmlx): 0.32.2 +// generator revision: 2 +// cases: 11 + +import Foundation +import MLX +import MLXNN +import Testing + +@testable import MLXOptimizers + +@Suite("generated: OptimizerFunctions") +struct GeneratedOptimizerFunctionsTests { + + @Test("clipGradNorm/norm") + func test_clipGradNorm_norm() throws { + // the returned norm is the norm *before* clipping + try withIntegrationState(seed: 74913) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let result = clipGradNorm( + gradients: ModuleParameters.unflattened([("a", a), ("b", b)]), maxNorm: 1.0 + ).1 + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 4.2577900886535645, + minimum: 4.2577900886535645, + maximum: 4.2577900886535645, + absoluteSum: 4.2577900886535645, + positionChecksum: 4.2577900886535645, + sampleIndices: [0], + samples: [4.2577900886535645]), + tolerance: .float32) + } + } + + @Test("clipGradNorm/clipped") + func test_clipGradNorm_clipped() throws { + // the norm of a normal [4, 3] + [5] is well above 1, so this clips + try withIntegrationState(seed: 32730) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let result = clipGradNorm( + gradients: ModuleParameters.unflattened([("a", a), ("b", b)]), maxNorm: 1.0 + ).0[unwrapping: "a"]! + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.141945943236351, + minimum: -0.134745791554451, + maximum: 0.4862045347690582, + absoluteSum: 2.1906065940856934, + positionChecksum: 1.2294964790344238, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.13243113458156586, 0.17849785089492798, 0.1051086038351059, + 0.24804984033107758, 0.32634231448173523, -0.0005411746096797287, + ]), + tolerance: .float32) + } + } + + @Test("clipGradNorm/underTheLimit") + func test_clipGradNorm_underTheLimit() throws { + // already inside the limit, so the gradients pass through unchanged + try withIntegrationState(seed: 41005) { + let a = MLXRandom.uniform(low: -0.02, high: 0.02, [4, 3], dtype: .float32) + let result = clipGradNorm( + gradients: ModuleParameters.unflattened([("a", a)]), maxNorm: 1.0 + ).0[unwrapping: "a"]! + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.0004688164626713842, + minimum: -0.016733653843402863, + maximum: 0.018089659512043, + absoluteSum: 0.12719818949699402, + positionChecksum: 0.06582144896189372, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.0020372532308101654, 0.01661718264222145, 0.008448302745819092, + -0.016733653843402863, -0.01313621737062931, -0.006226498633623123, + ]), + tolerance: .float32) + } + } + + @Test("clipGradNorm/array") + func test_clipGradNorm_array() throws { + // the Collection overload should agree with the ModuleParameters one + try withIntegrationState(seed: 30787) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let result = clipGradNorm(gradients: [a, b], maxNorm: 0.5).1 + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 4.24692964553833, + minimum: 4.24692964553833, + maximum: 4.24692964553833, + absoluteSum: 4.24692964553833, + positionChecksum: 4.24692964553833, + sampleIndices: [0], + samples: [4.24692964553833]), + tolerance: .float32) + } + } + + @Test("newtonSchulz/tall/steps1") + func test_newtonSchulz_tall_steps1() throws { + try withIntegrationState(seed: 35922) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = Muon(learningRate: 0.1).zeropowerViaNewtonSchulz5(a, steps: 1) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.14183175563812256, + minimum: -0.7796086072921753, + maximum: 0.9717819690704346, + absoluteSum: 5.605131149291992, + positionChecksum: 3.1666199366251626, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.34602782130241394, 0.6664655804634094, 0.5670061111450195, + 0.9717819690704346, 0.24358342587947845, -0.7796086072921753, + ]), + tolerance: .float32) + } + } + + @Test("newtonSchulz/tall/steps5") + func test_newtonSchulz_tall_steps5() throws { + try withIntegrationState(seed: 97931) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = Muon(learningRate: 0.1).zeropowerViaNewtonSchulz5(a, steps: 5) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.055028900504112244, + minimum: -0.9158430099487305, + maximum: 0.7314736247062683, + absoluteSum: 5.023744106292725, + positionChecksum: 2.7080958684285483, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.7314736247062683, -0.21115955710411072, -0.19221490621566772, + 0.6600404977798462, 0.5243890881538391, -0.13377171754837036, + ]), + tolerance: .float32) + } + } + + @Test("newtonSchulz/wide/steps1") + func test_newtonSchulz_wide_steps1() throws { + try withIntegrationState(seed: 50292) { + let a = MLXRandom.normal([3, 6], dtype: .float32, loc: 0.0, scale: 1.0) + let result = Muon(learningRate: 0.1).zeropowerViaNewtonSchulz5(a, steps: 1) + expectSummary( + result, + ArraySummary( + shape: [3, 6], + dtype: .float32, + mean: 0.029883000999689102, + minimum: -0.7599508762359619, + maximum: 0.7960574626922607, + absoluteSum: 6.28695821762085, + positionChecksum: 3.4281453026665583, + sampleIndices: [0, 3, 7, 10, 14, 17], + samples: [ + -0.05090735852718353, 0.15048637986183167, 0.5751805305480957, + 0.04742121696472168, -0.30298084020614624, 0.017879322171211243, + ]), + tolerance: .float32) + } + } + + @Test("newtonSchulz/wide/steps5") + func test_newtonSchulz_wide_steps5() throws { + try withIntegrationState(seed: 47885) { + let a = MLXRandom.normal([3, 6], dtype: .float32, loc: 0.0, scale: 1.0) + let result = Muon(learningRate: 0.1).zeropowerViaNewtonSchulz5(a, steps: 5) + expectSummary( + result, + ArraySummary( + shape: [3, 6], + dtype: .float32, + mean: 0.11007862538099289, + minimum: -0.3809030055999756, + maximum: 0.8513897657394409, + absoluteSum: 5.038133144378662, + positionChecksum: 2.663041008843316, + sampleIndices: [0, 3, 7, 10, 14, 17], + samples: [ + -0.2114892601966858, -0.3663442134857178, 0.32633036375045776, + 0.10934531688690186, 0.2657369375228882, 0.4214053153991699, + ]), + tolerance: .float32) + } + } + + @Test("newtonSchulz/square/steps1") + func test_newtonSchulz_square_steps1() throws { + try withIntegrationState(seed: 54544) { + let a = MLXRandom.normal([4, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let result = Muon(learningRate: 0.1).zeropowerViaNewtonSchulz5(a, steps: 1) + expectSummary( + result, + ArraySummary( + shape: [4, 4], + dtype: .float32, + mean: -0.01559077575802803, + minimum: -0.7211068868637085, + maximum: 0.7834442853927612, + absoluteSum: 5.679539680480957, + positionChecksum: 3.048487663269043, + sampleIndices: [0, 3, 6, 9, 12, 15], + samples: [ + -0.13931389153003693, -0.3704991936683655, -0.3554251492023468, + 0.6534769535064697, -0.014799904078245163, -0.7211068868637085, + ]), + tolerance: .float32) + } + } + + @Test("newtonSchulz/square/steps5") + func test_newtonSchulz_square_steps5() throws { + try withIntegrationState(seed: 28105) { + let a = MLXRandom.normal([4, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let result = Muon(learningRate: 0.1).zeropowerViaNewtonSchulz5(a, steps: 5) + expectSummary( + result, + ArraySummary( + shape: [4, 4], + dtype: .float32, + mean: -0.1895444095134735, + minimum: -0.8144998550415039, + maximum: 0.6692322492599487, + absoluteSum: 6.853260040283203, + positionChecksum: 3.5909621715545654, + sampleIndices: [0, 3, 6, 9, 12, 15], + samples: [ + -0.8144998550415039, -0.253656804561615, -0.3669908046722412, + -0.5989289283752441, 0.261578232049942, -0.6159399747848511, + ]), + tolerance: .float32) + } + } + + @Test("newtonSchulz/nearlyOrthogonal") + func test_newtonSchulz_nearlyOrthogonal() throws { + // an already orthogonal input is the fixed point of the iteration + try withIntegrationState(seed: 28031) { + let a = MLXRandom.normal([4, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let q = MLX.qr(a, stream: .cpu).0 + let result = Muon(learningRate: 0.1).zeropowerViaNewtonSchulz5(q, steps: 5) + expectSummary( + result, + ArraySummary( + shape: [4, 4], + dtype: .float32, + mean: 0.11143779754638672, + minimum: -0.5158382654190063, + maximum: 0.572691798210144, + absoluteSum: 5.199581146240234, + positionChecksum: 2.7055320739746094, + sampleIndices: [0, 3, 6, 9, 12, 15], + samples: [ + -0.16971194744110107, 0.50611412525177, -0.3549635410308838, + 0.5651206970214844, 0.20236217975616455, 0.572691798210144, + ]), + tolerance: .float32) + } + } +} diff --git a/Tests/MLXIntegrationTests/Generated/GeneratedOptimizersTests.swift b/Tests/MLXIntegrationTests/Generated/GeneratedOptimizersTests.swift new file mode 100644 index 000000000..a64b903fb --- /dev/null +++ b/Tests/MLXIntegrationTests/Generated/GeneratedOptimizersTests.swift @@ -0,0 +1,2113 @@ +// Copyright © 2026 Apple Inc. +// +// GENERATED by tools/integration_tests -- DO NOT EDIT. +// +// Values in this file were produced by python mlx and are compared against +// the Swift API. Regenerate with: +// +// python3 tools/integration_tests/generate.py +// +// python mlx: 0.32.2.dev20260910+1f8e74e3f +// vendored mlx (Cmlx): 0.32.2 +// generator revision: 2 +// cases: 42 + +import Foundation +import MLX +import MLXNN +import Testing + +@testable import MLXOptimizers + +@Suite("generated: Optimizers") +struct GeneratedOptimizersTests { + + @Test("SGD") + func test_SGD() throws { + try withIntegrationState(seed: 52904) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let bias = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let biasTarget = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = SGD(learningRate: 0.1) + var parameters = ModuleParameters.unflattened([("weight", weight), ("bias", bias)]) + for _ in 0 ..< 3 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)), + ("bias", 2 * (parameters[unwrapping: "bias"]! - biasTarget)), + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.1919216811656952, + minimum: -1.3976359367370605, + maximum: 1.300575613975525, + absoluteSum: 8.8001127243042, + positionChecksum: 4.43962033589681, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.10629726201295853, -0.980926513671875, -1.0590994358062744, + -0.5857878923416138, -0.16218313574790955, -0.2747551202774048, + ]), + tolerance: .float32) + expectSummary( + parameters[unwrapping: "bias"]!, + ArraySummary( + shape: [5], + dtype: .float32, + mean: -0.2919028103351593, + minimum: -0.7022152543067932, + maximum: 0.2861098051071167, + absoluteSum: 2.031733512878418, + positionChecksum: 0.9501703262329102, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + -0.5892770290374756, -0.7022152543067932, 0.2861098051071167, + -0.3718429207801819, -0.08228857815265656, + ]), + tolerance: .float32) + } + } + + @Test("SGD/momentum") + func test_SGD_momentum() throws { + try withIntegrationState(seed: 68261) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let bias = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let biasTarget = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = SGD(learningRate: 0.1, momentum: 0.9) + var parameters = ModuleParameters.unflattened([("weight", weight), ("bias", bias)]) + for _ in 0 ..< 3 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)), + ("bias", 2 * (parameters[unwrapping: "bias"]! - biasTarget)), + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.20715034008026123, + minimum: -2.3862385749816895, + maximum: 1.9399197101593018, + absoluteSum: 13.50024700164795, + positionChecksum: 7.357227325439453, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 1.9399197101593018, -1.117470622062683, -0.3328934907913208, + 1.0880004167556763, 0.503385066986084, 0.4228379726409912, + ]), + tolerance: .float32) + expectSummary( + parameters[unwrapping: "bias"]!, + ArraySummary( + shape: [5], + dtype: .float32, + mean: -0.07929762452840805, + minimum: -1.334601640701294, + maximum: 1.1424953937530518, + absoluteSum: 3.4348740577697754, + positionChecksum: 2.365402412414551, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + 0.3766975402832031, -0.07365846633911133, 1.1424953937530518, + -1.334601640701294, -0.5074209570884705, + ]), + tolerance: .float32) + } + } + + @Test("SGD/dampening") + func test_SGD_dampening() throws { + try withIntegrationState(seed: 79936) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let bias = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let biasTarget = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = SGD(learningRate: 0.1, momentum: 0.9, dampening: 0.1) + var parameters = ModuleParameters.unflattened([("weight", weight), ("bias", bias)]) + for _ in 0 ..< 3 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)), + ("bias", 2 * (parameters[unwrapping: "bias"]! - biasTarget)), + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.12276717275381088, + minimum: -1.7722910642623901, + maximum: 2.0786659717559814, + absoluteSum: 11.036983489990234, + positionChecksum: 5.762930552164714, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.274435818195343, 1.7900145053863525, -1.7722910642623901, + -1.1252398490905762, -1.0930812358856201, -0.4445021152496338, + ]), + tolerance: .float32) + expectSummary( + parameters[unwrapping: "bias"]!, + ArraySummary( + shape: [5], + dtype: .float32, + mean: -0.7147542834281921, + minimum: -0.8840757012367249, + maximum: -0.5355558395385742, + absoluteSum: 3.5737712383270264, + positionChecksum: 2.2165500640869142, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + -0.6281787157058716, -0.6859133839607239, -0.8400475978851318, + -0.5355558395385742, -0.8840757012367249, + ]), + tolerance: .float32) + } + } + + @Test("SGD/weightDecay") + func test_SGD_weightDecay() throws { + try withIntegrationState(seed: 54772) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let bias = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let biasTarget = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = SGD(learningRate: 0.1, momentum: 0.9, weightDecay: 0.1) + var parameters = ModuleParameters.unflattened([("weight", weight), ("bias", bias)]) + for _ in 0 ..< 3 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)), + ("bias", 2 * (parameters[unwrapping: "bias"]! - biasTarget)), + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.14776979386806488, + minimum: -1.3177533149719238, + maximum: 1.8211688995361328, + absoluteSum: 9.807472229003906, + positionChecksum: 5.342832565307617, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.6273982524871826, -0.31063804030418396, -0.7584969401359558, + -0.4017985761165619, -0.48426440358161926, 1.0547175407409668, + ]), + tolerance: .float32) + expectSummary( + parameters[unwrapping: "bias"]!, + ArraySummary( + shape: [5], + dtype: .float32, + mean: 0.48087412118911743, + minimum: -0.46386152505874634, + maximum: 2.3275372982025146, + absoluteSum: 3.509000539779663, + positionChecksum: 2.3551183700561524, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + -0.46386152505874634, 0.32813137769699097, 0.3010169267654419, + 2.3275372982025146, -0.08845347166061401, + ]), + tolerance: .float32) + } + } + + @Test("SGD/nesterov") + func test_SGD_nesterov() throws { + try withIntegrationState(seed: 34485) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let bias = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let biasTarget = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = SGD(learningRate: 0.1, momentum: 0.9, nesterov: true) + var parameters = ModuleParameters.unflattened([("weight", weight), ("bias", bias)]) + for _ in 0 ..< 3 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)), + ("bias", 2 * (parameters[unwrapping: "bias"]! - biasTarget)), + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.5082465410232544, + minimum: -1.6663074493408203, + maximum: 1.9156103134155273, + absoluteSum: 11.580924987792969, + positionChecksum: 6.680131276448567, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.5632967352867126, 0.5782409310340881, 1.9156103134155273, + -1.0746759176254272, 0.9145888686180115, -1.6663074493408203, + ]), + tolerance: .float32) + expectSummary( + parameters[unwrapping: "bias"]!, + ArraySummary( + shape: [5], + dtype: .float32, + mean: 0.14118611812591553, + minimum: -0.5578364133834839, + maximum: 1.4228016138076782, + absoluteSum: 2.7750954627990723, + positionChecksum: 1.6474948883056642, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + -0.476746141910553, 0.2838525176048279, 1.4228016138076782, + 0.03385895490646362, -0.5578364133834839, + ]), + tolerance: .float32) + } + } + + @Test("RMSprop") + func test_RMSprop() throws { + try withIntegrationState(seed: 83905) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let bias = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let biasTarget = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = RMSprop(learningRate: 0.1) + var parameters = ModuleParameters.unflattened([("weight", weight), ("bias", bias)]) + for _ in 0 ..< 3 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)), + ("bias", 2 * (parameters[unwrapping: "bias"]! - biasTarget)), + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.2316233515739441, + minimum: -1.397747278213501, + maximum: 1.4456602334976196, + absoluteSum: 6.713398456573486, + positionChecksum: 4.262948671976726, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.13057762384414673, 0.1284933090209961, 0.578632652759552, + 0.40558111667633057, 1.171437382698059, -0.2595949172973633, + ]), + tolerance: .float32) + expectSummary( + parameters[unwrapping: "bias"]!, + ArraySummary( + shape: [5], + dtype: .float32, + mean: -0.2978183627128601, + minimum: -0.9918203353881836, + maximum: 1.4725697040557861, + absoluteSum: 4.434231281280518, + positionChecksum: 2.373613166809082, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + -0.9918203353881836, 1.4725697040557861, -0.9049973487854004, + -0.10810491442680359, -0.9567388296127319, + ]), + tolerance: .float32) + } + } + + @Test("RMSprop/alpha") + func test_RMSprop_alpha() throws { + try withIntegrationState(seed: 78027) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let bias = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let biasTarget = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = RMSprop(learningRate: 0.1, alpha: 0.5, eps: 1e-6) + var parameters = ModuleParameters.unflattened([("weight", weight), ("bias", bias)]) + for _ in 0 ..< 3 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)), + ("bias", 2 * (parameters[unwrapping: "bias"]! - biasTarget)), + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.48947882652282715, + minimum: -2.4106788635253906, + maximum: 1.5219806432724, + absoluteSum: 11.106720924377441, + positionChecksum: 4.809648513793945, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -2.4106788635253906, 0.7164310812950134, -1.4400393962860107, + 0.128189355134964, 1.5219806432724, -0.10650129616260529, + ]), + tolerance: .float32) + expectSummary( + parameters[unwrapping: "bias"]!, + ArraySummary( + shape: [5], + dtype: .float32, + mean: 0.14489522576332092, + minimum: -0.25065889954566956, + maximum: 0.5770130753517151, + absoluteSum: 1.6556437015533447, + positionChecksum: 1.1530281066894532, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + -0.2149249017238617, 0.08289781957864761, 0.5770130753517151, + -0.25065889954566956, 0.530148983001709, + ]), + tolerance: .float32) + } + } + + @Test("AdaGrad") + func test_AdaGrad() throws { + try withIntegrationState(seed: 47873) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let bias = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let biasTarget = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = AdaGrad(learningRate: 0.1) + var parameters = ModuleParameters.unflattened([("weight", weight), ("bias", bias)]) + for _ in 0 ..< 3 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)), + ("bias", 2 * (parameters[unwrapping: "bias"]! - biasTarget)), + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.5292344093322754, + minimum: -0.2918749749660492, + maximum: 1.540786623954773, + absoluteSum: 7.637607574462891, + positionChecksum: 4.699300448099772, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.7778294086456299, -0.2918749749660492, 0.20912456512451172, + -0.27893269062042236, 0.9240051507949829, 0.8926200270652771, + ]), + tolerance: .float32) + expectSummary( + parameters[unwrapping: "bias"]!, + ArraySummary( + shape: [5], + dtype: .float32, + mean: -0.32091256976127625, + minimum: -2.3398289680480957, + maximum: 1.4234572649002075, + absoluteSum: 4.687746047973633, + positionChecksum: 3.28100700378418, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + 1.4234572649002075, 0.1181345283985138, -0.17913809418678284, + -0.6271874904632568, -2.3398289680480957, + ]), + tolerance: .float32) + } + } + + @Test("AdaDelta") + func test_AdaDelta() throws { + try withIntegrationState(seed: 4025) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let bias = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let biasTarget = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = AdaDelta(learningRate: 0.1) + var parameters = ModuleParameters.unflattened([("weight", weight), ("bias", bias)]) + for _ in 0 ..< 3 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)), + ("bias", 2 * (parameters[unwrapping: "bias"]! - biasTarget)), + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.22225290536880493, + minimum: -1.8476370573043823, + maximum: 1.6021405458450317, + absoluteSum: 11.606254577636719, + positionChecksum: 5.977849960327148, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -1.8476370573043823, -1.1789054870605469, 1.6021405458450317, + 0.3608337938785553, 0.5647247433662415, -1.187432050704956, + ]), + tolerance: .float32) + expectSummary( + parameters[unwrapping: "bias"]!, + ArraySummary( + shape: [5], + dtype: .float32, + mean: -0.5898528099060059, + minimum: -1.4875011444091797, + maximum: 0.8329288363456726, + absoluteSum: 4.615121841430664, + positionChecksum: 3.0648353576660154, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + -0.6340699791908264, -0.4064069092273712, -1.254214882850647, + -1.4875011444091797, 0.8329288363456726, + ]), + tolerance: .float32) + } + } + + @Test("AdaDelta/rho") + func test_AdaDelta_rho() throws { + try withIntegrationState(seed: 56156) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let bias = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let biasTarget = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = AdaDelta(learningRate: 0.1, rho: 0.5, eps: 1e-5) + var parameters = ModuleParameters.unflattened([("weight", weight), ("bias", bias)]) + for _ in 0 ..< 3 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)), + ("bias", 2 * (parameters[unwrapping: "bias"]! - biasTarget)), + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.6433796882629395, + minimum: -1.97148597240448, + maximum: 3.078878402709961, + absoluteSum: 11.844547271728516, + positionChecksum: 6.222557703653972, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.07173069566488266, 0.7374876737594604, 0.23356150090694427, + -0.018778571859002113, 0.8372526168823242, 0.6959485411643982, + ]), + tolerance: .float32) + expectSummary( + parameters[unwrapping: "bias"]!, + ArraySummary( + shape: [5], + dtype: .float32, + mean: 0.2940554618835449, + minimum: -0.3743702173233032, + maximum: 0.8940013647079468, + absoluteSum: 2.21901798248291, + positionChecksum: 1.5835657119750977, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + 0.320523202419281, -0.3743702173233032, 0.14193440973758698, + 0.48818862438201904, 0.8940013647079468, + ]), + tolerance: .float32) + } + } + + @Test("Adam") + func test_Adam() throws { + try withIntegrationState(seed: 2855) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let bias = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let biasTarget = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = Adam(learningRate: 0.1) + var parameters = ModuleParameters.unflattened([("weight", weight), ("bias", bias)]) + for _ in 0 ..< 3 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)), + ("bias", 2 * (parameters[unwrapping: "bias"]! - biasTarget)), + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.20924310386180878, + minimum: -1.4468461275100708, + maximum: 1.958080768585205, + absoluteSum: 11.411166191101074, + positionChecksum: 5.717811584472656, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 1.958080768585205, -0.2732287645339966, 1.9326785802841187, + 0.7665133476257324, -1.0778547525405884, -0.7260398268699646, + ]), + tolerance: .float32) + expectSummary( + parameters[unwrapping: "bias"]!, + ArraySummary( + shape: [5], + dtype: .float32, + mean: 0.43796560168266296, + minimum: -0.16814911365509033, + maximum: 0.8766471147537231, + absoluteSum: 2.5261263847351074, + positionChecksum: 1.248355770111084, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + 0.8766471147537231, 0.6090426445007324, 0.18284839391708374, + 0.689439058303833, -0.16814911365509033, + ]), + tolerance: .float32) + } + } + + @Test("Adam/betas") + func test_Adam_betas() throws { + try withIntegrationState(seed: 26177) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let bias = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let biasTarget = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = Adam(learningRate: 0.1, betas: (0.8, 0.9)) + var parameters = ModuleParameters.unflattened([("weight", weight), ("bias", bias)]) + for _ in 0 ..< 3 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)), + ("bias", 2 * (parameters[unwrapping: "bias"]! - biasTarget)), + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.13137023150920868, + minimum: -2.1283836364746094, + maximum: 1.7582262754440308, + absoluteSum: 9.839652061462402, + positionChecksum: 5.175012588500977, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -1.0437983274459839, 0.7374716997146606, -2.1283836364746094, + 0.43236860632896423, 1.7582262754440308, 0.6263670325279236, + ]), + tolerance: .float32) + expectSummary( + parameters[unwrapping: "bias"]!, + ArraySummary( + shape: [5], + dtype: .float32, + mean: -0.47237634658813477, + minimum: -2.5258002281188965, + maximum: 0.687131404876709, + absoluteSum: 3.736144542694092, + positionChecksum: 3.006717872619629, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + 0.687131404876709, -0.10542242228984833, -0.1645507514476776, + -0.25323984026908875, -2.5258002281188965, + ]), + tolerance: .float32) + } + } + + @Test("Adam/biasCorrection") + func test_Adam_biasCorrection() throws { + // bias correction is step dependent: identical at step 1, not later + try withIntegrationState(seed: 75773) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let bias = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let biasTarget = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = Adam(learningRate: 0.1, biasCorrection: true) + var parameters = ModuleParameters.unflattened([("weight", weight), ("bias", bias)]) + for _ in 0 ..< 3 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)), + ("bias", 2 * (parameters[unwrapping: "bias"]! - biasTarget)), + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.38266444206237793, + minimum: -0.3790029287338257, + maximum: 1.114660620689392, + absoluteSum: 5.856780052185059, + positionChecksum: 3.123239199320475, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.1374974548816681, 0.47027307748794556, 0.7321585416793823, + 1.114660620689392, -0.11590324342250824, 0.2759057581424713, + ]), + tolerance: .float32) + expectSummary( + parameters[unwrapping: "bias"]!, + ArraySummary( + shape: [5], + dtype: .float32, + mean: -0.7057530879974365, + minimum: -1.7636712789535522, + maximum: 1.1074949502944946, + absoluteSum: 5.743755340576172, + positionChecksum: 2.8100845336914064, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + -1.7636712789535522, 1.1074949502944946, -1.479508638381958, + -1.3321667909622192, -0.060913801193237305, + ]), + tolerance: .float32) + } + } + + @Test("AdamW") + func test_AdamW() throws { + try withIntegrationState(seed: 75439) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let bias = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let biasTarget = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = AdamW(learningRate: 0.1) + var parameters = ModuleParameters.unflattened([("weight", weight), ("bias", bias)]) + for _ in 0 ..< 3 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)), + ("bias", 2 * (parameters[unwrapping: "bias"]! - biasTarget)), + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.5645228028297424, + minimum: -0.8113094568252563, + maximum: 1.9090187549591064, + absoluteSum: 9.240545272827148, + positionChecksum: 4.90969975789388, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 1.8444397449493408, 0.6620670557022095, 0.06166234612464905, + -0.19734594225883484, -0.22448092699050903, -0.8113094568252563, + ]), + tolerance: .float32) + expectSummary( + parameters[unwrapping: "bias"]!, + ArraySummary( + shape: [5], + dtype: .float32, + mean: 0.11740782111883163, + minimum: -1.336305022239685, + maximum: 1.2092561721801758, + absoluteSum: 4.171524524688721, + positionChecksum: 2.539629364013672, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + -1.336305022239685, -0.45593777298927307, 0.11859375238418579, + 1.2092561721801758, 1.0514320135116577, + ]), + tolerance: .float32) + } + } + + @Test("AdamW/weightDecay") + func test_AdamW_weightDecay() throws { + try withIntegrationState(seed: 45988) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let bias = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let biasTarget = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = AdamW(learningRate: 0.1, weightDecay: 0.1, biasCorrection: true) + var parameters = ModuleParameters.unflattened([("weight", weight), ("bias", bias)]) + for _ in 0 ..< 3 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)), + ("bias", 2 * (parameters[unwrapping: "bias"]! - biasTarget)), + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.020839249715209007, + minimum: -0.9159260392189026, + maximum: 0.964216411113739, + absoluteSum: 5.242161750793457, + positionChecksum: 2.8174708684285483, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.24218687415122986, -0.09974238276481628, -0.7899527549743652, + 0.964216411113739, 0.6594827175140381, 0.3614733815193176, + ]), + tolerance: .float32) + expectSummary( + parameters[unwrapping: "bias"]!, + ArraySummary( + shape: [5], + dtype: .float32, + mean: 0.12586472928524017, + minimum: -0.33768603205680847, + maximum: 0.8633996248245239, + absoluteSum: 2.1387057304382324, + positionChecksum: 1.5357303619384766, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + -0.33768603205680847, -0.10294844210147858, 0.5206149816513062, + -0.3140564560890198, 0.8633996248245239, + ]), + tolerance: .float32) + } + } + + @Test("Adamax") + func test_Adamax() throws { + try withIntegrationState(seed: 15957) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let bias = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let biasTarget = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = Adamax(learningRate: 0.1) + var parameters = ModuleParameters.unflattened([("weight", weight), ("bias", bias)]) + for _ in 0 ..< 3 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)), + ("bias", 2 * (parameters[unwrapping: "bias"]! - biasTarget)), + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.09984465688467026, + minimum: -0.9413069486618042, + maximum: 1.3880528211593628, + absoluteSum: 9.504159927368164, + positionChecksum: 4.476319630940755, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 1.34331476688385, -0.785368800163269, -0.5558609962463379, + 0.9454747438430786, 0.5932865738868713, -0.9413069486618042, + ]), + tolerance: .float32) + expectSummary( + parameters[unwrapping: "bias"]!, + ArraySummary( + shape: [5], + dtype: .float32, + mean: 0.4311213493347168, + minimum: -0.5994178056716919, + maximum: 1.8316594362258911, + absoluteSum: 3.5902762413024902, + positionChecksum: 2.48956184387207, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + 0.6177985668182373, -0.11791691929101944, 0.4234835207462311, + 1.8316594362258911, -0.5994178056716919, + ]), + tolerance: .float32) + } + } + + @Test("Lion") + func test_Lion() throws { + try withIntegrationState(seed: 36413) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let bias = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let biasTarget = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = Lion(learningRate: 0.1) + var parameters = ModuleParameters.unflattened([("weight", weight), ("bias", bias)]) + for _ in 0 ..< 3 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)), + ("bias", 2 * (parameters[unwrapping: "bias"]! - biasTarget)), + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.13024353981018066, + minimum: -1.2807341814041138, + maximum: 0.7626199722290039, + absoluteSum: 6.7521820068359375, + positionChecksum: 3.5021870930989585, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.7626199722290039, -0.6869360208511353, 0.7073222398757935, + 0.5880645513534546, 0.3303239941596985, -0.37053045630455017, + ]), + tolerance: .float32) + expectSummary( + parameters[unwrapping: "bias"]!, + ArraySummary( + shape: [5], + dtype: .float32, + mean: 0.6855365633964539, + minimum: -0.7185075879096985, + maximum: 2.3334498405456543, + absoluteSum: 4.864697456359863, + positionChecksum: 2.822071075439453, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + 0.7441189289093018, -0.7185075879096985, 2.3334498405456543, + 0.4142351746559143, 0.6543862223625183, + ]), + tolerance: .float32) + } + } + + @Test("Lion/weightDecay") + func test_Lion_weightDecay() throws { + try withIntegrationState(seed: 7472) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let bias = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let biasTarget = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = Lion(learningRate: 0.1, betas: (0.8, 0.9), weightDecay: 0.1) + var parameters = ModuleParameters.unflattened([("weight", weight), ("bias", bias)]) + for _ in 0 ..< 3 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)), + ("bias", 2 * (parameters[unwrapping: "bias"]! - biasTarget)), + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.3012966513633728, + minimum: -1.460852026939392, + maximum: 1.3098098039627075, + absoluteSum: 6.7537384033203125, + positionChecksum: 3.571783701578776, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.10823702067136765, 0.02829088270664215, 0.49554669857025146, + 0.34063977003097534, 0.1592463254928589, 0.7434209585189819, + ]), + tolerance: .float32) + expectSummary( + parameters[unwrapping: "bias"]!, + ArraySummary( + shape: [5], + dtype: .float32, + mean: -0.19140489399433136, + minimum: -1.1581501960754395, + maximum: 0.6329450011253357, + absoluteSum: 3.3792479038238525, + positionChecksum: 1.5446331024169921, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + -1.1581501960754395, -0.917064905166626, 0.5781667828559875, + 0.6329450011253357, -0.09292108565568924, + ]), + tolerance: .float32) + } + } + + @Test("Adafactor") + func test_Adafactor() throws { + // the 2-D parameter is factored and the 1-D one is not + try withIntegrationState(seed: 1795) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let bias = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let biasTarget = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = Adafactor(learningRate: 0.1, relativeStep: false) + var parameters = ModuleParameters.unflattened([("weight", weight), ("bias", bias)]) + for _ in 0 ..< 3 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)), + ("bias", 2 * (parameters[unwrapping: "bias"]! - biasTarget)), + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.25985509157180786, + minimum: -1.102706789970398, + maximum: 1.0535223484039307, + absoluteSum: 6.386711120605469, + positionChecksum: 3.502171516418457, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.13309255242347717, 0.447610080242157, -1.102706789970398, + -0.7704825401306152, -0.013223405927419662, -0.8421300053596497, + ]), + tolerance: .float32) + expectSummary( + parameters[unwrapping: "bias"]!, + ArraySummary( + shape: [5], + dtype: .float32, + mean: -0.07308216392993927, + minimum: -1.8284789323806763, + maximum: 1.3594447374343872, + absoluteSum: 4.476809978485107, + positionChecksum: 3.4291236877441404, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + 0.1359250843524933, 0.5603299140930176, -0.5926315784454346, + -1.8284789323806763, 1.3594447374343872, + ]), + tolerance: .float32) + } + } + + @Test("Adafactor/relativeStep") + func test_Adafactor_relativeStep() throws { + // without a learning rate the step size comes from the step count + try withIntegrationState(seed: 77705) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let bias = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let biasTarget = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = Adafactor() + var parameters = ModuleParameters.unflattened([("weight", weight), ("bias", bias)]) + for _ in 0 ..< 3 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)), + ("bias", 2 * (parameters[unwrapping: "bias"]! - biasTarget)), + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.11738912761211395, + minimum: -0.6956945657730103, + maximum: 1.4269239902496338, + absoluteSum: 6.289972305297852, + positionChecksum: 4.015622456868489, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.15101845562458038, -0.4634033441543579, -0.4472257196903229, + -0.2334236353635788, 1.4269239902496338, -0.6956945657730103, + ]), + tolerance: .float32) + expectSummary( + parameters[unwrapping: "bias"]!, + ArraySummary( + shape: [5], + dtype: .float32, + mean: -0.5099145174026489, + minimum: -2.1994707584381104, + maximum: 1.0201630592346191, + absoluteSum: 4.589898586273193, + positionChecksum: 2.507386016845703, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + -0.40986695885658264, -2.1994707584381104, -0.5772594213485718, + 1.0201630592346191, -0.3831382393836975, + ]), + tolerance: .float32) + } + } + + @Test("Adafactor/beta1") + func test_Adafactor_beta1() throws { + try withIntegrationState(seed: 48053) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let bias = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let biasTarget = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = Adafactor(learningRate: 0.1, beta1: 0.9, relativeStep: false) + var parameters = ModuleParameters.unflattened([("weight", weight), ("bias", bias)]) + for _ in 0 ..< 3 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)), + ("bias", 2 * (parameters[unwrapping: "bias"]! - biasTarget)), + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.16388410329818726, + minimum: -0.9698700308799744, + maximum: 1.4167040586471558, + absoluteSum: 8.533632278442383, + positionChecksum: 4.301989237467448, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.9218981862068176, 1.1756681203842163, 1.4167040586471558, + -0.41768917441368103, -0.8146899342536926, -0.12940138578414917, + ]), + tolerance: .float32) + expectSummary( + parameters[unwrapping: "bias"]!, + ArraySummary( + shape: [5], + dtype: .float32, + mean: -0.5856785774230957, + minimum: -1.6530791521072388, + maximum: 0.16239207983016968, + absoluteSum: 3.2531771659851074, + positionChecksum: 1.8036436080932616, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + -0.9667976498603821, -0.4675377607345581, 0.16239207983016968, + -1.6530791521072388, -0.0033703986555337906, + ]), + tolerance: .float32) + } + } + + @Test("Adafactor/noScaleParameter") + func test_Adafactor_noScaleParameter() throws { + try withIntegrationState(seed: 81017) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let bias = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let biasTarget = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = Adafactor( + learningRate: 0.1, weightDecay: 0.1, scaleParameter: false, relativeStep: false) + var parameters = ModuleParameters.unflattened([("weight", weight), ("bias", bias)]) + for _ in 0 ..< 3 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)), + ("bias", 2 * (parameters[unwrapping: "bias"]! - biasTarget)), + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.14465275406837463, + minimum: -1.1899194717407227, + maximum: 1.5379257202148438, + absoluteSum: 8.267154693603516, + positionChecksum: 5.085835138956706, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.026530228555202484, -0.9377098083496094, 0.4836030900478363, + 0.6172863245010376, 0.016763836145401, -1.1380316019058228, + ]), + tolerance: .float32) + expectSummary( + parameters[unwrapping: "bias"]!, + ArraySummary( + shape: [5], + dtype: .float32, + mean: 0.34424692392349243, + minimum: -0.3117382228374481, + maximum: 1.2032966613769531, + absoluteSum: 2.567047357559204, + positionChecksum: 1.3552351951599122, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + 0.42866355180740356, 1.2032966613769531, -0.11116818338632584, + 0.5121809244155884, -0.3117382228374481, + ]), + tolerance: .float32) + } + } + + @Test("Adafactor/warmupInit") + func test_Adafactor_warmupInit() throws { + try withIntegrationState(seed: 48640) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let bias = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let biasTarget = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = Adafactor(warmupInit: true) + var parameters = ModuleParameters.unflattened([("weight", weight), ("bias", bias)]) + for _ in 0 ..< 3 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)), + ("bias", 2 * (parameters[unwrapping: "bias"]! - biasTarget)), + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.3013237416744232, + minimum: -3.0539402961730957, + maximum: 1.6986980438232422, + absoluteSum: 13.95006275177002, + positionChecksum: 7.426012674967448, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 1.6986980438232422, -2.1536543369293213, -0.3216652274131775, + 0.16456238925457, 0.7074159383773804, -1.3977259397506714, + ]), + tolerance: .float32) + expectSummary( + parameters[unwrapping: "bias"]!, + ArraySummary( + shape: [5], + dtype: .float32, + mean: -0.565890908241272, + minimum: -1.6157656908035278, + maximum: 0.5499652028083801, + absoluteSum: 4.476120471954346, + positionChecksum: 2.3359392166137694, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + -1.2608189582824707, 0.5499652028083801, -1.6157656908035278, + -0.7762027978897095, 0.2733677625656128, + ]), + tolerance: .float32) + } + } + + @Test("Muon") + func test_Muon() throws { + // loose: Muon compounds the Float-vs-double hyperparameter difference (todos/muon-04) to ~1e-4; momentumZero/oneStep/newtonSchulz are the controls + try withIntegrationState(seed: 20620) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let bias = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let biasTarget = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = Muon(learningRate: 0.1) + var parameters = ModuleParameters.unflattened([("weight", weight), ("bias", bias)]) + for _ in 0 ..< 3 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)), + ("bias", 2 * (parameters[unwrapping: "bias"]! - biasTarget)), + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.03163754940032959, + minimum: -2.5139663219451904, + maximum: 2.240089178085327, + absoluteSum: 11.950052261352539, + positionChecksum: 6.561954498291016, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.03799111396074295, 2.12414288520813, -1.5080902576446533, + 2.240089178085327, -2.5139663219451904, -0.06351476162672043, + ]), + tolerance: .loose) + expectSummary( + parameters[unwrapping: "bias"]!, + ArraySummary( + shape: [5], + dtype: .float32, + mean: -0.4560411870479584, + minimum: -2.053658962249756, + maximum: 0.9637644290924072, + absoluteSum: 4.570568561553955, + positionChecksum: 2.6469955444335938, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + 0.18141686916351318, -1.2737054824829102, -2.053658962249756, + 0.9637644290924072, -0.09802297502756119, + ]), + tolerance: .loose) + } + } + + @Test("Muon/vectorOnly") + func test_Muon_vectorOnly() throws { + // rank 1: plain momentum update, no orthogonalization + try withIntegrationState(seed: 69306) { + let bias = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let biasTarget = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = Muon(learningRate: 0.1) + var parameters = ModuleParameters.unflattened([("bias", bias)]) + for _ in 0 ..< 3 { + let gradients = ModuleParameters.unflattened([ + ("bias", 2 * (parameters[unwrapping: "bias"]! - biasTarget)) + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "bias"]!, + ArraySummary( + shape: [5], + dtype: .float32, + mean: 0.20210061967372894, + minimum: -1.4918893575668335, + maximum: 2.1059176921844482, + absoluteSum: 6.641729354858398, + positionChecksum: 3.698236846923828, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + 2.1059176921844482, -0.45383501052856445, 1.720198631286621, + -1.4918893575668335, -0.8698888421058655, + ]), + tolerance: .float32) + } + } + + @Test("Muon/matrixOnly") + func test_Muon_matrixOnly() throws { + // tall matrix: the Newton-Schulz iteration transposes first + try withIntegrationState(seed: 75197) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = Muon(learningRate: 0.1) + var parameters = ModuleParameters.unflattened([("weight", weight)]) + for _ in 0 ..< 3 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)) + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.08751313388347626, + minimum: -1.5594749450683594, + maximum: 0.9060735702514648, + absoluteSum: 6.130833625793457, + positionChecksum: 3.4926980336507163, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.7616370916366577, 0.08012694865465164, 0.3566833436489105, + -1.5594749450683594, 0.5005760788917542, 0.8296913504600525, + ]), + tolerance: .loose) + } + } + + @Test("Muon/wide") + func test_Muon_wide() throws { + // wide matrix: no transpose in the iteration + try withIntegrationState(seed: 50653) { + let weight = MLXRandom.normal([3, 6], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([3, 6], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = Muon(learningRate: 0.1) + var parameters = ModuleParameters.unflattened([("weight", weight)]) + for _ in 0 ..< 3 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)) + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [3, 6], + dtype: .float32, + mean: 0.018045898526906967, + minimum: -1.1826281547546387, + maximum: 2.2200303077697754, + absoluteSum: 13.78824234008789, + positionChecksum: 5.91921149359809, + sampleIndices: [0, 3, 7, 10, 14, 17], + samples: [ + 2.2200303077697754, 1.1188658475875854, -0.8831031918525696, + -0.7170143723487854, -0.5960718989372253, -0.6579134464263916, + ]), + tolerance: .loose) + } + } + + @Test("Muon/oneStep") + func test_Muon_oneStep() throws { + // first update only, so the momentum state does not carry + try withIntegrationState(seed: 18547) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = Muon(learningRate: 0.1) + var parameters = ModuleParameters.unflattened([("weight", weight)]) + for _ in 0 ..< 1 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)) + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.04324345290660858, + minimum: -2.483652353286743, + maximum: 1.6655775308609009, + absoluteSum: 6.838813781738281, + positionChecksum: 2.857765515645345, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.38655349612236023, 0.19337239861488342, -0.26349616050720215, + 0.9585155248641968, 0.10933685302734375, -0.05299711227416992, + ]), + tolerance: .float32) + } + } + + @Test("Muon/oneNewtonSchulzStep") + func test_Muon_oneNewtonSchulzStep() throws { + // a single iteration: if this agrees and the 5 step default does not, the difference is amplification rather than a formula error + try withIntegrationState(seed: 64243) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = Muon(learningRate: 0.1, nsSteps: 1) + var parameters = ModuleParameters.unflattened([("weight", weight)]) + for _ in 0 ..< 3 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)) + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.06221400946378708, + minimum: -2.0203006267547607, + maximum: 2.590712785720825, + absoluteSum: 10.942645072937012, + positionChecksum: 6.706094741821289, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.2853350043296814, -0.027788713574409485, 1.4993284940719604, + -0.9181387424468994, -1.357496976852417, -0.14571359753608704, + ]), + tolerance: .float32) + } + } + + @Test("Muon/twoSteps") + func test_Muon_twoSteps() throws { + // with oneStep and threeSteps this shows how fast a difference grows + try withIntegrationState(seed: 45980) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = Muon(learningRate: 0.1) + var parameters = ModuleParameters.unflattened([("weight", weight)]) + for _ in 0 ..< 2 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)) + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.028408613055944443, + minimum: -0.846992552280426, + maximum: 1.4628791809082031, + absoluteSum: 7.093116760253906, + positionChecksum: 3.6778895060221353, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.5931099057197571, 1.4628791809082031, -0.8363688588142395, + 0.7522042393684387, -0.7285213470458984, 0.362490713596344, + ]), + tolerance: .loose) + } + } + + @Test("Muon/momentumZero") + func test_Muon_momentumZero() throws { + // momentum 0 removes the (1 - momentum) coefficient, which python computes in double and Swift in Float; if this agrees and the default does not, that coefficient is the seed of the difference + try withIntegrationState(seed: 18459) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = Muon(learningRate: 0.1, momentum: 0.0, nesterov: false) + var parameters = ModuleParameters.unflattened([("weight", weight)]) + for _ in 0 ..< 3 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)) + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.597805380821228, + minimum: -2.283677816390991, + maximum: 0.8412001132965088, + absoluteSum: 10.496516227722168, + positionChecksum: 5.905821482340495, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -1.0269737243652344, 0.06046082079410553, -0.9181715846061707, + -2.283677816390991, 0.8412001132965088, -0.649782121181488, + ]), + tolerance: .float32) + } + } + + @Test("Muon/noWeightDecay") + func test_Muon_noWeightDecay() throws { + // isolates the weight decay term (python defaults it to 0.01) + try withIntegrationState(seed: 15172) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let bias = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let biasTarget = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = Muon(learningRate: 0.1, weightDecay: 0.0) + var parameters = ModuleParameters.unflattened([("weight", weight), ("bias", bias)]) + for _ in 0 ..< 3 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)), + ("bias", 2 * (parameters[unwrapping: "bias"]! - biasTarget)), + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.041766803711652756, + minimum: -1.3217413425445557, + maximum: 1.0777087211608887, + absoluteSum: 8.2870512008667, + positionChecksum: 4.308595657348633, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.8723462224006653, -0.9842115640640259, -1.3217413425445557, + -0.7982552647590637, -1.110180377960205, 0.4663331210613251, + ]), + tolerance: .loose) + expectSummary( + parameters[unwrapping: "bias"]!, + ArraySummary( + shape: [5], + dtype: .float32, + mean: 0.23301398754119873, + minimum: -0.6206360459327698, + maximum: 1.185815691947937, + absoluteSum: 3.5016415119171143, + positionChecksum: 2.000499153137207, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + 0.7082110643386841, -0.6206360459327698, 1.185815691947937, + 0.43932899832725525, -0.5476498603820801, + ]), + tolerance: .loose) + } + } + + @Test("Lion/oneStep") + func test_Lion_oneStep() throws { + // Lion's update is lr * sign(c), so a mismatched beta shows up as a difference of exactly 2 * lr per element + try withIntegrationState(seed: 34768) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let bias = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let biasTarget = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = Lion(learningRate: 0.1) + var parameters = ModuleParameters.unflattened([("weight", weight), ("bias", bias)]) + for _ in 0 ..< 1 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)), + ("bias", 2 * (parameters[unwrapping: "bias"]! - biasTarget)), + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.19433651864528656, + minimum: -1.3650480508804321, + maximum: 0.5978026390075684, + absoluteSum: 5.644954681396484, + positionChecksum: 2.9201577504475913, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.2685225009918213, -0.7982114553451538, -1.3650480508804321, + -0.06641935557126999, -0.21257930994033813, 0.5896715521812439, + ]), + tolerance: .float32) + expectSummary( + parameters[unwrapping: "bias"]!, + ArraySummary( + shape: [5], + dtype: .float32, + mean: -0.07899947464466095, + minimum: -1.4680547714233398, + maximum: 0.9950604438781738, + absoluteSum: 3.8233253955841064, + positionChecksum: 2.4469171524047852, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + -0.6411066055297852, 0.9950604438781738, 0.6133296489715576, + 0.10577392578125, -1.4680547714233398, + ]), + tolerance: .float32) + } + } + + @Test("Muon/noNesterov") + func test_Muon_noNesterov() throws { + try withIntegrationState(seed: 59801) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let bias = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let biasTarget = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = Muon( + learningRate: 0.1, momentum: 0.8, weightDecay: 0.0, nesterov: false, nsSteps: 3) + var parameters = ModuleParameters.unflattened([("weight", weight), ("bias", bias)]) + for _ in 0 ..< 3 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)), + ("bias", 2 * (parameters[unwrapping: "bias"]! - biasTarget)), + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.2554815411567688, + minimum: -0.9771368503570557, + maximum: 2.7309908866882324, + absoluteSum: 8.80552864074707, + positionChecksum: 4.27720578511556, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.04743572324514389, -0.7059667706489563, 0.5713173151016235, + 0.7252044677734375, -0.2173776626586914, -0.25553134083747864, + ]), + tolerance: .loose) + expectSummary( + parameters[unwrapping: "bias"]!, + ArraySummary( + shape: [5], + dtype: .float32, + mean: -0.5632168650627136, + minimum: -1.19088876247406, + maximum: -0.19167444109916687, + absoluteSum: 2.816084146499634, + positionChecksum: 2.11552734375, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + -0.19167444109916687, -0.280126690864563, -0.35240861773490906, + -1.19088876247406, -0.80098557472229, + ]), + tolerance: .loose) + } + } + + @Test("Adam/tenSteps") + func test_Adam_tenSteps() throws { + try withIntegrationState(seed: 36115) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let bias = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let biasTarget = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = Adam(learningRate: 0.05, biasCorrection: true) + var parameters = ModuleParameters.unflattened([("weight", weight), ("bias", bias)]) + for _ in 0 ..< 10 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)), + ("bias", 2 * (parameters[unwrapping: "bias"]! - biasTarget)), + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.02142399549484253, + minimum: -0.9353051781654358, + maximum: 0.8552088737487793, + absoluteSum: 5.777521133422852, + positionChecksum: 2.9126103719075522, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.24671006202697754, 0.8552088737487793, -0.9353051781654358, + 0.46673938632011414, -0.26257777214050293, -0.3439560532569885, + ]), + tolerance: .float32) + expectSummary( + parameters[unwrapping: "bias"]!, + ArraySummary( + shape: [5], + dtype: .float32, + mean: -0.20665812492370605, + minimum: -1.7892528772354126, + maximum: 0.8451172709465027, + absoluteSum: 3.4166171550750732, + positionChecksum: 1.5547866821289062, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + -1.7892528772354126, -0.013102374039590359, 0.8451172709465027, + -0.42259863018989563, 0.3465459942817688, + ]), + tolerance: .float32) + } + } + + @Test("SGD/singleStep") + func test_SGD_singleStep() throws { + // the first step alone, for comparison with the multi-step cases + try withIntegrationState(seed: 38685) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let bias = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let biasTarget = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = SGD(learningRate: 0.1, momentum: 0.9) + var parameters = ModuleParameters.unflattened([("weight", weight), ("bias", bias)]) + for _ in 0 ..< 1 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)), + ("bias", 2 * (parameters[unwrapping: "bias"]! - biasTarget)), + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.0076949745416641235, + minimum: -0.7979379296302795, + maximum: 1.260810375213623, + absoluteSum: 6.1943678855896, + positionChecksum: 3.0758225123087564, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 1.127612590789795, -0.1752968430519104, 0.3452015519142151, + -0.7125493288040161, 0.18754345178604126, -0.3442728519439697, + ]), + tolerance: .float32) + expectSummary( + parameters[unwrapping: "bias"]!, + ArraySummary( + shape: [5], + dtype: .float32, + mean: -0.0020482302643358707, + minimum: -0.7621648907661438, + maximum: 0.3773851990699768, + absoluteSum: 1.540561318397522, + positionChecksum: 0.814434814453125, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + -0.7621648907661438, -0.013236373662948608, 0.08243922144174576, + 0.3773851990699768, 0.3053356409072876, + ]), + tolerance: .float32) + } + } + + @Test("Adam/shapes") + func test_Adam_shapes() throws { + try withIntegrationState(seed: 98800) { + let scalar = MLXRandom.normal([Int](), dtype: .float32, loc: 0.0, scale: 1.0) + let vector = MLXRandom.normal([7], dtype: .float32, loc: 0.0, scale: 1.0) + let matrix = MLXRandom.normal([3, 5], dtype: .float32, loc: 0.0, scale: 1.0) + let tensor = MLXRandom.normal([2, 3, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let scalarTarget = MLXRandom.normal([Int](), dtype: .float32, loc: 0.0, scale: 1.0) + let vectorTarget = MLXRandom.normal([7], dtype: .float32, loc: 0.0, scale: 1.0) + let matrixTarget = MLXRandom.normal([3, 5], dtype: .float32, loc: 0.0, scale: 1.0) + let tensorTarget = MLXRandom.normal([2, 3, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = Adam(learningRate: 0.1) + var parameters = ModuleParameters.unflattened([ + ("scalar", scalar), ("vector", vector), ("matrix", matrix), ("tensor", tensor), + ]) + for _ in 0 ..< 3 { + let gradients = ModuleParameters.unflattened([ + ("scalar", 2 * (parameters[unwrapping: "scalar"]! - scalarTarget)), + ("vector", 2 * (parameters[unwrapping: "vector"]! - vectorTarget)), + ("matrix", 2 * (parameters[unwrapping: "matrix"]! - matrixTarget)), + ("tensor", 2 * (parameters[unwrapping: "tensor"]! - tensorTarget)), + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "scalar"]!, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: -0.04785709083080292, + minimum: -0.04785709083080292, + maximum: -0.04785709083080292, + absoluteSum: 0.04785709083080292, + positionChecksum: 0.04785709083080292, + sampleIndices: [0], + samples: [-0.04785709083080292]), + tolerance: .float32) + expectSummary( + parameters[unwrapping: "vector"]!, + ArraySummary( + shape: [7], + dtype: .float32, + mean: -0.5064614415168762, + minimum: -2.039198875427246, + maximum: 0.8779197931289673, + absoluteSum: 6.674189567565918, + positionChecksum: 2.9789466857910156, + sampleIndices: [0, 1, 2, 4, 5, 6], + samples: [ + -1.4115877151489258, -2.039198875427246, -0.9141958951950073, + 0.8779197931289673, -0.7447269558906555, 0.33861857652664185, + ]), + tolerance: .float32) + expectSummary( + parameters[unwrapping: "matrix"]!, + ArraySummary( + shape: [3, 5], + dtype: .float32, + mean: -0.028072740882635117, + minimum: -1.7022106647491455, + maximum: 1.6618009805679321, + absoluteSum: 14.534392356872559, + positionChecksum: 7.248495483398438, + sampleIndices: [0, 3, 6, 8, 11, 14], + samples: [ + -1.1987392902374268, 0.766374945640564, -1.7022106647491455, + 0.8694322109222412, -0.6111583113670349, -0.4536212980747223, + ]), + tolerance: .float32) + expectSummary( + parameters[unwrapping: "tensor"]!, + ArraySummary( + shape: [2, 3, 4], + dtype: .float32, + mean: -0.10296154022216797, + minimum: -1.9619927406311035, + maximum: 1.3071478605270386, + absoluteSum: 13.969619750976562, + positionChecksum: 7.8157094319661455, + sampleIndices: [0, 5, 9, 14, 18, 23], + samples: [ + -0.4101582169532776, 1.3071478605270386, 0.08434449136257172, + 0.14656518399715424, -1.9619927406311035, -1.6114184856414795, + ]), + tolerance: .float32) + } + } + + @Test("Adafactor/shapes") + func test_Adafactor_shapes() throws { + try withIntegrationState(seed: 41330) { + let vector = MLXRandom.normal([7], dtype: .float32, loc: 0.0, scale: 1.0) + let matrix = MLXRandom.normal([3, 5], dtype: .float32, loc: 0.0, scale: 1.0) + let tensor = MLXRandom.normal([2, 3, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let vectorTarget = MLXRandom.normal([7], dtype: .float32, loc: 0.0, scale: 1.0) + let matrixTarget = MLXRandom.normal([3, 5], dtype: .float32, loc: 0.0, scale: 1.0) + let tensorTarget = MLXRandom.normal([2, 3, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let optimizer = Adafactor(learningRate: 0.1, relativeStep: false) + var parameters = ModuleParameters.unflattened([ + ("vector", vector), ("matrix", matrix), ("tensor", tensor), + ]) + for _ in 0 ..< 3 { + let gradients = ModuleParameters.unflattened([ + ("vector", 2 * (parameters[unwrapping: "vector"]! - vectorTarget)), + ("matrix", 2 * (parameters[unwrapping: "matrix"]! - matrixTarget)), + ("tensor", 2 * (parameters[unwrapping: "tensor"]! - tensorTarget)), + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "vector"]!, + ArraySummary( + shape: [7], + dtype: .float32, + mean: 0.24444937705993652, + minimum: -0.7275797724723816, + maximum: 1.3433971405029297, + absoluteSum: 3.5701165199279785, + positionChecksum: 2.086233820234026, + sampleIndices: [0, 1, 2, 4, 5, 6], + samples: [ + 0.7943839430809021, 0.034648530185222626, -0.20190563797950745, + 0.14161652326583862, 0.3265848159790039, -0.7275797724723816, + ]), + tolerance: .float32) + expectSummary( + parameters[unwrapping: "matrix"]!, + ArraySummary( + shape: [3, 5], + dtype: .float32, + mean: -0.022888900712132454, + minimum: -1.4636974334716797, + maximum: 1.163638949394226, + absoluteSum: 8.799198150634766, + positionChecksum: 4.993682861328125, + sampleIndices: [0, 3, 6, 8, 11, 14], + samples: [ + 0.019102375954389572, 0.8244732618331909, -0.2120092213153839, + 0.47228384017944336, 0.7905675172805786, -0.47224366664886475, + ]), + tolerance: .float32) + expectSummary( + parameters[unwrapping: "tensor"]!, + ArraySummary( + shape: [2, 3, 4], + dtype: .float32, + mean: 0.1986042559146881, + minimum: -0.7255576252937317, + maximum: 1.7102957963943481, + absoluteSum: 12.307641983032227, + positionChecksum: 6.778558095296224, + sampleIndices: [0, 5, 9, 14, 18, 23], + samples: [ + -0.2963857650756836, 0.9113672375679016, -0.4616476595401764, + 0.8774430155754089, 0.8562024831771851, 0.14194099605083466, + ]), + tolerance: .float32) + } + } + + @Test("SGD/positive") + func test_SGD_positive() throws { + try withIntegrationState(seed: 1980) { + let weight = MLXRandom.uniform(low: 0.5, high: 1.5, [4, 3], dtype: .float32) + let bias = MLXRandom.uniform(low: 0.5, high: 1.5, [5], dtype: .float32) + let weightTarget = MLXRandom.uniform(low: 0.5, high: 1.5, [4, 3], dtype: .float32) + let biasTarget = MLXRandom.uniform(low: 0.5, high: 1.5, [5], dtype: .float32) + let optimizer = SGD(learningRate: 0.05, momentum: 0.9, weightDecay: 0.2) + var parameters = ModuleParameters.unflattened([("weight", weight), ("bias", bias)]) + for _ in 0 ..< 3 { + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)), + ("bias", 2 * (parameters[unwrapping: "bias"]! - biasTarget)), + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 1.0045102834701538, + minimum: 0.7203009724617004, + maximum: 1.2829526662826538, + absoluteSum: 12.054122924804688, + positionChecksum: 6.408329010009766, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 1.152671456336975, 0.759495735168457, 0.9144626259803772, + 0.8906338214874268, 1.2829526662826538, 0.7203009724617004, + ]), + tolerance: .float32) + expectSummary( + parameters[unwrapping: "bias"]!, + ArraySummary( + shape: [5], + dtype: .float32, + mean: 0.9440234303474426, + minimum: 0.7767097353935242, + maximum: 1.3086878061294556, + absoluteSum: 4.720117092132568, + positionChecksum: 2.8720470428466798, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + 0.8757483959197998, 0.9107287526130676, 0.8482428789138794, + 1.3086878061294556, 0.7767097353935242, + ]), + tolerance: .float32) + } + } + + @Test("SGD/cosineDecay") + func test_SGD_cosineDecay() throws { + try withIntegrationState(seed: 70093) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let bias = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let biasTarget = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let schedule = cosineDecay(0.1, decaySteps: 10) + let optimizer = SGD(learningRate: schedule(0), momentum: 0.9) + var parameters = ModuleParameters.unflattened([("weight", weight), ("bias", bias)]) + for step in 0 ..< 5 { + optimizer.learningRate = schedule(step) + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)), + ("bias", 2 * (parameters[unwrapping: "bias"]! - biasTarget)), + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.5896511077880859, + minimum: -1.0688788890838623, + maximum: 2.913727045059204, + absoluteSum: 14.172250747680664, + positionChecksum: 6.361111958821614, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -1.0688788890838623, 1.9024821519851685, 0.9821493029594421, + -0.4725768268108368, -0.4211561679840088, 0.9433878660202026, + ]), + tolerance: .float32) + expectSummary( + parameters[unwrapping: "bias"]!, + ArraySummary( + shape: [5], + dtype: .float32, + mean: 0.006274223327636719, + minimum: -2.107053756713867, + maximum: 1.7379882335662842, + absoluteSum: 4.676799774169922, + positionChecksum: 2.9456262588500977, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + 1.7379882335662842, -0.21566040813922882, 0.44083574414253235, + 0.1752614676952362, -2.107053756713867, + ]), + tolerance: .float32) + } + } + + @Test("Adam/exponentialDecay") + func test_Adam_exponentialDecay() throws { + try withIntegrationState(seed: 60243) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let bias = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let biasTarget = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let schedule = exponentialDecay(0.1, decayRate: 0.9) + let optimizer = Adam(learningRate: schedule(0)) + var parameters = ModuleParameters.unflattened([("weight", weight), ("bias", bias)]) + for step in 0 ..< 5 { + optimizer.learningRate = schedule(step) + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)), + ("bias", 2 * (parameters[unwrapping: "bias"]! - biasTarget)), + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.3028128147125244, + minimum: -1.2761456966400146, + maximum: 0.6678332686424255, + absoluteSum: 7.528482913970947, + positionChecksum: 4.288489977518718, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -1.0958678722381592, 0.2696917951107025, -0.22582115232944489, + -1.2761456966400146, 0.34309184551239014, -1.0258384943008423, + ]), + tolerance: .float32) + expectSummary( + parameters[unwrapping: "bias"]!, + ArraySummary( + shape: [5], + dtype: .float32, + mean: 0.3606184124946594, + minimum: -0.35021162033081055, + maximum: 1.2646141052246094, + absoluteSum: 2.5035152435302734, + positionChecksum: 1.0792693138122558, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + 1.2646141052246094, 0.3743942081928253, 0.425295889377594, + 0.08899958431720734, -0.35021162033081055, + ]), + tolerance: .float32) + } + } + + @Test("SGD/stepDecay") + func test_SGD_stepDecay() throws { + try withIntegrationState(seed: 56543) { + let weight = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let bias = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let weightTarget = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let biasTarget = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let schedule = stepDecay(0.1, decayRate: 0.5, stepSize: 2) + let optimizer = SGD(learningRate: schedule(0)) + var parameters = ModuleParameters.unflattened([("weight", weight), ("bias", bias)]) + for step in 0 ..< 6 { + optimizer.learningRate = schedule(step) + let gradients = ModuleParameters.unflattened([ + ("weight", 2 * (parameters[unwrapping: "weight"]! - weightTarget)), + ("bias", 2 * (parameters[unwrapping: "bias"]! - biasTarget)), + ]) + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + } + expectSummary( + parameters[unwrapping: "weight"]!, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.45779770612716675, + minimum: -0.9003390669822693, + maximum: 2.1724798679351807, + absoluteSum: 9.183296203613281, + positionChecksum: 5.28756841023763, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.4551224708557129, 0.24549242854118347, 2.1724798679351807, + 1.0964405536651611, -0.5553304553031921, 0.46668076515197754, + ]), + tolerance: .float32) + expectSummary( + parameters[unwrapping: "bias"]!, + ArraySummary( + shape: [5], + dtype: .float32, + mean: -0.051898252218961716, + minimum: -0.9068479537963867, + maximum: 1.0226328372955322, + absoluteSum: 3.2517483234405518, + positionChecksum: 1.760629653930664, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + 1.0226328372955322, -0.9068479537963867, -0.085511215031147, + 0.47349560260772705, -0.763260543346405, + ]), + tolerance: .float32) + } + } +} diff --git a/Tests/MLXIntegrationTests/Generated/GeneratedQuantizationTests.swift b/Tests/MLXIntegrationTests/Generated/GeneratedQuantizationTests.swift new file mode 100644 index 000000000..11ce7fffa --- /dev/null +++ b/Tests/MLXIntegrationTests/Generated/GeneratedQuantizationTests.swift @@ -0,0 +1,571 @@ +// Copyright © 2026 Apple Inc. +// +// GENERATED by tools/integration_tests -- DO NOT EDIT. +// +// Values in this file were produced by python mlx and are compared against +// the Swift API. Regenerate with: +// +// python3 tools/integration_tests/generate.py +// +// python mlx: 0.32.2.dev20260910+1f8e74e3f +// vendored mlx (Cmlx): 0.32.2 +// generator revision: 2 +// cases: 22 + +import Foundation +import MLX +import Testing + +@Suite("generated: Quantization") +struct GeneratedQuantizationTests { + + @Test("quantized/bits2/wq") + func test_quantized_bits2_wq() throws { + try withIntegrationState(seed: 44380) { + let w = MLXRandom.uniform(low: -1.0, high: 1.0, [64, 128], dtype: .float32) + let result = MLX.quantized(w, groupSize: 64, bits: 2).wq + expectSummary( + result, + ArraySummary( + shape: [64, 8], + dtype: .uint32, + mean: 2706155520.0, + minimum: 55539704.0, + maximum: 4294343168.0, + absoluteSum: 1385551626240.0, + positionChecksum: 690314412032.0, + sampleIndices: [0, 102, 204, 307, 409, 511], + samples: [ + 1849313664.0, 1907358848.0, 2144858240.0, 328123136.0, 3683949312.0, + 1981644544.0, + ]), + tolerance: .exact) + } + } + + @Test("quantized/bits2/scales") + func test_quantized_bits2_scales() throws { + try withIntegrationState(seed: 66112) { + let w = MLXRandom.uniform(low: -1.0, high: 1.0, [64, 128], dtype: .float32) + let result = MLX.quantized(w, groupSize: 64, bits: 2).scales + expectSummary( + result, + ArraySummary( + shape: [64, 2], + dtype: .float32, + mean: -0.006201982498168945, + minimum: -0.49998176097869873, + maximum: 0.4999878704547882, + absoluteSum: 62.84138107299805, + positionChecksum: 31.624141693115234, + sampleIndices: [0, 25, 51, 76, 102, 127], + samples: [ + 0.4971696436405182, 0.4996141493320465, 0.49712875485420227, + 0.4998661279678345, 0.49900156259536743, -0.48680394887924194, + ]), + tolerance: .float32) + } + } + + @Test("quantized/bits2/biases") + func test_quantized_bits2_biases() throws { + try withIntegrationState(seed: 88914) { + let w = MLXRandom.uniform(low: -1.0, high: 1.0, [64, 128], dtype: .float32) + let result = MLX.quantized(w, groupSize: 64, bits: 2).biases! + expectSummary( + result, + ArraySummary( + shape: [64, 2], + dtype: .float32, + mean: -0.04578951746225357, + minimum: -0.9995979070663452, + maximum: 0.9998750686645508, + absoluteSum: 126.25718688964844, + positionChecksum: 63.589569091796875, + sampleIndices: [0, 25, 51, 76, 102, 127], + samples: [ + 0.9953732490539551, -0.9974033832550049, 0.9865076541900635, + -0.9958232045173645, -0.9863418936729431, -0.9527243375778198, + ]), + tolerance: .float32) + } + } + + @Test("dequantized/bits2") + func test_dequantized_bits2() throws { + try withIntegrationState(seed: 42894) { + let w = MLXRandom.uniform(low: -1.0, high: 1.0, [64, 128], dtype: .float32) + let q = MLX.quantized(w, groupSize: 64, bits: 2) + let result = MLX.dequantized( + q.wq, scales: q.scales, biases: q.biases, groupSize: 64, bits: 2) + expectSummary( + result, + ArraySummary( + shape: [64, 128], + dtype: .float32, + mean: 0.02056698314845562, + minimum: -0.9999780654907227, + maximum: 0.9999353885650635, + absoluteSum: 3605.364501953125, + positionChecksum: 1786.4384765625, + sampleIndices: [0, 1638, 3276, 4915, 6553, 8191], + samples: [ + -0.4797353148460388, 0.0, 0.9931421279907227, 0.49951058626174927, + -0.4997556805610657, -0.4949713349342346, + ]), + tolerance: .float32) + } + } + + @Test("quantizedMM/bits2") + func test_quantizedMM_bits2() throws { + try withIntegrationState(seed: 40462) { + let w = MLXRandom.uniform(low: -1.0, high: 1.0, [64, 128], dtype: .float32) + let x = MLXRandom.normal([8, 128], dtype: .float32, loc: 0.0, scale: 1.0) + let q = MLX.quantized(w, groupSize: 64, bits: 2) + let result = MLX.quantizedMM( + x, q.wq, scales: q.scales, biases: q.biases, groupSize: 64, bits: 2) + expectSummary( + result, + ArraySummary( + shape: [8, 64], + dtype: .float32, + mean: -0.2884669303894043, + minimum: -21.208293914794922, + maximum: 16.252973556518555, + absoluteSum: 2394.50634765625, + positionChecksum: 1203.15673828125, + sampleIndices: [0, 102, 204, 307, 409, 511], + samples: [ + 5.748183727264404, -3.221513271331787, -0.056482791900634766, + -0.08748602867126465, -2.5631000995635986, 7.229959487915039, + ]), + tolerance: .float32) + } + } + + @Test("quantized/bits4/wq") + func test_quantized_bits4_wq() throws { + try withIntegrationState(seed: 44608) { + let w = MLXRandom.uniform(low: -1.0, high: 1.0, [64, 128], dtype: .float32) + let result = MLX.quantized(w, groupSize: 64, bits: 4).wq + expectSummary( + result, + ArraySummary( + shape: [64, 16], + dtype: .uint32, + mean: 2183851520.0, + minimum: 21900200.0, + maximum: 4294455040.0, + absoluteSum: 2236263956480.0, + positionChecksum: 1127263502336.0, + sampleIndices: [0, 205, 409, 614, 818, 1023], + samples: [ + 1132585600.0, 573891328.0, 1776925312.0, 2403129856.0, 2522929152.0, + 1568619904.0, + ]), + tolerance: .exact) + } + } + + @Test("quantized/bits4/scales") + func test_quantized_bits4_scales() throws { + try withIntegrationState(seed: 46023) { + let w = MLXRandom.uniform(low: -1.0, high: 1.0, [64, 128], dtype: .float32) + let result = MLX.quantized(w, groupSize: 64, bits: 4).scales + expectSummary( + result, + ArraySummary( + shape: [64, 2], + dtype: .float32, + mean: -0.003454235615208745, + minimum: -0.12493161857128143, + maximum: 0.1249997615814209, + absoluteSum: 15.703900337219238, + positionChecksum: 7.915512561798096, + sampleIndices: [0, 25, 51, 76, 102, 127], + samples: [ + -0.12353004515171051, 0.12238717824220657, -0.11830480396747589, + 0.12461955100297928, -0.1237143725156784, 0.12481806427240372, + ]), + tolerance: .float32) + } + } + + @Test("quantized/bits4/biases") + func test_quantized_bits4_biases() throws { + try withIntegrationState(seed: 34133) { + let w = MLXRandom.uniform(low: -1.0, high: 1.0, [64, 128], dtype: .float32) + let result = MLX.quantized(w, groupSize: 64, bits: 4).biases! + expectSummary( + result, + ArraySummary( + shape: [64, 2], + dtype: .float32, + mean: -0.09647571295499802, + minimum: -0.9999547600746155, + maximum: 0.9996299743652344, + absoluteSum: 126.08390808105469, + positionChecksum: 63.495304107666016, + sampleIndices: [0, 25, 51, 76, 102, 127], + samples: [ + -0.9987098574638367, 0.9933186769485474, -0.9994155168533325, + 0.9901624917984009, -0.9994679093360901, 0.9919745922088623, + ]), + tolerance: .float32) + } + } + + @Test("dequantized/bits4") + func test_dequantized_bits4() throws { + try withIntegrationState(seed: 62267) { + let w = MLXRandom.uniform(low: -1.0, high: 1.0, [64, 128], dtype: .float32) + let q = MLX.quantized(w, groupSize: 64, bits: 4) + let result = MLX.dequantized( + q.wq, scales: q.scales, biases: q.biases, groupSize: 64, bits: 4) + expectSummary( + result, + ArraySummary( + shape: [64, 128], + dtype: .float32, + mean: 0.004197360016405582, + minimum: -0.9999241828918457, + maximum: 0.9990442991256714, + absoluteSum: 4095.062255859375, + positionChecksum: 2047.1114501953125, + sampleIndices: [0, 1638, 3276, 4915, 6553, 8191], + samples: [ + 0.24011309444904327, 0.12464109063148499, 0.6011228561401367, + 0.6184362769126892, 0.12476050853729248, 0.0, + ]), + tolerance: .float32) + } + } + + @Test("quantizedMM/bits4") + func test_quantizedMM_bits4() throws { + try withIntegrationState(seed: 15163) { + let w = MLXRandom.uniform(low: -1.0, high: 1.0, [64, 128], dtype: .float32) + let x = MLXRandom.normal([8, 128], dtype: .float32, loc: 0.0, scale: 1.0) + let q = MLX.quantized(w, groupSize: 64, bits: 4) + let result = MLX.quantizedMM( + x, q.wq, scales: q.scales, biases: q.biases, groupSize: 64, bits: 4) + expectSummary( + result, + ArraySummary( + shape: [8, 64], + dtype: .float32, + mean: -0.029651537537574768, + minimum: -16.354642868041992, + maximum: 19.714792251586914, + absoluteSum: 2415.8642578125, + positionChecksum: 1229.1107177734375, + sampleIndices: [0, 102, 204, 307, 409, 511], + samples: [ + 3.3405404090881348, 1.769765019416809, -1.461564540863037, + -1.8853719234466553, -13.134235382080078, 0.8762807846069336, + ]), + tolerance: .float32) + } + } + + @Test("quantized/bits8/wq") + func test_quantized_bits8_wq() throws { + try withIntegrationState(seed: 11800) { + let w = MLXRandom.uniform(low: -1.0, high: 1.0, [64, 128], dtype: .float32) + let result = MLX.quantized(w, groupSize: 64, bits: 8).wq + expectSummary( + result, + ArraySummary( + shape: [64, 32], + dtype: .uint32, + mean: 2115291136.0, + minimum: 169010.0, + maximum: 4294641408.0, + absoluteSum: 4332116246528.0, + positionChecksum: 2162359140352.0, + sampleIndices: [0, 409, 819, 1228, 1638, 2047], + samples: [ + 3760455424.0, 1464697344.0, 4021152000.0, 3835848704.0, 36983752.0, + 3903625984.0, + ]), + tolerance: .exact) + } + } + + @Test("quantized/bits8/scales") + func test_quantized_bits8_scales() throws { + try withIntegrationState(seed: 81192) { + let w = MLXRandom.uniform(low: -1.0, high: 1.0, [64, 128], dtype: .float32) + let result = MLX.quantized(w, groupSize: 64, bits: 8).scales + expectSummary( + result, + ArraySummary( + shape: [64, 2], + dtype: .float32, + mean: -0.00018952885875478387, + minimum: -0.007808024063706398, + maximum: 0.007807521149516106, + absoluteSum: 0.9724078178405762, + positionChecksum: 0.490849107503891, + sampleIndices: [0, 25, 51, 76, 102, 127], + samples: [ + 0.007772293407469988, 0.007766560185700655, -0.007544763386249542, + -0.007312706205993891, -0.00779489241540432, 0.007582233287394047, + ]), + tolerance: .float32) + } + } + + @Test("quantized/bits8/biases") + func test_quantized_bits8_biases() throws { + try withIntegrationState(seed: 59258) { + let w = MLXRandom.uniform(low: -1.0, high: 1.0, [64, 128], dtype: .float32) + let result = MLX.quantized(w, groupSize: 64, bits: 8).biases! + expectSummary( + result, + ArraySummary( + shape: [64, 2], + dtype: .float32, + mean: 0.09264960139989853, + minimum: -0.999893069267273, + maximum: 0.9999345541000366, + absoluteSum: 125.88665771484375, + positionChecksum: 63.497344970703125, + sampleIndices: [0, 25, 51, 76, 102, 127], + samples: [ + 0.9969533681869507, 0.9536901712417603, 0.9930459260940552, + -0.9810078740119934, 0.9552137851715088, -0.996850848197937, + ]), + tolerance: .float32) + } + } + + @Test("dequantized/bits8") + func test_dequantized_bits8() throws { + try withIntegrationState(seed: 81104) { + let w = MLXRandom.uniform(low: -1.0, high: 1.0, [64, 128], dtype: .float32) + let q = MLX.quantized(w, groupSize: 64, bits: 8) + let result = MLX.dequantized( + q.wq, scales: q.scales, biases: q.biases, groupSize: 64, bits: 8) + expectSummary( + result, + ArraySummary( + shape: [64, 128], + dtype: .float32, + mean: 0.004389249254018068, + minimum: -0.9991215467453003, + maximum: 0.9998486042022705, + absoluteSum: 4120.2646484375, + positionChecksum: 2047.842041015625, + sampleIndices: [0, 1638, 3276, 4915, 6553, 8191], + samples: [ + -0.3541724383831024, -0.645488440990448, 0.8101414442062378, + 0.7884412407875061, 0.3327997326850891, 0.5721215605735779, + ]), + tolerance: .float32) + } + } + + @Test("quantizedMM/bits8") + func test_quantizedMM_bits8() throws { + try withIntegrationState(seed: 75568) { + let w = MLXRandom.uniform(low: -1.0, high: 1.0, [64, 128], dtype: .float32) + let x = MLXRandom.normal([8, 128], dtype: .float32, loc: 0.0, scale: 1.0) + let q = MLX.quantized(w, groupSize: 64, bits: 8) + let result = MLX.quantizedMM( + x, q.wq, scales: q.scales, biases: q.biases, groupSize: 64, bits: 8) + expectSummary( + result, + ArraySummary( + shape: [8, 64], + dtype: .float32, + mean: 0.16322675347328186, + minimum: -19.405149459838867, + maximum: 15.277058601379395, + absoluteSum: 2676.939697265625, + positionChecksum: 1373.523681640625, + sampleIndices: [0, 102, 204, 307, 409, 511], + samples: [ + 5.699150085449219, 4.081239223480225, -7.194596290588379, 5.197780609130859, + 1.424774169921875, 0.47141361236572266, + ]), + tolerance: .float32) + } + } + + @Test("quantized/groupSize32") + func test_quantized_groupSize32() throws { + try withIntegrationState(seed: 81977) { + let w = MLXRandom.uniform(low: -1.0, high: 1.0, [64, 128], dtype: .float32) + let result = MLX.quantized(w, groupSize: 32, bits: 4).wq + expectSummary( + result, + ArraySummary( + shape: [64, 16], + dtype: .uint32, + mean: 2162841344.0, + minimum: 1575843.0, + maximum: 4294780416.0, + absoluteSum: 2214749536256.0, + positionChecksum: 1112428249088.0, + sampleIndices: [0, 205, 409, 614, 818, 1023], + samples: [ + 1404200832.0, 4060153600.0, 2282566144.0, 787873920.0, 2414542848.0, + 3172955136.0, + ]), + tolerance: .exact) + } + } + + @Test("quantizedMM/noTranspose") + func test_quantizedMM_noTranspose() throws { + try withIntegrationState(seed: 6989) { + let w = MLXRandom.uniform(low: -1.0, high: 1.0, [128, 64], dtype: .float32) + let x = MLXRandom.normal([8, 128], dtype: .float32, loc: 0.0, scale: 1.0) + let q = MLX.quantized(w, groupSize: 64, bits: 4) + let result = MLX.quantizedMM( + x, q.wq, scales: q.scales, biases: q.biases, transpose: false, groupSize: 64, + bits: 4) + expectSummary( + result, + ArraySummary( + shape: [8, 64], + dtype: .float32, + mean: 0.04205028712749481, + minimum: -21.537073135375977, + maximum: 19.365976333618164, + absoluteSum: 2747.372314453125, + positionChecksum: 1449.293701171875, + sampleIndices: [0, 102, 204, 307, 409, 511], + samples: [ + -1.0757472515106201, 7.758847713470459, 5.602324485778809, + -0.715898871421814, -0.12151119112968445, -1.0324546098709106, + ]), + tolerance: .float32) + } + } + + @Test("gatherQuantizedMM") + func test_gatherQuantizedMM() throws { + // without indices this is quantizedMM; the indexed form needs the multi-input generator + try withIntegrationState(seed: 93943) { + let w = MLXRandom.uniform(low: -1.0, high: 1.0, [64, 128], dtype: .float32) + let x = MLXRandom.normal([8, 128], dtype: .float32, loc: 0.0, scale: 1.0) + let q = MLX.quantized(w, groupSize: 64, bits: 4) + let result = MLX.gatherQuantizedMM( + x, q.wq, scales: q.scales, biases: q.biases, groupSize: 64, bits: 4) + expectSummary( + result, + ArraySummary( + shape: [8, 64], + dtype: .float32, + mean: -0.27394434809684753, + minimum: -17.840972900390625, + maximum: 17.40959358215332, + absoluteSum: 2546.3671875, + positionChecksum: 1242.66064453125, + sampleIndices: [0, 102, 204, 307, 409, 511], + samples: [ + -3.621098041534424, -0.9873456954956055, 0.8476672172546387, + 3.559670925140381, -2.1285338401794434, 4.59012508392334, + ]), + tolerance: .float32) + } + } + + @Test("blockMaskedMM") + func test_blockMaskedMM() throws { + try withIntegrationState(seed: 73180) { + let a = MLXRandom.normal([32, 64], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([64, 32], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.blockMaskedMM(a, b, blockSize: 32) + expectSummary( + result, + ArraySummary( + shape: [32, 32], + dtype: .float32, + mean: 0.3000313639640808, + minimum: -27.04250717163086, + maximum: 25.29352569580078, + absoluteSum: 6362.05419921875, + positionChecksum: 3143.78466796875, + sampleIndices: [0, 205, 409, 614, 818, 1023], + samples: [ + -11.411656379699707, 16.519243240356445, -3.2464497089385986, + -3.709771156311035, 6.376284122467041, -3.328993558883667, + ]), + tolerance: .float32) + } + } + + @Test("gatherMM") + func test_gatherMM() throws { + try withIntegrationState(seed: 60084) { + let a = MLXRandom.normal([4, 3, 5], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([6, 5, 2], dtype: .float32, loc: 0.0, scale: 1.0) + let lhsIndices = MLXRandom.randInt(low: 0, high: 4, [2], type: UInt32.self) + let rhsIndices = MLXRandom.randInt(low: 0, high: 6, [2], type: UInt32.self) + let result = MLX.gatherMM(a, b, lhsIndices: lhsIndices, rhsIndices: rhsIndices) + expectSummary( + result, + ArraySummary( + shape: [2, 3, 2], + dtype: .float32, + mean: 1.3322674036026, + minimum: -0.8716369867324829, + maximum: 4.639623641967773, + absoluteSum: 19.419553756713867, + positionChecksum: 11.04852040608724, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 2.538201093673706, 0.6413965225219727, 1.8107318878173828, + -0.8716369867324829, 0.9088269472122192, 4.639623641967773, + ]), + tolerance: .float32) + } + } + + @Test("toFP8") + func test_toFP8() throws { + try withIntegrationState(seed: 19517) { + let a = MLXRandom.uniform(low: -4.0, high: 4.0, [4, 3], dtype: .float32) + let result = MLX.toFP8(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .uint8, + mean: 64.33333587646484, + minimum: 2.0, + maximum: 198.0, + absoluteSum: 772.0, + positionChecksum: 462.4166666666667, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [72.0, 10.0, 58.0, 56.0, 50.0, 198.0]), + tolerance: .exact) + } + } + + @Test("fromFP8") + func test_fromFP8() throws { + try withIntegrationState(seed: 87362) { + let a = MLXRandom.uniform(low: -4.0, high: 4.0, [4, 3], dtype: .float32) + let encoded = MLX.toFP8(a) + let result = MLX.fromFP8(encoded, dtype: .float32) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.1614583432674408, + minimum: -3.75, + maximum: 4.0, + absoluteSum: 26.3125, + positionChecksum: 14.755208333333334, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [-0.625, 2.75, 3.25, -1.125, -3.75, 4.0]), + tolerance: .float32) + } + } +} diff --git a/Tests/MLXIntegrationTests/Generated/GeneratedRandomInputsTests.swift b/Tests/MLXIntegrationTests/Generated/GeneratedRandomInputsTests.swift new file mode 100644 index 000000000..fa4f56c43 --- /dev/null +++ b/Tests/MLXIntegrationTests/Generated/GeneratedRandomInputsTests.swift @@ -0,0 +1,596 @@ +// Copyright © 2026 Apple Inc. +// +// GENERATED by tools/integration_tests -- DO NOT EDIT. +// +// Values in this file were produced by python mlx and are compared against +// the Swift API. Regenerate with: +// +// python3 tools/integration_tests/generate.py +// +// python mlx: 0.32.2.dev20260910+1f8e74e3f +// vendored mlx (Cmlx): 0.32.2 +// generator revision: 2 +// cases: 21 + +import Foundation +import MLX +import Testing + +@Suite("generated: RandomInputs") +struct GeneratedRandomInputsTests { + + @Test("normal/2d") + func test_normal_2d() throws { + try withIntegrationState(seed: 75481) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + expectSummary( + a, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.5905870795249939, + minimum: -1.3774175643920898, + maximum: 2.2758142948150635, + absoluteSum: 12.038883209228516, + positionChecksum: 5.29741636912028, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 2.2758142948150635, 1.4398239850997925, 0.3749207854270935, + 1.3119254112243652, -0.48573410511016846, 0.871476948261261, + ]), + tolerance: .float32) + } + } + + @Test("normal/2d/seed2") + func test_normal_2d_seed2() throws { + try withIntegrationState(seed: 89536) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + expectSummary( + a, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.13405632972717285, + minimum: -3.1497466564178467, + maximum: 1.7531814575195312, + absoluteSum: 12.751688957214355, + positionChecksum: 6.7695051829020185, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.06993529200553894, 0.6983082294464111, -3.1497466564178467, + 1.1703077554702759, 0.8777164220809937, -1.2233566045761108, + ]), + tolerance: .float32) + } + } + + @Test("normal/1d") + func test_normal_1d() throws { + try withIntegrationState(seed: 80954) { + let a = MLXRandom.normal([17], dtype: .float32, loc: 0.0, scale: 1.0) + expectSummary( + a, + ArraySummary( + shape: [17], + dtype: .float32, + mean: 0.21538208425045013, + minimum: -1.4181712865829468, + maximum: 2.0330264568328857, + absoluteSum: 14.34508991241455, + positionChecksum: 7.750569063074448, + sampleIndices: [0, 3, 6, 10, 13, 16], + samples: [ + -0.5277025103569031, 1.0268818140029907, -0.07764638960361481, + 0.72825026512146, -1.4181712865829468, 0.6145756840705872, + ]), + tolerance: .float32) + } + } + + @Test("normal/4d") + func test_normal_4d() throws { + try withIntegrationState(seed: 41791) { + let a = MLXRandom.normal([2, 3, 4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + expectSummary( + a, + ArraySummary( + shape: [2, 3, 4, 3], + dtype: .float32, + mean: -0.12563791871070862, + minimum: -2.510462522506714, + maximum: 2.0979740619659424, + absoluteSum: 60.029022216796875, + positionChecksum: 28.137715657552082, + sampleIndices: [0, 14, 28, 43, 57, 71], + samples: [ + -0.9431903958320618, -2.1896116733551025, 0.5952138900756836, + -0.40866464376449585, -0.6583279967308044, 0.3055444657802582, + ]), + tolerance: .float32) + } + } + + @Test("normal/scalar") + func test_normal_scalar() throws { + try withIntegrationState(seed: 67446) { + let a = MLXRandom.normal([Int](), dtype: .float32, loc: 0.0, scale: 1.0) + expectSummary( + a, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 0.35509586334228516, + minimum: 0.35509586334228516, + maximum: 0.35509586334228516, + absoluteSum: 0.35509586334228516, + positionChecksum: 0.35509586334228516, + sampleIndices: [0], + samples: [0.35509586334228516]), + tolerance: .float32) + } + } + + @Test("normal/locScale") + func test_normal_locScale() throws { + try withIntegrationState(seed: 74498) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 2.5, scale: 0.25) + expectSummary( + a, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 2.4573264122009277, + minimum: 2.1506776809692383, + maximum: 2.6786084175109863, + absoluteSum: 29.4879150390625, + positionChecksum: 15.964505513509115, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 2.6650614738464355, 2.2944018840789795, 2.3229405879974365, + 2.448261260986328, 2.595139741897583, 2.4941651821136475, + ]), + tolerance: .float32) + } + } + + @Test("normal/float16") + func test_normal_float16() throws { + try withIntegrationState(seed: 44715) { + let a = MLXRandom.normal([4, 3], dtype: .float16, loc: 0.0, scale: 1.0) + expectSummary( + a, + ArraySummary( + shape: [4, 3], + dtype: .float16, + mean: 0.323333740234375, + minimum: -1.005859375, + maximum: 1.4765625, + absoluteSum: 8.0045166015625, + positionChecksum: 4.085835774739583, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -1.005859375, 0.306396484375, -0.6796875, 0.25390625, 0.66552734375, + 0.1114501953125, + ]), + tolerance: .float16) + } + } + + @Test("normal/bfloat16") + func test_normal_bfloat16() throws { + try withIntegrationState(seed: 69686) { + let a = MLXRandom.normal([4, 3], dtype: .bfloat16, loc: 0.0, scale: 1.0) + expectSummary( + a, + ArraySummary( + shape: [4, 3], + dtype: .bfloat16, + mean: 0.1966959685087204, + minimum: -2.875, + maximum: 2.546875, + absoluteSum: 14.2568359375, + positionChecksum: 9.154378255208334, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [0.64453125, -1.6484375, -0.0439453125, -1.140625, 0.53125, -2.875]), + tolerance: .float16) + } + } + + @Test("uniform/unit") + func test_uniform_unit() throws { + try withIntegrationState(seed: 62798) { + let a = MLXRandom.uniform(low: 0.0, high: 1.0, [4, 3], dtype: .float32) + expectSummary( + a, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.4539196491241455, + minimum: 0.06912913173437119, + maximum: 0.986918032169342, + absoluteSum: 5.447035789489746, + positionChecksum: 2.776052474975586, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.6409809589385986, 0.14968015253543854, 0.986918032169342, + 0.7626613974571228, 0.16480962932109833, 0.13554799556732178, + ]), + tolerance: .float32) + } + } + + @Test("uniform/unit/seed2") + func test_uniform_unit_seed2() throws { + try withIntegrationState(seed: 41931) { + let a = MLXRandom.uniform(low: 0.0, high: 1.0, [4, 3], dtype: .float32) + expectSummary( + a, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.4274345636367798, + minimum: 0.010180529206991196, + maximum: 0.9993993639945984, + absoluteSum: 5.129214763641357, + positionChecksum: 2.8637450536092124, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.41757380962371826, 0.4386035203933716, 0.2980853021144867, + 0.9993993639945984, 0.8430754542350769, 0.0188215933740139, + ]), + tolerance: .float32) + } + } + + @Test("uniform/range") + func test_uniform_range() throws { + try withIntegrationState(seed: 23538) { + let a = MLXRandom.uniform(low: 0.1, high: 2.0, [4, 3], dtype: .float32) + expectSummary( + a, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 1.189245581626892, + minimum: 0.2704201340675354, + maximum: 1.9868327379226685, + absoluteSum: 14.270946502685547, + positionChecksum: 7.686832427978516, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.8097621202468872, 1.754885196685791, 0.6641294956207275, + 1.9868327379226685, 1.5307133197784424, 1.0133975744247437, + ]), + tolerance: .float32) + } + } + + @Test("uniform/negative") + func test_uniform_negative() throws { + try withIntegrationState(seed: 40772) { + let a = MLXRandom.uniform(low: -0.9, high: 0.9, [4, 3], dtype: .float32) + expectSummary( + a, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.019281506538391113, + minimum: -0.8162036538124084, + maximum: 0.7615352869033813, + absoluteSum: 5.340817451477051, + positionChecksum: 3.0421508153279624, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.31404179334640503, 0.5715023279190063, -0.10399436950683594, + -0.7621757388114929, 0.17554616928100586, 0.7615352869033813, + ]), + tolerance: .float32) + } + } + + @Test("uniform/float16") + func test_uniform_float16() throws { + try withIntegrationState(seed: 82054) { + let a = MLXRandom.uniform(low: 0.0, high: 1.0, [4, 3], dtype: .float16) + expectSummary( + a, + ArraySummary( + shape: [4, 3], + dtype: .float16, + mean: 0.43940991163253784, + minimum: 0.038726806640625, + maximum: 0.875, + absoluteSum: 5.272918701171875, + positionChecksum: 2.8272933959960938, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.23974609375, 0.08013916015625, 0.29248046875, 0.37353515625, + 0.72900390625, 0.177978515625, + ]), + tolerance: .float16) + } + } + + @Test("randInt/small") + func test_randInt_small() throws { + try withIntegrationState(seed: 21039) { + let a = MLXRandom.randInt(low: 0, high: 4, [4, 3], type: Int32.self) + expectSummary( + a, + ArraySummary( + shape: [4, 3], + dtype: .int32, + mean: 1.75, + minimum: 0.0, + maximum: 3.0, + absoluteSum: 21.0, + positionChecksum: 13.083333333333334, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [0.0, 1.0, 3.0, 3.0, 3.0, 1.0]), + tolerance: .exact) + } + } + + @Test("randInt/large") + func test_randInt_large() throws { + try withIntegrationState(seed: 4514) { + let a = MLXRandom.randInt(low: -100, high: 1024, [9], type: Int32.self) + expectSummary( + a, + ArraySummary( + shape: [9], + dtype: .int32, + mean: 518.4444580078125, + minimum: -83.0, + maximum: 998.0, + absoluteSum: 4832.0, + positionChecksum: 2413.1111111111113, + sampleIndices: [0, 2, 3, 5, 6, 8], + samples: [998.0, 463.0, 975.0, 870.0, 361.0, 140.0]), + tolerance: .exact) + } + } + + @Test("randInt/uint32") + func test_randInt_uint32() throws { + try withIntegrationState(seed: 4708) { + let a = MLXRandom.randInt(low: 0, high: 6, [4], type: UInt32.self) + expectSummary( + a, + ArraySummary( + shape: [4], + dtype: .uint32, + mean: 1.5, + minimum: 0.0, + maximum: 3.0, + absoluteSum: 6.0, + positionChecksum: 4.25, + sampleIndices: [0, 1, 2, 3], + samples: [2.0, 0.0, 1.0, 3.0]), + tolerance: .exact) + } + } + + @Test("bernoulli/half") + func test_bernoulli_half() throws { + try withIntegrationState(seed: 17534) { + let a = MLXRandom.bernoulli(0.5, [4, 3]) + expectSummary( + a, + ArraySummary( + shape: [4, 3], + dtype: .bool, + mean: 0.4166666865348816, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 5.0, + positionChecksum: 2.6666666666666665, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [1.0, 0.0, 1.0, 1.0, 0.0, 1.0]), + tolerance: .exact) + } + } + + @Test("bernoulli/quarter") + func test_bernoulli_quarter() throws { + try withIntegrationState(seed: 82624) { + let a = MLXRandom.bernoulli(0.25, [4, 3]) + expectSummary( + a, + ArraySummary( + shape: [4, 3], + dtype: .bool, + mean: 0.25, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 3.0, + positionChecksum: 1.9166666666666667, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [0.0, 0.0, 0.0, 0.0, 1.0, 0.0]), + tolerance: .exact) + } + } + + @Test("sequence/normal") + func test_sequence_normal() throws { + try withIntegrationState(seed: 91718) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + expectSummary( + a, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.6608368158340454, + minimum: -3.117097854614258, + maximum: 0.4458395838737488, + absoluteSum: 9.928923606872559, + positionChecksum: 4.727120081583659, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -3.117097854614258, -1.0726861953735352, -0.45238417387008667, + -0.7030599117279053, -0.81221604347229, -0.6846268773078918, + ]), + tolerance: .float32) + let b = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + expectSummary( + b, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.05415819212794304, + minimum: -2.0514211654663086, + maximum: 2.580822229385376, + absoluteSum: 12.921215057373047, + positionChecksum: 8.349390029907227, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.852959156036377, 0.06281861662864685, 1.4253885746002197, + 0.7632550597190857, 2.580822229385376, -1.264439344406128, + ]), + tolerance: .float32) + let c = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + expectSummary( + c, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.2736375033855438, + minimum: -1.6628878116607666, + maximum: 1.6362630128860474, + absoluteSum: 10.95200252532959, + positionChecksum: 5.341580708821614, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 1.2123521566390991, 1.0446020364761353, 0.8530498147010803, + 0.12905484437942505, -0.4385685622692108, -0.6690441370010376, + ]), + tolerance: .float32) + } + } + + @Test("sequence/mixed") + func test_sequence_mixed() throws { + try withIntegrationState(seed: 1510) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + expectSummary( + a, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.33849453926086426, + minimum: -2.251725673675537, + maximum: 1.393682599067688, + absoluteSum: 12.362427711486816, + positionChecksum: 6.745829264322917, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -1.7234569787979126, -0.23261785507202148, -0.3641354739665985, + -2.251725673675537, -0.11936888098716736, 0.1819031834602356, + ]), + tolerance: .float32) + let b = MLXRandom.uniform(low: 0.1, high: 2.0, [4, 3], dtype: .float32) + expectSummary( + b, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 1.373868703842163, + minimum: 0.1270606517791748, + maximum: 1.9089659452438354, + absoluteSum: 16.48642349243164, + positionChecksum: 10.1610959370931, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.1270606517791748, 1.5031565427780151, 1.8709282875061035, + 1.6418914794921875, 1.796919345855713, 1.742146611213684, + ]), + tolerance: .float32) + let c = MLXRandom.randInt(low: 0, high: 4, [4, 3], type: Int32.self) + expectSummary( + c, + ArraySummary( + shape: [4, 3], + dtype: .int32, + mean: 1.4166667461395264, + minimum: 0.0, + maximum: 3.0, + absoluteSum: 17.0, + positionChecksum: 10.0, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [0.0, 3.0, 1.0, 2.0, 0.0, 3.0]), + tolerance: .exact) + let d = MLXRandom.bernoulli(0.5, [4, 3]) + expectSummary( + d, + ArraySummary( + shape: [4, 3], + dtype: .bool, + mean: 0.4166666865348816, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 5.0, + positionChecksum: 2.8333333333333335, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [1.0, 0.0, 0.0, 0.0, 1.0, 1.0]), + tolerance: .exact) + } + } + + @Test("sequence/shapes") + func test_sequence_shapes() throws { + try withIntegrationState(seed: 74672) { + let a = MLXRandom.normal([10, 8], dtype: .float32, loc: 0.0, scale: 1.0) + expectSummary( + a, + ArraySummary( + shape: [10, 8], + dtype: .float32, + mean: 0.131831094622612, + minimum: -2.4943573474884033, + maximum: 1.8212450742721558, + absoluteSum: 63.549659729003906, + positionChecksum: 34.66279296875, + sampleIndices: [0, 16, 32, 47, 63, 79], + samples: [ + 0.22473129630088806, -0.6478517651557922, -0.24332140386104584, + -1.8043724298477173, 0.17549262940883636, -2.4943573474884033, + ]), + tolerance: .float32) + let b = MLXRandom.normal([8, 13], dtype: .float32, loc: 0.0, scale: 1.0) + expectSummary( + b, + ArraySummary( + shape: [8, 13], + dtype: .float32, + mean: 0.02560247667133808, + minimum: -2.513389825820923, + maximum: 2.4972832202911377, + absoluteSum: 82.86776733398438, + positionChecksum: 42.1690673828125, + sampleIndices: [0, 21, 41, 62, 82, 103], + samples: [ + -1.0536662340164185, -0.4492167830467224, -0.45552271604537964, + 1.7218302488327026, 1.0314418077468872, -1.0060495138168335, + ]), + tolerance: .float32) + let c = MLXRandom.normal([2, 3, 4], dtype: .float32, loc: 0.0, scale: 1.0) + expectSummary( + c, + ArraySummary( + shape: [2, 3, 4], + dtype: .float32, + mean: -0.17977163195610046, + minimum: -2.1054158210754395, + maximum: 1.6299241781234741, + absoluteSum: 17.145761489868164, + positionChecksum: 8.302991231282553, + sampleIndices: [0, 5, 9, 14, 18, 23], + samples: [ + -1.1952673196792603, -1.2896186113357544, 0.006688464432954788, + -0.3580016493797302, -1.3585586547851562, 0.4489697217941284, + ]), + tolerance: .float32) + } + } +} diff --git a/Tests/MLXIntegrationTests/Generated/GeneratedRandomTests.swift b/Tests/MLXIntegrationTests/Generated/GeneratedRandomTests.swift new file mode 100644 index 000000000..c12688c85 --- /dev/null +++ b/Tests/MLXIntegrationTests/Generated/GeneratedRandomTests.swift @@ -0,0 +1,269 @@ +// Copyright © 2026 Apple Inc. +// +// GENERATED by tools/integration_tests -- DO NOT EDIT. +// +// Values in this file were produced by python mlx and are compared against +// the Swift API. Regenerate with: +// +// python3 tools/integration_tests/generate.py +// +// python mlx: 0.32.2.dev20260910+1f8e74e3f +// vendored mlx (Cmlx): 0.32.2 +// generator revision: 2 +// cases: 11 + +import Foundation +import MLX +import Testing + +@Suite("generated: Random") +struct GeneratedRandomTests { + + @Test("gumbel") + func test_gumbel() throws { + try withIntegrationState(seed: 65021) { + let result = MLXRandom.gumbel([4, 3]) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.846875786781311, + minimum: -0.9293802976608276, + maximum: 3.7100210189819336, + absoluteSum: 15.268047332763672, + positionChecksum: 6.564333597819011, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 1.0602260828018188, 0.20688922703266144, 3.7100210189819336, + 1.311169147491455, 0.7526293396949768, -0.8609374761581421, + ]), + tolerance: .float32) + } + } + + @Test("gumbel/dtype") + func test_gumbel_dtype() throws { + try withIntegrationState(seed: 90296) { + let result = MLXRandom.gumbel([4, 3], dtype: .float16) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float16, + mean: 0.8638967275619507, + minimum: -1.18359375, + maximum: 5.22265625, + absoluteSum: 17.61090087890625, + positionChecksum: 8.492075602213541, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 5.22265625, 0.91845703125, -0.77490234375, 3.275390625, 1.8720703125, + -1.18359375, + ]), + tolerance: .float16) + } + } + + @Test("laplace") + func test_laplace() throws { + // the Swift dtype: parameter has no default + try withIntegrationState(seed: 62794) { + let result = MLXRandom.laplace([4, 3], dtype: .float32) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.37702733278274536, + minimum: -2.6754982471466064, + maximum: 2.45698618888855, + absoluteSum: 12.961724281311035, + positionChecksum: 8.14331309000651, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.23149043321609497, -0.45272037386894226, 0.8567035794258118, + -2.6754982471466064, 1.9753601551055908, 0.19610698521137238, + ]), + tolerance: .float32) + } + } + + @Test("laplace/locScale") + func test_laplace_locScale() throws { + try withIntegrationState(seed: 37443) { + let result = MLXRandom.laplace([4, 3], dtype: .float32, loc: 1.0, scale: 2.0) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.0226888470351696, + minimum: -4.823460578918457, + maximum: 4.867783546447754, + absoluteSum: 30.37203025817871, + positionChecksum: 17.33179982503255, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -1.4936223030090332, 0.259247362613678, 4.175117492675781, + -4.638853549957275, -1.1265568733215332, 4.867783546447754, + ]), + tolerance: .float32) + } + } + + @Test("truncatedNormal") + func test_truncatedNormal() throws { + try withIntegrationState(seed: 37638) { + let result = MLXRandom.truncatedNormal(low: -1.0, high: 1.0, [4, 3]) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.05009012669324875, + minimum: -0.8407586216926575, + maximum: 0.7136946320533752, + absoluteSum: 6.225048542022705, + positionChecksum: 3.2496751149495444, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.3011457920074463, 0.6435555815696716, -0.3013858199119568, + -0.6681084036827087, 0.2138742208480835, -0.4170948266983032, + ]), + tolerance: .float32) + } + } + + @Test("categorical") + func test_categorical() throws { + try withIntegrationState(seed: 30654) { + let logits = MLXRandom.normal([4, 5], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXRandom.categorical(logits) + expectSummary( + result, + ArraySummary( + shape: [4], + dtype: .uint32, + mean: 2.75, + minimum: 0.0, + maximum: 4.0, + absoluteSum: 11.0, + positionChecksum: 8.25, + sampleIndices: [0, 1, 2, 3], + samples: [0.0, 4.0, 3.0, 4.0]), + tolerance: .exact) + } + } + + @Test("categorical/count") + func test_categorical_count() throws { + try withIntegrationState(seed: 71811) { + let logits = MLXRandom.normal([4, 5], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXRandom.categorical(logits, count: 3) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .uint32, + mean: 1.75, + minimum: 0.0, + maximum: 3.0, + absoluteSum: 21.0, + positionChecksum: 12.333333333333334, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [0.0, 2.0, 0.0, 3.0, 0.0, 3.0]), + tolerance: .exact) + } + } + + @Test("permutation/count") + func test_permutation_count() throws { + try withIntegrationState(seed: 53709) { + let result = MLXRandom.permutation(10) + expectSummary( + result, + ArraySummary( + shape: [10], + dtype: .uint32, + mean: 4.5, + minimum: 0.0, + maximum: 9.0, + absoluteSum: 45.0, + positionChecksum: 24.7, + sampleIndices: [0, 2, 4, 5, 7, 9], + samples: [9.0, 3.0, 2.0, 6.0, 0.0, 7.0]), + tolerance: .exact) + } + } + + @Test("permutation/array") + func test_permutation_array() throws { + try withIntegrationState(seed: 9976) { + let a = MLXRandom.normal([6, 2], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLXRandom.permutation(a) + expectSummary( + result, + ArraySummary( + shape: [6, 2], + dtype: .float32, + mean: 0.30458083748817444, + minimum: -1.2631405591964722, + maximum: 3.427476167678833, + absoluteSum: 9.221382141113281, + positionChecksum: 5.262217839558919, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.9780406951904297, -1.2631405591964722, -0.6314418315887451, + 0.45676153898239136, 1.034820795059204, -0.07250352948904037, + ]), + tolerance: .float32) + } + } + + @Test("multivariateNormal") + func test_multivariateNormal() throws { + try withIntegrationState(seed: 597) { + let mean = MLX.zeros([3]) + let covariance = MLX.eye(3) * 2.0 + let result = MLXRandom.multivariateNormal( + mean: mean, covariance: covariance, shape: [4], dtype: .float32, stream: .cpu) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.37988409399986267, + minimum: -3.102618932723999, + maximum: 2.8767926692962646, + absoluteSum: 18.955251693725586, + positionChecksum: 9.738911946614584, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -1.462942361831665, 0.007976667955517769, 2.2923548221588135, + -3.102618932723999, -0.5047874450683594, 1.6479119062423706, + ]), + tolerance: .float32) + } + } + + @Test("key") + func test_key() throws { + try withIntegrationState(seed: 20155) { + let result = MLXRandom.key(42) + expectSummary( + result, + ArraySummary( + shape: [2], + dtype: .uint32, + mean: 21.0, + minimum: 0.0, + maximum: 42.0, + absoluteSum: 42.0, + positionChecksum: 42.0, + sampleIndices: [0, 1], + samples: [0.0, 42.0]), + tolerance: .exact) + } + } +} diff --git a/Tests/MLXIntegrationTests/Generated/GeneratedReductionTests.swift b/Tests/MLXIntegrationTests/Generated/GeneratedReductionTests.swift new file mode 100644 index 000000000..ea27372e7 --- /dev/null +++ b/Tests/MLXIntegrationTests/Generated/GeneratedReductionTests.swift @@ -0,0 +1,1506 @@ +// Copyright © 2026 Apple Inc. +// +// GENERATED by tools/integration_tests -- DO NOT EDIT. +// +// Values in this file were produced by python mlx and are compared against +// the Swift API. Regenerate with: +// +// python3 tools/integration_tests/generate.py +// +// python mlx: 0.32.2.dev20260910+1f8e74e3f +// vendored mlx (Cmlx): 0.32.2 +// generator revision: 2 +// cases: 64 + +import Foundation +import MLX +import Testing + +@Suite("generated: Reduction") +struct GeneratedReductionTests { + + @Test("sum") + func test_sum() throws { + try withIntegrationState(seed: 22575) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.sum(a) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 0.7057861685752869, + minimum: 0.7057861685752869, + maximum: 0.7057861685752869, + absoluteSum: 0.7057861685752869, + positionChecksum: 0.7057861685752869, + sampleIndices: [0], + samples: [0.7057861685752869]), + tolerance: .float32) + } + } + + @Test("sum/axis") + func test_sum_axis() throws { + try withIntegrationState(seed: 37581) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.sum(a, axis: -1) + expectSummary( + result, + ArraySummary( + shape: [4], + dtype: .float32, + mean: -0.5632065534591675, + minimum: -2.874988555908203, + maximum: 0.6334655284881592, + absoluteSum: 4.267636775970459, + positionChecksum: 2.3913214206695557, + sampleIndices: [0, 1, 2, 3], + samples: [ + 0.37393978238105774, -2.874988555908203, 0.6334655284881592, + -0.38524290919303894, + ]), + tolerance: .float32) + } + } + + @Test("sum/axes") + func test_sum_axes() throws { + try withIntegrationState(seed: 15457) { + let a = MLXRandom.normal([2, 3, 4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.sum(a, axes: [0, -1]) + expectSummary( + result, + ArraySummary( + shape: [3, 4], + dtype: .float32, + mean: 0.44520968198776245, + minimum: -6.6333723068237305, + maximum: 7.999377727508545, + absoluteSum: 31.566116333007812, + positionChecksum: 17.56298573811849, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.32674166560173035, 7.999377727508545, -0.22247016429901123, + -6.6333723068237305, 1.6439025402069092, -2.4121196269989014, + ]), + tolerance: .float32) + } + } + + @Test("sum/keepDims") + func test_sum_keepDims() throws { + try withIntegrationState(seed: 30341) { + let a = MLXRandom.normal([2, 3, 4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.sum(a, axes: [0, -1], keepDims: true) + expectSummary( + result, + ArraySummary( + shape: [1, 3, 4, 1], + dtype: .float32, + mean: 0.3402286171913147, + minimum: -5.441669464111328, + maximum: 6.115833282470703, + absoluteSum: 33.430877685546875, + positionChecksum: 17.420809427897137, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -3.12631893157959, 1.0557702779769897, -2.0913119316101074, + 2.757711410522461, -5.441669464111328, -0.5352753400802612, + ]), + tolerance: .float32) + } + } + + @Test("sum/method") + func test_sum_method() throws { + try withIntegrationState(seed: 92671) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = a.sum(axis: -1) + expectSummary( + result, + ArraySummary( + shape: [4], + dtype: .float32, + mean: -0.891804575920105, + minimum: -1.5811764001846313, + maximum: -0.49883174896240234, + absoluteSum: 3.56721830368042, + positionChecksum: 2.2104756832122803, + sampleIndices: [0, 1, 2, 3], + samples: [ + -0.8713735938072205, -0.6158364415168762, -1.5811764001846313, + -0.49883174896240234, + ]), + tolerance: .float32) + } + } + + @Test("mean") + func test_mean() throws { + try withIntegrationState(seed: 78073) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.mean(a) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: -0.23587274551391602, + minimum: -0.23587274551391602, + maximum: -0.23587274551391602, + absoluteSum: 0.23587274551391602, + positionChecksum: 0.23587274551391602, + sampleIndices: [0], + samples: [-0.23587274551391602]), + tolerance: .float32) + } + } + + @Test("mean/axis") + func test_mean_axis() throws { + try withIntegrationState(seed: 14185) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.mean(a, axis: -1) + expectSummary( + result, + ArraySummary( + shape: [4], + dtype: .float32, + mean: -0.2951846718788147, + minimum: -1.14504075050354, + maximum: 0.3134152293205261, + absoluteSum: 2.174551010131836, + positionChecksum: 1.1131901741027832, + sampleIndices: [0, 1, 2, 3], + samples: [ + -1.14504075050354, 0.3134152293205261, 0.18349099159240723, + -0.5326040983200073, + ]), + tolerance: .float32) + } + } + + @Test("mean/axes") + func test_mean_axes() throws { + try withIntegrationState(seed: 90885) { + let a = MLXRandom.normal([2, 3, 4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.mean(a, axes: [0, -1]) + expectSummary( + result, + ArraySummary( + shape: [3, 4], + dtype: .float32, + mean: 0.09846246242523193, + minimum: -0.46042436361312866, + maximum: 0.7632273435592651, + absoluteSum: 3.7009592056274414, + positionChecksum: 2.2039335568745932, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.17073218524456024, -0.06515144556760788, 0.4433436095714569, + -0.46042436361312866, -0.4089309573173523, 0.05317772924900055, + ]), + tolerance: .float32) + } + } + + @Test("mean/keepDims") + func test_mean_keepDims() throws { + try withIntegrationState(seed: 64348) { + let a = MLXRandom.normal([2, 3, 4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.mean(a, axes: [0, -1], keepDims: true) + expectSummary( + result, + ArraySummary( + shape: [1, 3, 4, 1], + dtype: .float32, + mean: -0.10993402451276779, + minimum: -0.8585461378097534, + maximum: 0.3680957555770874, + absoluteSum: 3.5669546127319336, + positionChecksum: 2.2769004503885903, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.138699010014534, 0.04041095823049545, -0.17920581996440887, + 0.15857860445976257, -0.15389952063560486, -0.5889500975608826, + ]), + tolerance: .float32) + } + } + + @Test("mean/method") + func test_mean_method() throws { + try withIntegrationState(seed: 64427) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = a.mean(axis: -1) + expectSummary( + result, + ArraySummary( + shape: [4], + dtype: .float32, + mean: -0.2153686285018921, + minimum: -0.5716754794120789, + maximum: 0.36203593015670776, + absoluteSum: 1.5855462551116943, + positionChecksum: 1.1746270656585693, + sampleIndices: [0, 1, 2, 3], + samples: [ + -0.13388535380363464, 0.36203593015670776, -0.5179495811462402, + -0.5716754794120789, + ]), + tolerance: .float32) + } + } + + @Test("min") + func test_min() throws { + try withIntegrationState(seed: 27506) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.min(a) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: -1.6234827041625977, + minimum: -1.6234827041625977, + maximum: -1.6234827041625977, + absoluteSum: 1.6234827041625977, + positionChecksum: 1.6234827041625977, + sampleIndices: [0], + samples: [-1.6234827041625977]), + tolerance: .float32) + } + } + + @Test("min/axis") + func test_min_axis() throws { + try withIntegrationState(seed: 17376) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.min(a, axis: -1) + expectSummary( + result, + ArraySummary( + shape: [4], + dtype: .float32, + mean: -1.527522087097168, + minimum: -3.202768564224243, + maximum: -0.5903887152671814, + absoluteSum: 6.110088348388672, + positionChecksum: 3.523585319519043, + sampleIndices: [0, 1, 2, 3], + samples: [ + -1.1166950464248657, -3.202768564224243, -0.5903887152671814, + -1.2002359628677368, + ]), + tolerance: .float32) + } + } + + @Test("min/axes") + func test_min_axes() throws { + try withIntegrationState(seed: 1356) { + let a = MLXRandom.normal([2, 3, 4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.min(a, axes: [0, -1]) + expectSummary( + result, + ArraySummary( + shape: [3, 4], + dtype: .float32, + mean: -1.2078721523284912, + minimum: -2.452378988265991, + maximum: -0.4843156933784485, + absoluteSum: 14.494464874267578, + positionChecksum: 8.333949406941732, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -1.393852710723877, -0.9224222302436829, -1.487610936164856, + -1.094190001487732, -0.6395735144615173, -2.452378988265991, + ]), + tolerance: .float32) + } + } + + @Test("min/keepDims") + func test_min_keepDims() throws { + try withIntegrationState(seed: 78536) { + let a = MLXRandom.normal([2, 3, 4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.min(a, axes: [0, -1], keepDims: true) + expectSummary( + result, + ArraySummary( + shape: [1, 3, 4, 1], + dtype: .float32, + mean: -1.075392484664917, + minimum: -2.119365692138672, + maximum: -0.35634124279022217, + absoluteSum: 12.904708862304688, + positionChecksum: 7.181203842163086, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -1.1877634525299072, -1.355299949645996, -0.35634124279022217, + -0.6733994483947754, -0.7651041150093079, -1.3274004459381104, + ]), + tolerance: .float32) + } + } + + @Test("min/method") + func test_min_method() throws { + try withIntegrationState(seed: 56994) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = a.min(axis: -1) + expectSummary( + result, + ArraySummary( + shape: [4], + dtype: .float32, + mean: -0.8903401494026184, + minimum: -2.030451536178589, + maximum: 0.40401411056518555, + absoluteSum: 4.369388580322266, + positionChecksum: 2.9881086349487305, + sampleIndices: [0, 1, 2, 3], + samples: [ + 0.40401411056518555, -2.030451536178589, -0.25217559933662415, + -1.6827476024627686, + ]), + tolerance: .float32) + } + } + + @Test("max") + func test_max() throws { + try withIntegrationState(seed: 63179) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.max(a) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 1.1173732280731201, + minimum: 1.1173732280731201, + maximum: 1.1173732280731201, + absoluteSum: 1.1173732280731201, + positionChecksum: 1.1173732280731201, + sampleIndices: [0], + samples: [1.1173732280731201]), + tolerance: .float32) + } + } + + @Test("max/axis") + func test_max_axis() throws { + try withIntegrationState(seed: 6753) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.max(a, axis: -1) + expectSummary( + result, + ArraySummary( + shape: [4], + dtype: .float32, + mean: 0.34814202785491943, + minimum: -0.42853012681007385, + maximum: 1.4546650648117065, + absoluteSum: 2.6034679412841797, + positionChecksum: 1.8354352712631226, + sampleIndices: [0, 1, 2, 3], + samples: [ + -0.17691978812217712, 0.5433529615402222, 1.4546650648117065, + -0.42853012681007385, + ]), + tolerance: .float32) + } + } + + @Test("max/axes") + func test_max_axes() throws { + try withIntegrationState(seed: 90061) { + let a = MLXRandom.normal([2, 3, 4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.max(a, axes: [0, -1]) + expectSummary( + result, + ArraySummary( + shape: [3, 4], + dtype: .float32, + mean: 1.2743048667907715, + minimum: 0.052372973412275314, + maximum: 2.9989659786224365, + absoluteSum: 15.291658401489258, + positionChecksum: 8.274401346842447, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 1.4796899557113647, 0.4484626054763794, 2.9989659786224365, + 0.052372973412275314, 1.8464957475662231, 0.9091650247573853, + ]), + tolerance: .float32) + } + } + + @Test("max/keepDims") + func test_max_keepDims() throws { + try withIntegrationState(seed: 80866) { + let a = MLXRandom.normal([2, 3, 4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.max(a, axes: [0, -1], keepDims: true) + expectSummary( + result, + ArraySummary( + shape: [1, 3, 4, 1], + dtype: .float32, + mean: 1.187777042388916, + minimum: 0.4238319396972656, + maximum: 2.3939080238342285, + absoluteSum: 14.253324508666992, + positionChecksum: 7.477830251057942, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 1.456592321395874, 0.9078659415245056, 0.4238319396972656, + 0.6892192959785461, 0.9878159761428833, 0.4668642282485962, + ]), + tolerance: .float32) + } + } + + @Test("max/method") + func test_max_method() throws { + try withIntegrationState(seed: 52822) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = a.max(axis: -1) + expectSummary( + result, + ArraySummary( + shape: [4], + dtype: .float32, + mean: 0.5650159120559692, + minimum: 0.33188390731811523, + maximum: 1.0610710382461548, + absoluteSum: 2.260063648223877, + positionChecksum: 1.1534028053283691, + sampleIndices: [0, 1, 2, 3], + samples: [ + 1.0610710382461548, 0.37632203102111816, 0.49078676104545593, + 0.33188390731811523, + ]), + tolerance: .float32) + } + } + + @Test("product") + func test_product() throws { + try withIntegrationState(seed: 1646) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.product(a) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: -0.09494685381650925, + minimum: -0.09494685381650925, + maximum: -0.09494685381650925, + absoluteSum: 0.09494685381650925, + positionChecksum: 0.09494685381650925, + sampleIndices: [0], + samples: [-0.09494685381650925]), + tolerance: .float32) + } + } + + @Test("product/axis") + func test_product_axis() throws { + try withIntegrationState(seed: 51321) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.product(a, axis: -1) + expectSummary( + result, + ArraySummary( + shape: [4], + dtype: .float32, + mean: 0.2498457133769989, + minimum: -0.070404052734375, + maximum: 0.6425824761390686, + absoluteSum: 1.1401909589767456, + positionChecksum: 0.6402937173843384, + sampleIndices: [0, 1, 2, 3], + samples: [ + 0.6425824761390686, 0.0007186831790022552, -0.070404052734375, + 0.42648571729660034, + ]), + tolerance: .float32) + } + } + + @Test("product/axes") + func test_product_axes() throws { + try withIntegrationState(seed: 89717) { + let a = MLXRandom.normal([2, 3, 4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.product(a, axes: [0, -1]) + expectSummary( + result, + ArraySummary( + shape: [3, 4], + dtype: .float32, + mean: -0.003022553399205208, + minimum: -0.1403750628232956, + maximum: 0.047358572483062744, + absoluteSum: 0.3018551468849182, + positionChecksum: 0.08646868666013081, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.1403750628232956, -0.023114101961255074, 0.047358572483062744, + -0.0006650111172348261, 0.007170831318944693, 0.02495034784078598, + ]), + tolerance: .float32) + } + } + + @Test("product/keepDims") + func test_product_keepDims() throws { + try withIntegrationState(seed: 24620) { + let a = MLXRandom.normal([2, 3, 4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.product(a, axes: [0, -1], keepDims: true) + expectSummary( + result, + ArraySummary( + shape: [1, 3, 4, 1], + dtype: .float32, + mean: -0.01248251087963581, + minimum: -0.30200156569480896, + maximum: 0.234603613615036, + absoluteSum: 0.7817379832267761, + positionChecksum: 0.4315507411956787, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.0029180373530834913, 0.0015591266565024853, 0.01657288521528244, + 0.031356602907180786, -0.10858920961618423, 0.029305333271622658, + ]), + tolerance: .float32) + } + } + + @Test("product/method") + func test_product_method() throws { + try withIntegrationState(seed: 43022) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = a.product(axis: -1) + expectSummary( + result, + ArraySummary( + shape: [4], + dtype: .float32, + mean: -0.7211005687713623, + minimum: -2.585625171661377, + maximum: 0.2048765867948532, + absoluteSum: 3.434208393096924, + positionChecksum: 1.9454559087753296, + sampleIndices: [0, 1, 2, 3], + samples: [ + 0.07002640515565872, -2.585625171661377, -0.5736801624298096, + 0.2048765867948532, + ]), + tolerance: .float32) + } + } + + @Test("logSumExp") + func test_logSumExp() throws { + try withIntegrationState(seed: 10288) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.logSumExp(a) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 3.6624972820281982, + minimum: 3.6624972820281982, + maximum: 3.6624972820281982, + absoluteSum: 3.6624972820281982, + positionChecksum: 3.6624972820281982, + sampleIndices: [0], + samples: [3.6624972820281982]), + tolerance: .float32) + } + } + + @Test("logSumExp/axis") + func test_logSumExp_axis() throws { + try withIntegrationState(seed: 78003) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.logSumExp(a, axis: -1) + expectSummary( + result, + ArraySummary( + shape: [4], + dtype: .float32, + mean: 1.3402810096740723, + minimum: 0.5028432607650757, + maximum: 1.6911296844482422, + absoluteSum: 5.361124038696289, + positionChecksum: 3.1376407146453857, + sampleIndices: [0, 1, 2, 3], + samples: [ + 1.6911296844482422, 1.6588503122329712, 0.5028432607650757, 1.50830078125, + ]), + tolerance: .float32) + } + } + + @Test("logSumExp/axes") + func test_logSumExp_axes() throws { + try withIntegrationState(seed: 33247) { + let a = MLXRandom.normal([2, 3, 4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.logSumExp(a, axes: [0, -1]) + expectSummary( + result, + ArraySummary( + shape: [3, 4], + dtype: .float32, + mean: 2.3641576766967773, + minimum: 1.0966796875, + maximum: 3.5836031436920166, + absoluteSum: 28.369890213012695, + positionChecksum: 15.860939025878906, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 2.2214109897613525, 2.3977997303009033, 2.9473743438720703, + 2.1123032569885254, 2.5586776733398438, 2.292973041534424, + ]), + tolerance: .float32) + } + } + + @Test("logSumExp/keepDims") + func test_logSumExp_keepDims() throws { + try withIntegrationState(seed: 5045) { + let a = MLXRandom.normal([2, 3, 4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.logSumExp(a, axes: [0, -1], keepDims: true) + expectSummary( + result, + ArraySummary( + shape: [1, 3, 4, 1], + dtype: .float32, + mean: 1.9824609756469727, + minimum: 0.9807406663894653, + maximum: 3.0540881156921387, + absoluteSum: 23.789531707763672, + positionChecksum: 12.198070526123047, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 1.818997859954834, 2.2178637981414795, 1.7742036581039429, + 1.2414398193359375, 0.9807406663894653, 1.9903795719146729, + ]), + tolerance: .float32) + } + } + + @Test("logSumExp/method") + func test_logSumExp_method() throws { + try withIntegrationState(seed: 98817) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = a.logSumExp(axis: -1) + expectSummary( + result, + ArraySummary( + shape: [4], + dtype: .float32, + mean: 1.328906536102295, + minimum: 0.47593313455581665, + maximum: 1.9842123985290527, + absoluteSum: 5.31562614440918, + positionChecksum: 3.2436769008636475, + sampleIndices: [0, 1, 2, 3], + samples: [ + 1.28114652633667, 1.9842123985290527, 0.47593313455581665, + 1.5743341445922852, + ]), + tolerance: .float32) + } + } + + @Test("variance") + func test_variance() throws { + try withIntegrationState(seed: 24490) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.variance(a) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 0.3176921606063843, + minimum: 0.3176921606063843, + maximum: 0.3176921606063843, + absoluteSum: 0.3176921606063843, + positionChecksum: 0.3176921606063843, + sampleIndices: [0], + samples: [0.3176921606063843]), + tolerance: .float32) + } + } + + @Test("variance/axis") + func test_variance_axis() throws { + try withIntegrationState(seed: 94775) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.variance(a, axis: -1) + expectSummary( + result, + ArraySummary( + shape: [4], + dtype: .float32, + mean: 0.5114014148712158, + minimum: 0.12645000219345093, + maximum: 1.533726453781128, + absoluteSum: 2.0456056594848633, + positionChecksum: 1.4190025329589844, + sampleIndices: [0, 1, 2, 3], + samples: [ + 0.2399284839630127, 0.12645000219345093, 1.533726453781128, + 0.14550063014030457, + ]), + tolerance: .float32) + } + } + + @Test("variance/ddof") + func test_variance_ddof() throws { + try withIntegrationState(seed: 73532) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.variance(a, axis: -1, ddof: 1) + expectSummary( + result, + ArraySummary( + shape: [4], + dtype: .float32, + mean: 1.3407645225524902, + minimum: 0.09674856066703796, + maximum: 2.520488739013672, + absoluteSum: 5.363058090209961, + positionChecksum: 3.5370237827301025, + sampleIndices: [0, 1, 2, 3], + samples: [ + 0.09674856066703796, 2.520488739013672, 1.9729126691818237, + 0.7729080319404602, + ]), + tolerance: .float32) + } + } + + @Test("std") + func test_std() throws { + try withIntegrationState(seed: 24170) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.std(a) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 1.0559873580932617, + minimum: 1.0559873580932617, + maximum: 1.0559873580932617, + absoluteSum: 1.0559873580932617, + positionChecksum: 1.0559873580932617, + sampleIndices: [0], + samples: [1.0559873580932617]), + tolerance: .float32) + } + } + + @Test("std/axis") + func test_std_axis() throws { + try withIntegrationState(seed: 15089) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.std(a, axis: -1) + expectSummary( + result, + ArraySummary( + shape: [4], + dtype: .float32, + mean: 0.6319688558578491, + minimum: 0.5627264380455017, + maximum: 0.6639640927314758, + absoluteSum: 2.5278754234313965, + positionChecksum: 1.5501139163970947, + sampleIndices: [0, 1, 2, 3], + samples: [ + 0.6447125673294067, 0.6564723253250122, 0.6639640927314758, + 0.5627264380455017, + ]), + tolerance: .float32) + } + } + + @Test("std/ddof") + func test_std_ddof() throws { + try withIntegrationState(seed: 23642) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.std(a, axis: -1, ddof: 1) + expectSummary( + result, + ArraySummary( + shape: [4], + dtype: .float32, + mean: 0.6409159898757935, + minimum: 0.15130622684955597, + maximum: 1.0482271909713745, + absoluteSum: 2.563663959503174, + positionChecksum: 1.684018611907959, + sampleIndices: [0, 1, 2, 3], + samples: [ + 0.6390142440795898, 0.7251163721084595, 0.15130622684955597, + 1.0482271909713745, + ]), + tolerance: .float32) + } + } + + @Test("median") + func test_median() throws { + try withIntegrationState(seed: 76670) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.median(a) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: -0.48609113693237305, + minimum: -0.48609113693237305, + maximum: -0.48609113693237305, + absoluteSum: 0.48609113693237305, + positionChecksum: 0.48609113693237305, + sampleIndices: [0], + samples: [-0.48609113693237305]), + tolerance: .float32) + } + } + + @Test("median/axis") + func test_median_axis() throws { + try withIntegrationState(seed: 54563) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.median(a, axis: -1) + expectSummary( + result, + ArraySummary( + shape: [4], + dtype: .float32, + mean: 0.21008887887001038, + minimum: -0.07639828324317932, + maximum: 0.35667484998703003, + absoluteSum: 0.9931520223617554, + positionChecksum: 0.6930854320526123, + sampleIndices: [0, 1, 2, 3], + samples: [ + 0.2302650660276413, -0.07639828324317932, 0.35667484998703003, + 0.3298138678073883, + ]), + tolerance: .float32) + } + } + + @Test("all") + func test_all() throws { + try withIntegrationState(seed: 78751) { + let a = MLXRandom.bernoulli(0.5, [4, 3]) + let result = MLX.all(a) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .bool, + mean: 0.0, + minimum: 0.0, + maximum: 0.0, + absoluteSum: 0.0, + positionChecksum: 0.0, + sampleIndices: [0], + samples: [0.0]), + tolerance: .exact) + } + } + + @Test("all/axis") + func test_all_axis() throws { + try withIntegrationState(seed: 95363) { + let a = MLXRandom.bernoulli(0.5, [4, 3]) + let result = MLX.all(a, axis: -1) + expectSummary( + result, + ArraySummary( + shape: [4], + dtype: .bool, + mean: 0.0, + minimum: 0.0, + maximum: 0.0, + absoluteSum: 0.0, + positionChecksum: 0.0, + sampleIndices: [0, 1, 2, 3], + samples: [0.0, 0.0, 0.0, 0.0]), + tolerance: .exact) + } + } + + @Test("any") + func test_any() throws { + try withIntegrationState(seed: 12694) { + let a = MLXRandom.bernoulli(0.5, [4, 3]) + let result = MLX.any(a) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .bool, + mean: 1.0, + minimum: 1.0, + maximum: 1.0, + absoluteSum: 1.0, + positionChecksum: 1.0, + sampleIndices: [0], + samples: [1.0]), + tolerance: .exact) + } + } + + @Test("any/axis") + func test_any_axis() throws { + try withIntegrationState(seed: 52226) { + let a = MLXRandom.bernoulli(0.5, [4, 3]) + let result = MLX.any(a, axis: -1) + expectSummary( + result, + ArraySummary( + shape: [4], + dtype: .bool, + mean: 0.75, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 3.0, + positionChecksum: 2.0, + sampleIndices: [0, 1, 2, 3], + samples: [1.0, 0.0, 1.0, 1.0]), + tolerance: .exact) + } + } + + @Test("argMin") + func test_argMin() throws { + try withIntegrationState(seed: 22413) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.argMin(a) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .uint32, + mean: 1.0, + minimum: 1.0, + maximum: 1.0, + absoluteSum: 1.0, + positionChecksum: 1.0, + sampleIndices: [0], + samples: [1.0]), + tolerance: .exact) + } + } + + @Test("argMin/axis") + func test_argMin_axis() throws { + try withIntegrationState(seed: 96731) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.argMin(a, axis: -1) + expectSummary( + result, + ArraySummary( + shape: [4], + dtype: .uint32, + mean: 1.25, + minimum: 0.0, + maximum: 2.0, + absoluteSum: 5.0, + positionChecksum: 2.75, + sampleIndices: [0, 1, 2, 3], + samples: [1.0, 2.0, 2.0, 0.0]), + tolerance: .exact) + } + } + + @Test("argMax") + func test_argMax() throws { + try withIntegrationState(seed: 30164) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.argMax(a) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .uint32, + mean: 8.0, + minimum: 8.0, + maximum: 8.0, + absoluteSum: 8.0, + positionChecksum: 8.0, + sampleIndices: [0], + samples: [8.0]), + tolerance: .exact) + } + } + + @Test("argMax/axis") + func test_argMax_axis() throws { + try withIntegrationState(seed: 36698) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.argMax(a, axis: -1) + expectSummary( + result, + ArraySummary( + shape: [4], + dtype: .uint32, + mean: 1.75, + minimum: 1.0, + maximum: 2.0, + absoluteSum: 7.0, + positionChecksum: 4.5, + sampleIndices: [0, 1, 2, 3], + samples: [2.0, 1.0, 2.0, 2.0]), + tolerance: .exact) + } + } + + @Test("cumsum/axis") + func test_cumsum_axis() throws { + try withIntegrationState(seed: 2662) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.cumsum(a, axis: -1) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.10426246374845505, + minimum: -2.4091546535491943, + maximum: 2.405540943145752, + absoluteSum: 10.7410888671875, + positionChecksum: 8.23110834757487, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.11669416725635529, -0.1358266919851303, 0.07263723015785217, + 1.0716676712036133, -1.257238745689392, -2.4091546535491943, + ]), + tolerance: .float32) + } + } + + @Test("cumsum/reverse") + func test_cumsum_reverse() throws { + try withIntegrationState(seed: 60267) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.cumsum(a, axis: -1, reverse: true) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -1.1823921203613281, + minimum: -3.4265153408050537, + maximum: 1.0207629203796387, + absoluteSum: 17.93625259399414, + positionChecksum: 7.194801330566406, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -3.4265153408050537, -1.890081763267517, -2.832381248474121, + -0.3406890034675598, -1.968650460243225, 0.3049541711807251, + ]), + tolerance: .float32) + } + } + + @Test("cumsum/exclusive") + func test_cumsum_exclusive() throws { + try withIntegrationState(seed: 81804) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.cumsum(a, axis: -1, inclusive: false) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.3773741126060486, + minimum: -3.1300058364868164, + maximum: 1.900256633758545, + absoluteSum: 11.56186294555664, + positionChecksum: 7.195838928222656, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.0, -1.3729093074798584, 1.349768877029419, -2.0088613033294678, 0.0, + 1.900256633758545, + ]), + tolerance: .float32) + } + } + + @Test("cumprod/axis") + func test_cumprod_axis() throws { + try withIntegrationState(seed: 92128) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.cumprod(a, axis: -1) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.6749899387359619, + minimum: -3.232989549636841, + maximum: 1.2675939798355103, + absoluteSum: 13.808900833129883, + positionChecksum: 6.150314966837565, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -2.8824124336242676, 0.42125698924064636, 0.020262163132429123, + -2.222736358642578, 0.544766902923584, 1.2675939798355103, + ]), + tolerance: .float32) + } + } + + @Test("cumprod/reverse") + func test_cumprod_reverse() throws { + try withIntegrationState(seed: 88787) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.cumprod(a, axis: -1, reverse: true) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.013112984597682953, + minimum: -1.2184795141220093, + maximum: 1.3128567934036255, + absoluteSum: 7.380435943603516, + positionChecksum: 4.74104372660319, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.35103660821914673, 0.676842451095581, 0.7622838020324707, + 0.04095830023288727, -1.0784687995910645, -1.2184795141220093, + ]), + tolerance: .float32) + } + } + + @Test("cumprod/exclusive") + func test_cumprod_exclusive() throws { + try withIntegrationState(seed: 35949) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.cumprod(a, axis: -1, inclusive: false) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.18877795338630676, + minimum: -1.8461512327194214, + maximum: 1.2107491493225098, + absoluteSum: 10.803426742553711, + positionChecksum: 4.96627680460612, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 1.0, -1.8461512327194214, -0.90610271692276, 0.6101259589195251, 1.0, + 0.3891559839248657, + ]), + tolerance: .float32) + } + } + + @Test("cummax/axis") + func test_cummax_axis() throws { + try withIntegrationState(seed: 48106) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.cummax(a, axis: -1) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.3044331669807434, + minimum: -1.4819244146347046, + maximum: 1.6274610757827759, + absoluteSum: 8.384651184082031, + positionChecksum: 5.544460296630859, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.7419719696044922, 0.7419719696044922, -0.04854029417037964, + -0.3050103485584259, -1.4819244146347046, 1.6274610757827759, + ]), + tolerance: .float32) + } + } + + @Test("cummax/reverse") + func test_cummax_reverse() throws { + try withIntegrationState(seed: 62634) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.cummax(a, axis: -1, reverse: true) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.5222153663635254, + minimum: -1.3533363342285156, + maximum: 1.8465720415115356, + absoluteSum: 10.072877883911133, + positionChecksum: 5.809856414794922, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.820504367351532, 0.17221690714359283, -0.27490541338920593, + 1.8465720415115356, 0.864747166633606, 0.24279813468456268, + ]), + tolerance: .float32) + } + } + + @Test("cummax/exclusive") + func test_cummax_exclusive() throws { + try withIntegrationState(seed: 25900) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.cummax(a, axis: -1, inclusive: false) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -Double.infinity, + minimum: -Double.infinity, + maximum: 0.9814871549606323, + absoluteSum: Double.infinity, + positionChecksum: Double.infinity, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -Double.infinity, 0.9814871549606323, 0.16378529369831085, + -0.5856958031654358, -Double.infinity, -1.9102404117584229, + ]), + tolerance: .float32) + } + } + + @Test("cummin/axis") + func test_cummin_axis() throws { + try withIntegrationState(seed: 94795) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.cummin(a, axis: -1) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.7058196067810059, + minimum: -1.1635726690292358, + maximum: 0.8255786299705505, + absoluteSum: 10.120992660522461, + positionChecksum: 5.628684997558594, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.8448988199234009, -0.8448988199234009, -1.0783499479293823, + -0.9089512825012207, 0.8255786299705505, -0.8575506806373596, + ]), + tolerance: .float32) + } + } + + @Test("cummin/reverse") + func test_cummin_reverse() throws { + try withIntegrationState(seed: 80709) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.cummin(a, axis: -1, reverse: true) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.37917232513427734, + minimum: -2.6307315826416016, + maximum: 2.1301822662353516, + absoluteSum: 12.440810203552246, + positionChecksum: 6.774248123168945, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.9820541143417358, 0.7694716453552246, -0.08122467249631882, + -2.6307315826416016, 0.2754596769809723, 0.44595789909362793, + ]), + tolerance: .float32) + } + } + + @Test("cummin/exclusive") + func test_cummin_exclusive() throws { + try withIntegrationState(seed: 21818) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.cummin(a, axis: -1, inclusive: false) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: Double.infinity, + minimum: -0.9863429665565491, + maximum: Double.infinity, + absoluteSum: Double.infinity, + positionChecksum: Double.infinity, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + Double.infinity, -0.716964066028595, 0.8032096028327942, + -0.37494421005249023, Double.infinity, -0.9684954285621643, + ]), + tolerance: .float32) + } + } + + @Test("logCumsumExp/axis") + func test_logCumsumExp_axis() throws { + try withIntegrationState(seed: 68907) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.logCumsumExp(a, axis: -1) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 1.124143123626709, + minimum: -0.18204832077026367, + maximum: 2.4726126194000244, + absoluteSum: 14.127202987670898, + positionChecksum: 9.762095133463541, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.18204832077026367, 0.888251006603241, 0.21538907289505005, + 2.3542518615722656, 1.4005087614059448, 1.8368728160858154, + ]), + tolerance: .float32) + } + } + + @Test("logCumsumExp/reverse") + func test_logCumsumExp_reverse() throws { + try withIntegrationState(seed: 2059) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.logCumsumExp(a, axis: -1, reverse: true) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.058185141533613205, + minimum: -1.9204189777374268, + maximum: 1.5514507293701172, + absoluteSum: 13.139307022094727, + positionChecksum: 6.322711944580078, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 1.4820340871810913, -1.6690078973770142, 1.4352582693099976, + -1.3616853952407837, 0.6133323311805725, -1.0922249555587769, + ]), + tolerance: .float32) + } + } + + @Test("logCumsumExp/exclusive") + func test_logCumsumExp_exclusive() throws { + try withIntegrationState(seed: 27959) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.logCumsumExp(a, axis: -1, inclusive: false) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -Double.infinity, + minimum: -Double.infinity, + maximum: 1.805098056793213, + absoluteSum: Double.infinity, + positionChecksum: Double.infinity, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -Double.infinity, 0.3031471371650696, -0.4027698040008545, + 1.8003222942352295, -Double.infinity, 0.24152666330337524, + ]), + tolerance: .float32) + } + } + + @Test("softmax/axis") + func test_softmax_axis() throws { + try withIntegrationState(seed: 96786) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.softmax(a, axis: -1) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.3333333432674408, + minimum: 0.02120232954621315, + maximum: 0.927946150302887, + absoluteSum: 4.0, + positionChecksum: 2.1860249837239585, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.060398709028959274, 0.887180745601654, 0.02120232954621315, + 0.5514998435974121, 0.10449923574924469, 0.4300844967365265, + ]), + tolerance: .float32) + } + } + + @Test("softmax/axes") + func test_softmax_axes() throws { + try withIntegrationState(seed: 10846) { + let a = MLXRandom.normal([2, 3, 4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.softmax(a, axes: [0, -1]) + expectSummary( + result, + ArraySummary( + shape: [2, 3, 4, 3], + dtype: .float32, + mean: 0.1666666716337204, + minimum: 0.02015802077949047, + maximum: 0.5400499105453491, + absoluteSum: 12.0, + positionChecksum: 6.345125834147136, + sampleIndices: [0, 14, 28, 43, 57, 71], + samples: [ + 0.22917267680168152, 0.04581531509757042, 0.4089527130126953, + 0.3178185224533081, 0.030428849160671234, 0.3574194312095642, + ]), + tolerance: .float32) + } + } + + @Test("softmax/precise") + func test_softmax_precise() throws { + // precise: computes in float32; the python side is already float32 + try withIntegrationState(seed: 68889) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.softmax(a, axis: -1, precise: true) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.3333333432674408, + minimum: 0.016829494386911392, + maximum: 0.6464029550552368, + absoluteSum: 4.0, + positionChecksum: 2.197746435801188, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.10401526838541031, 0.24958176910877228, 0.016829494386911392, + 0.5826361179351807, 0.26546579599380493, 0.6221785545349121, + ]), + tolerance: .float32) + } + } +} diff --git a/Tests/MLXIntegrationTests/Generated/GeneratedSchedulesTests.swift b/Tests/MLXIntegrationTests/Generated/GeneratedSchedulesTests.swift new file mode 100644 index 000000000..b774efd8c --- /dev/null +++ b/Tests/MLXIntegrationTests/Generated/GeneratedSchedulesTests.swift @@ -0,0 +1,224 @@ +// Copyright © 2026 Apple Inc. +// +// GENERATED by tools/integration_tests -- DO NOT EDIT. +// +// Values in this file were produced by python mlx and are compared against +// the Swift API. Regenerate with: +// +// python3 tools/integration_tests/generate.py +// +// python mlx: 0.32.2.dev20260910+1f8e74e3f +// vendored mlx (Cmlx): 0.32.2 +// generator revision: 2 +// cases: 8 + +import Foundation +import MLX +import Testing + +@testable import MLXOptimizers + +@Suite("generated: Schedules") +struct GeneratedSchedulesTests { + + @Test("exponentialDecay") + func test_exponentialDecay() throws { + try withIntegrationState(seed: 61845) { + let schedule = exponentialDecay(0.1, decayRate: 0.9) + let values = MLXArray((0 ..< 12).map { schedule($0) }) + expectSummary( + values, + ArraySummary( + shape: [12], + dtype: .float32, + mean: 0.059797536581754684, + minimum: 0.03138105943799019, + maximum: 0.10000000149011612, + absoluteSum: 0.717570424079895, + positionChecksum: 0.31554583708445233, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.10000000149011612, 0.08100000023841858, 0.06560999900102615, + 0.04782969132065773, 0.038742050528526306, 0.03138105943799019, + ]), + tolerance: .float32) + } + } + + @Test("stepDecay") + func test_stepDecay() throws { + try withIntegrationState(seed: 34315) { + let schedule = stepDecay(0.1, decayRate: 0.5, stepSize: 3) + let values = MLXArray((0 ..< 12).map { schedule($0) }) + expectSummary( + values, + ArraySummary( + shape: [12], + dtype: .float32, + mean: 0.046875, + minimum: 0.012500000186264515, + maximum: 0.10000000149011612, + absoluteSum: 0.5625, + positionChecksum: 0.19687501589457193, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.10000000149011612, 0.10000000149011612, 0.05000000074505806, + 0.02500000037252903, 0.012500000186264515, 0.012500000186264515, + ]), + tolerance: .float32) + } + } + + @Test("cosineDecay") + func test_cosineDecay() throws { + // constant at the end value beyond decaySteps, which the 12 samples cover + try withIntegrationState(seed: 58536) { + let schedule = cosineDecay(0.1, decaySteps: 8) + let values = MLXArray((0 ..< 12).map { schedule($0) }) + expectSummary( + values, + ArraySummary( + shape: [12], + dtype: .float32, + mean: 0.03750000149011612, + minimum: 0.0, + maximum: 0.10000000149011612, + absoluteSum: 0.45000001788139343, + positionChecksum: 0.11609552303949992, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.10000000149011612, 0.08535534143447876, 0.04999999701976776, + 0.0038060189690440893, 0.0, 0.0, + ]), + tolerance: .float32) + } + } + + @Test("cosineDecay/end") + func test_cosineDecay_end() throws { + try withIntegrationState(seed: 5429) { + let schedule = cosineDecay(0.1, decaySteps: 8, end: 0.01) + let values = MLXArray((0 ..< 12).map { schedule($0) }) + expectSummary( + values, + ArraySummary( + shape: [12], + dtype: .float32, + mean: 0.04374999925494194, + minimum: 0.009999999776482582, + maximum: 0.10000000149011612, + absoluteSum: 0.5249999761581421, + positionChecksum: 0.16948598623275757, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.10000000149011612, 0.08681980520486832, 0.054999999701976776, + 0.013425417244434357, 0.009999999776482582, 0.009999999776482582, + ]), + tolerance: .float32) + } + } + + @Test("linearSchedule") + func test_linearSchedule() throws { + try withIntegrationState(seed: 81538) { + let schedule = linearSchedule(0.0, end: 0.1, steps: 8) + let values = MLXArray((0 ..< 12).map { schedule($0) }) + expectSummary( + values, + ArraySummary( + shape: [12], + dtype: .float32, + mean: 0.0625, + minimum: 0.0, + maximum: 0.10000000149011612, + absoluteSum: 0.75, + positionChecksum: 0.5250000158945719, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.0, 0.02500000037252903, 0.05000000074505806, 0.08749999850988388, + 0.10000000149011612, 0.10000000149011612, + ]), + tolerance: .float32) + } + } + + @Test("linearSchedule/down") + func test_linearSchedule_down() throws { + try withIntegrationState(seed: 61284) { + let schedule = linearSchedule(0.1, end: 0.0, steps: 5) + let values = MLXArray((0 ..< 12).map { schedule($0) }) + expectSummary( + values, + ArraySummary( + shape: [12], + dtype: .float32, + mean: 0.02500000409781933, + minimum: 7.450580596923828e-09, + maximum: 0.10000000149011612, + absoluteSum: 0.30000004172325134, + positionChecksum: 0.05833337207635244, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.10000000149011612, 0.06000000238418579, 0.020000003278255463, + 7.450580596923828e-09, 7.450580596923828e-09, 7.450580596923828e-09, + ]), + tolerance: .float32) + } + } + + @Test("joinSchedules") + func test_joinSchedules() throws { + // warmup then decay -- the classic use, and the boundary handling is the part that is easy to get wrong + try withIntegrationState(seed: 21828) { + let schedule = joinSchedules( + [linearSchedule(0.0, end: 0.1, steps: 4), cosineDecay(0.1, decaySteps: 8)], + boundaries: [4]) + let values = MLXArray((0 ..< 12).map { schedule($0) }) + expectSummary( + values, + ArraySummary( + shape: [12], + dtype: .float32, + mean: 0.05000000447034836, + minimum: 0.0, + maximum: 0.10000000149011612, + absoluteSum: 0.6000000238418579, + positionChecksum: 0.3077622056007385, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.0, 0.05000000074505806, 0.10000000149011612, 0.0691341683268547, + 0.030865823850035667, 0.0038060189690440893, + ]), + tolerance: .float32) + } + } + + @Test("joinSchedules/three") + func test_joinSchedules_three() throws { + try withIntegrationState(seed: 92857) { + let schedule = joinSchedules( + [ + linearSchedule(0.0, end: 0.1, steps: 3), + stepDecay(0.1, decayRate: 0.5, stepSize: 2), + linearSchedule(0.05, end: 0.0, steps: 4), + ], boundaries: [3, 7]) + let values = MLXArray((0 ..< 12).map { schedule($0) }) + expectSummary( + values, + ArraySummary( + shape: [12], + dtype: .float32, + mean: 0.04375000298023224, + minimum: 0.0, + maximum: 0.10000000149011612, + absoluteSum: 0.5250000357627869, + positionChecksum: 0.24513888359069824, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.0, 0.06666667014360428, 0.10000000149011612, 0.05000000074505806, + 0.02500000037252903, 0.0, + ]), + tolerance: .float32) + } + } +} diff --git a/Tests/MLXIntegrationTests/Generated/GeneratedShapeTests.swift b/Tests/MLXIntegrationTests/Generated/GeneratedShapeTests.swift new file mode 100644 index 000000000..bed682645 --- /dev/null +++ b/Tests/MLXIntegrationTests/Generated/GeneratedShapeTests.swift @@ -0,0 +1,1172 @@ +// Copyright © 2026 Apple Inc. +// +// GENERATED by tools/integration_tests -- DO NOT EDIT. +// +// Values in this file were produced by python mlx and are compared against +// the Swift API. Regenerate with: +// +// python3 tools/integration_tests/generate.py +// +// python mlx: 0.32.2.dev20260910+1f8e74e3f +// vendored mlx (Cmlx): 0.32.2 +// generator revision: 2 +// cases: 49 + +import Foundation +import MLX +import Testing + +@Suite("generated: Shape") +struct GeneratedShapeTests { + + @Test("concatenated") + func test_concatenated() throws { + try withIntegrationState(seed: 31820) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.concatenated([a, b]) + expectSummary( + result, + ArraySummary( + shape: [8, 3], + dtype: .float32, + mean: 0.08590242266654968, + minimum: -1.2395046949386597, + maximum: 1.4470280408859253, + absoluteSum: 16.01652717590332, + positionChecksum: 7.683911005655925, + sampleIndices: [0, 5, 9, 14, 18, 23], + samples: [ + -0.6949948668479919, 1.133658766746521, -0.05453184247016907, + -0.12073975056409836, 0.3892839252948761, 0.24856531620025635, + ]), + tolerance: .float32) + } + } + + @Test("concatenated/axis") + func test_concatenated_axis() throws { + try withIntegrationState(seed: 20890) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.concatenated([a, b], axis: 1) + expectSummary( + result, + ArraySummary( + shape: [4, 6], + dtype: .float32, + mean: -0.08199743926525116, + minimum: -2.4462075233459473, + maximum: 1.5454013347625732, + absoluteSum: 17.703125, + positionChecksum: 8.136674880981445, + sampleIndices: [0, 5, 9, 14, 18, 23], + samples: [ + 1.5267335176467896, -1.4755029678344727, 0.2989487051963806, + -0.7911329865455627, -0.8274667263031006, -0.6284770965576172, + ]), + tolerance: .float32) + } + } + + @Test("stacked") + func test_stacked() throws { + try withIntegrationState(seed: 94057) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.stacked([a, b]) + expectSummary( + result, + ArraySummary( + shape: [2, 4, 3], + dtype: .float32, + mean: -0.2751633822917938, + minimum: -1.3530937433242798, + maximum: 1.9417942762374878, + absoluteSum: 18.261486053466797, + positionChecksum: 9.054578145345053, + sampleIndices: [0, 5, 9, 14, 18, 23], + samples: [ + 0.9709651470184326, -0.3056058883666992, -0.9813033938407898, + -1.3027057647705078, -0.97268146276474, 0.43900448083877563, + ]), + tolerance: .float32) + } + } + + @Test("stacked/axis") + func test_stacked_axis() throws { + try withIntegrationState(seed: 63795) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let b = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.stacked([a, b], axis: -1) + expectSummary( + result, + ArraySummary( + shape: [4, 3, 2], + dtype: .float32, + mean: -0.2341117560863495, + minimum: -2.2961699962615967, + maximum: 1.4487446546554565, + absoluteSum: 20.772422790527344, + positionChecksum: 11.166987101236979, + sampleIndices: [0, 5, 9, 14, 18, 23], + samples: [ + 0.1765039563179016, -0.06140340119600296, 1.031309962272644, + 1.1500154733657837, -2.2923014163970947, -1.327276587486267, + ]), + tolerance: .float32) + } + } + + @Test("flipped") + func test_flipped() throws { + try withIntegrationState(seed: 15319) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.flipped(a) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.2945154309272766, + minimum: -1.2866392135620117, + maximum: 1.67825448513031, + absoluteSum: 8.736289978027344, + positionChecksum: 5.237860679626465, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.4458805322647095, -0.5673931837081909, 0.14704769849777222, + 0.32986971735954285, -1.144845962524414, -0.7193918228149414, + ]), + tolerance: .float32) + } + } + + @Test("flipped/axis") + func test_flipped_axis() throws { + try withIntegrationState(seed: 45420) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.flipped(a, axis: 1) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.1344703733921051, + minimum: -1.337192416191101, + maximum: 1.5279310941696167, + absoluteSum: 11.764471054077148, + positionChecksum: 6.069742838541667, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 1.5263196229934692, 0.5950191020965576, -1.337192416191101, + -1.1254405975341797, 0.7738415002822876, -0.5284457206726074, + ]), + tolerance: .float32) + } + } + + @Test("tiled") + func test_tiled() throws { + try withIntegrationState(seed: 25082) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.tiled(a, repetitions: [2, 1]) + expectSummary( + result, + ArraySummary( + shape: [8, 3], + dtype: .float32, + mean: 0.13411381840705872, + minimum: -0.9933890700340271, + maximum: 1.0480691194534302, + absoluteSum: 13.16995620727539, + positionChecksum: 6.219711939493815, + sampleIndices: [0, 5, 9, 14, 18, 23], + samples: [ + -0.9933890700340271, -0.25454410910606384, -0.058666568249464035, + 0.957951545715332, 0.1598086655139923, 0.7156504392623901, + ]), + tolerance: .float32) + } + } + + @Test("repeated") + func test_repeated() throws { + try withIntegrationState(seed: 11706) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.repeated(a, count: 2, axis: 1) + expectSummary( + result, + ArraySummary( + shape: [4, 6], + dtype: .float32, + mean: -0.0626811534166336, + minimum: -1.3356250524520874, + maximum: 1.4082560539245605, + absoluteSum: 11.977014541625977, + positionChecksum: 7.969771067301433, + sampleIndices: [0, 5, 9, 14, 18, 23], + samples: [ + 0.03741425648331642, 0.7571091055870056, 0.26252982020378113, + -1.3356250524520874, 0.012510073371231556, 1.4082560539245605, + ]), + tolerance: .float32) + } + } + + @Test("padded") + func test_padded() throws { + try withIntegrationState(seed: 8104) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.padded(a, width: 2) + expectSummary( + result, + ArraySummary( + shape: [8, 7], + dtype: .float32, + mean: -0.016057996079325676, + minimum: -1.319933295249939, + maximum: 1.7083896398544312, + absoluteSum: 9.278562545776367, + positionChecksum: 4.887723650251116, + sampleIndices: [0, 11, 22, 33, 44, 55], + samples: [0.0, 0.0, 0.0, 0.0, 0.0, 0.0]), + tolerance: .float32) + } + } + + @Test("padded/edge") + func test_padded_edge() throws { + try withIntegrationState(seed: 78439) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.padded(a, width: 2, mode: .edge) + expectSummary( + result, + ArraySummary( + shape: [8, 7], + dtype: .float32, + mean: 0.650275707244873, + minimum: -1.475324273109436, + maximum: 1.795180082321167, + absoluteSum: 59.184226989746094, + positionChecksum: 35.910391671316965, + sampleIndices: [0, 11, 22, 33, 44, 55], + samples: [ + -0.6165251731872559, 0.5478887557983398, 0.3702576458454132, + 1.1758822202682495, 1.7887344360351562, 1.2311887741088867, + ]), + tolerance: .float32) + } + } + + @Test("padded/widths") + func test_padded_widths() throws { + try withIntegrationState(seed: 15378) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.padded(a, widths: [IntOrPair([1, 2]), IntOrPair([0, 1])]) + expectSummary( + result, + ArraySummary( + shape: [7, 4], + dtype: .float32, + mean: 0.03199724853038788, + minimum: -1.4607915878295898, + maximum: 1.7407435178756714, + absoluteSum: 8.945435523986816, + positionChecksum: 4.13947514125279, + sampleIndices: [0, 5, 11, 16, 22, 27], + samples: [0.0, -0.9670328497886658, 0.0, -0.6681327223777771, 0.0, 0.0]), + tolerance: .float32) + } + } + + @Test("padded/value") + func test_padded_value() throws { + try withIntegrationState(seed: 16948) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.padded(a, width: 1, value: MLXArray(Float(2.5))) + expectSummary( + result, + ArraySummary( + shape: [6, 5], + dtype: .float32, + mean: 1.8162819147109985, + minimum: -0.7915933132171631, + maximum: 2.8892898559570312, + absoluteSum: 56.07164001464844, + positionChecksum: 28.535982259114583, + sampleIndices: [0, 6, 12, 17, 23, 29], + samples: [ + 2.5, -0.7915933132171631, 2.8892898559570312, 0.3641357123851776, + 0.6989915370941162, 2.5, + ]), + tolerance: .float32) + } + } + + @Test("roll") + func test_roll() throws { + try withIntegrationState(seed: 19553) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.roll(a, shift: 2, axis: 0) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.26138874888420105, + minimum: -2.099207639694214, + maximum: 1.4279882907867432, + absoluteSum: 10.050413131713867, + positionChecksum: 5.608279546101888, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.016545826569199562, 1.4216101169586182, -0.5516468286514282, + 1.170422077178955, 0.0015867743641138077, 1.4279882907867432, + ]), + tolerance: .float32) + } + } + + @Test("transposed") + func test_transposed() throws { + try withIntegrationState(seed: 60390) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.transposed(a) + expectSummary( + result, + ArraySummary( + shape: [3, 4], + dtype: .float32, + mean: 0.0771273821592331, + minimum: -1.0791972875595093, + maximum: 2.472217082977295, + absoluteSum: 10.713860511779785, + positionChecksum: 5.417025248209636, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 1.5892202854156494, -1.0791972875595093, 0.30885958671569824, + -1.0246267318725586, 1.2981765270233154, -0.5083723068237305, + ]), + tolerance: .float32) + } + } + + @Test("transposed/axes") + func test_transposed_axes() throws { + try withIntegrationState(seed: 83812) { + let a = MLXRandom.normal([2, 3, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.transposed(a, axes: [2, 0, 1]) + expectSummary( + result, + ArraySummary( + shape: [4, 2, 3], + dtype: .float32, + mean: -0.02823314070701599, + minimum: -1.502172827720642, + maximum: 1.6366482973098755, + absoluteSum: 19.033903121948242, + positionChecksum: 10.281106313069662, + sampleIndices: [0, 5, 9, 14, 18, 23], + samples: [ + -0.9478650689125061, 1.2468520402908325, -0.5932059288024902, + -0.07760278135538101, -0.11842215806245804, -1.2484960556030273, + ]), + tolerance: .float32) + } + } + + @Test("reshaped") + func test_reshaped() throws { + try withIntegrationState(seed: 68729) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.reshaped(a, [3, 4]) + expectSummary( + result, + ArraySummary( + shape: [3, 4], + dtype: .float32, + mean: -0.26473572850227356, + minimum: -1.2564424276351929, + maximum: 1.1341665983200073, + absoluteSum: 10.037672996520996, + positionChecksum: 5.030018170674642, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -1.2564424276351929, -0.8097322583198547, 0.3184874653816223, + -1.0139796733856201, -0.5813552141189575, -0.4897230267524719, + ]), + tolerance: .float32) + } + } + + @Test("squeezed") + func test_squeezed() throws { + try withIntegrationState(seed: 41015) { + let a = MLXRandom.normal([4, 1, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.squeezed(a, axis: 1) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.01943860575556755, + minimum: -1.4817900657653809, + maximum: 1.1566555500030518, + absoluteSum: 7.114997863769531, + positionChecksum: 3.67331600189209, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.22527667880058289, -1.4817900657653809, -0.4570019245147705, + -0.14241276681423187, -0.3114512860774994, 0.7334579229354858, + ]), + tolerance: .float32) + } + } + + @Test("expandedDimensions") + func test_expandedDimensions() throws { + try withIntegrationState(seed: 19999) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.expandedDimensions(a, axes: [0, -1]) + expectSummary( + result, + ArraySummary( + shape: [1, 4, 3, 1], + dtype: .float32, + mean: 0.09923312067985535, + minimum: -1.8916503190994263, + maximum: 1.8866270780563354, + absoluteSum: 12.624441146850586, + positionChecksum: 7.273377736409505, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 1.5140305757522583, 1.8866270780563354, -0.9565528631210327, + 1.1745210886001587, -1.0232282876968384, 1.3436799049377441, + ]), + tolerance: .float32) + } + } + + @Test("movedAxis") + func test_movedAxis() throws { + try withIntegrationState(seed: 64658) { + let a = MLXRandom.normal([2, 3, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.movedAxis(a, source: 0, destination: 2) + expectSummary( + result, + ArraySummary( + shape: [3, 4, 2], + dtype: .float32, + mean: -0.14943131804466248, + minimum: -2.045391082763672, + maximum: 1.752381682395935, + absoluteSum: 20.457523345947266, + positionChecksum: 11.733778635660807, + sampleIndices: [0, 5, 9, 14, 18, 23], + samples: [ + 0.7180314660072327, -1.3107750415802002, 0.3325449526309967, + -1.3910497426986694, -0.08477483689785004, -2.045391082763672, + ]), + tolerance: .float32) + } + } + + @Test("swappedAxes") + func test_swappedAxes() throws { + try withIntegrationState(seed: 48013) { + let a = MLXRandom.normal([2, 3, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.swappedAxes(a, 0, 2) + expectSummary( + result, + ArraySummary( + shape: [4, 3, 2], + dtype: .float32, + mean: -0.1168517917394638, + minimum: -1.750725269317627, + maximum: 1.7686270475387573, + absoluteSum: 21.129409790039062, + positionChecksum: 10.469156901041666, + sampleIndices: [0, 5, 9, 14, 18, 23], + samples: [ + -1.2027404308319092, -1.109092354774475, -0.47926244139671326, + 1.6138956546783447, 0.576898992061615, 0.4427987039089203, + ]), + tolerance: .float32) + } + } + + @Test("flattened") + func test_flattened() throws { + try withIntegrationState(seed: 34896) { + let a = MLXRandom.normal([2, 3, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.flattened(a) + expectSummary( + result, + ArraySummary( + shape: [24], + dtype: .float32, + mean: 0.09383618831634521, + minimum: -2.187122344970703, + maximum: 2.043957471847534, + absoluteSum: 21.715774536132812, + positionChecksum: 11.356913248697916, + sampleIndices: [0, 5, 9, 14, 18, 23], + samples: [ + -2.187122344970703, -1.3866022825241089, -0.039398424327373505, + 1.7639521360397339, -1.4106767177581787, -0.17367228865623474, + ]), + tolerance: .float32) + } + } + + @Test("flattened/range") + func test_flattened_range() throws { + try withIntegrationState(seed: 32731) { + let a = MLXRandom.normal([2, 3, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.flattened(a, start: 1, end: 2) + expectSummary( + result, + ArraySummary( + shape: [2, 12], + dtype: .float32, + mean: 0.08807133883237839, + minimum: -1.6179569959640503, + maximum: 1.552121877670288, + absoluteSum: 17.646623611450195, + positionChecksum: 8.385078430175781, + sampleIndices: [0, 5, 9, 14, 18, 23], + samples: [ + -0.7485483288764954, -1.4518442153930664, -1.6179569959640503, + 0.10551830381155014, 1.3873385190963745, -0.1756194829940796, + ]), + tolerance: .float32) + } + } + + @Test("unflatten") + func test_unflatten() throws { + try withIntegrationState(seed: 91953) { + let a = MLXRandom.normal([2, 12], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.unflatten(a, axis: 1, shape: [3, 4]) + expectSummary( + result, + ArraySummary( + shape: [2, 3, 4], + dtype: .float32, + mean: 0.24676339328289032, + minimum: -1.1256226301193237, + maximum: 1.6207804679870605, + absoluteSum: 13.571483612060547, + positionChecksum: 6.478286107381185, + sampleIndices: [0, 5, 9, 14, 18, 23], + samples: [ + 1.2019108533859253, 0.9293265342712402, -0.8174373507499695, + 0.5026535987854004, 0.7566536068916321, 1.0857075452804565, + ]), + tolerance: .float32) + } + } + + @Test("broadcast") + func test_broadcast() throws { + try withIntegrationState(seed: 88800) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.broadcast(a, to: [2, 4, 3]) + expectSummary( + result, + ArraySummary( + shape: [2, 4, 3], + dtype: .float32, + mean: -0.19611158967018127, + minimum: -0.9691779613494873, + maximum: 1.3313171863555908, + absoluteSum: 14.432416915893555, + positionChecksum: 6.544670104980469, + sampleIndices: [0, 5, 9, 14, 18, 23], + samples: [ + -0.9691779613494873, -0.34518253803253174, -0.6659537553787231, + 1.3313171863555908, -0.8420485854148865, 0.033554743975400925, + ]), + tolerance: .float32) + } + } + + @Test("atLeast2D") + func test_atLeast2D() throws { + try withIntegrationState(seed: 76585) { + let a = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.atLeast2D(a) + expectSummary( + result, + ArraySummary( + shape: [1, 5], + dtype: .float32, + mean: 0.9482474327087402, + minimum: -0.18628670275211334, + maximum: 2.5243351459503174, + absoluteSum: 5.1138105392456055, + positionChecksum: 3.568856048583984, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + 0.47301429510116577, -0.18628670275211334, 2.5243351459503174, + 0.2251853346824646, 1.7049888372421265, + ]), + tolerance: .float32) + } + } + + @Test("take") + func test_take() throws { + try withIntegrationState(seed: 96628) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let indices = MLXRandom.randInt(low: 0, high: 4, [5], type: Int32.self) + let result = MLX.take(a, indices, axis: 0) + expectSummary( + result, + ArraySummary( + shape: [5, 3], + dtype: .float32, + mean: -0.13128289580345154, + minimum: -1.813438057899475, + maximum: 1.2486591339111328, + absoluteSum: 13.089598655700684, + positionChecksum: 6.21908213297526, + sampleIndices: [0, 3, 6, 8, 11, 14], + samples: [ + -1.4467936754226685, -1.4467936754226685, -1.813438057899475, + -1.3840922117233276, 0.2269812375307083, -1.2374604940414429, + ]), + tolerance: .float32) + } + } + + @Test("take/flat") + func test_take_flat() throws { + try withIntegrationState(seed: 20188) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let indices = MLXRandom.randInt(low: 0, high: 12, [4], type: Int32.self) + let result = MLX.take(a, indices) + expectSummary( + result, + ArraySummary( + shape: [4], + dtype: .float32, + mean: -0.20158261060714722, + minimum: -1.8244850635528564, + maximum: 0.7776317596435547, + absoluteSum: 2.9791224002838135, + positionChecksum: 1.8212910890579224, + sampleIndices: [0, 1, 2, 3], + samples: [ + -0.06824146211147308, -1.8244850635528564, 0.7776317596435547, + 0.3087643086910248, + ]), + tolerance: .float32) + } + } + + @Test("takeAlong") + func test_takeAlong() throws { + try withIntegrationState(seed: 75580) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let indices = MLXRandom.randInt(low: 0, high: 3, [4, 1], type: Int32.self) + let result = MLX.takeAlong(a, indices, axis: 1) + expectSummary( + result, + ArraySummary( + shape: [4, 1], + dtype: .float32, + mean: 0.9248526692390442, + minimum: -0.45212575793266296, + maximum: 2.0080971717834473, + absoluteSum: 4.603662490844727, + positionChecksum: 3.3495242595672607, + sampleIndices: [0, 1, 2, 3], + samples: [ + 0.18274369835853577, 2.0080971717834473, -0.45212575793266296, + 1.960695505142212, + ]), + tolerance: .float32) + } + } + + @Test("sorted") + func test_sorted() throws { + try withIntegrationState(seed: 9140) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.sorted(a, axis: -1) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.03630068898200989, + minimum: -1.7828558683395386, + maximum: 1.2648265361785889, + absoluteSum: 11.236763954162598, + positionChecksum: 5.297865549723308, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -1.7828558683395386, 1.2601059675216675, -1.2201143503189087, + 0.1945764720439911, -0.847751796245575, 1.2648265361785889, + ]), + tolerance: .float32) + } + } + + @Test("top") + func test_top() throws { + try withIntegrationState(seed: 29038) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.top(a, k: 2, axis: -1) + expectSummary( + result, + ArraySummary( + shape: [4, 2], + dtype: .float32, + mean: 0.7949435710906982, + minimum: -0.42735984921455383, + maximum: 1.7273166179656982, + absoluteSum: 7.214267730712891, + positionChecksum: 3.4773030281066895, + sampleIndices: [0, 1, 3, 4, 6, 7], + samples: [ + 0.8981940746307373, 1.7273166179656982, 1.2254910469055176, + -0.42735984921455383, 0.1055675819516182, 0.7501816153526306, + ]), + tolerance: .float32) + } + } + + @Test("tril") + func test_tril() throws { + try withIntegrationState(seed: 23707) { + let a = MLXRandom.normal([4, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.tril(a) + expectSummary( + result, + ArraySummary( + shape: [4, 4], + dtype: .float32, + mean: 0.0325717031955719, + minimum: -1.7094670534133911, + maximum: 1.558538556098938, + absoluteSum: 10.388333320617676, + positionChecksum: 6.133814334869385, + sampleIndices: [0, 3, 6, 9, 12, 15], + samples: [ + 1.5056425333023071, 0.0, 0.0, -0.28900396823883057, 1.558538556098938, + 0.8308115601539612, + ]), + tolerance: .float32) + } + } + + @Test("triu/k") + func test_triu_k() throws { + try withIntegrationState(seed: 57960) { + let a = MLXRandom.normal([4, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.triu(a, k: 1) + expectSummary( + result, + ArraySummary( + shape: [4, 4], + dtype: .float32, + mean: 0.08439695835113525, + minimum: -2.722848653793335, + maximum: 1.6600544452667236, + absoluteSum: 7.721507549285889, + positionChecksum: 2.9740920066833496, + sampleIndices: [0, 3, 6, 9, 12, 15], + samples: [0.0, 1.6600544452667236, -2.722848653793335, 0.0, 0.0, 0.0]), + tolerance: .float32) + } + } + + @Test("tri") + func test_tri() throws { + try withIntegrationState(seed: 39346) { + let result = MLX.tri(4, m: 3, k: 0, dtype: .float32) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.75, + minimum: 0.0, + maximum: 1.0, + absoluteSum: 9.0, + positionChecksum: 5.583333333333333, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [1.0, 0.0, 1.0, 1.0, 1.0, 1.0]), + tolerance: .float32) + } + } + + @Test("diag") + func test_diag() throws { + try withIntegrationState(seed: 26965) { + let a = MLXRandom.normal([4, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.diag(a) + expectSummary( + result, + ArraySummary( + shape: [4], + dtype: .float32, + mean: -0.21551969647407532, + minimum: -2.043041706085205, + maximum: 0.5067105889320374, + absoluteSum: 3.2240047454833984, + positionChecksum: 2.5592713356018066, + sampleIndices: [0, 1, 2, 3], + samples: [ + 0.5067105889320374, 0.4645487666130066, 0.20970351994037628, + -2.043041706085205, + ]), + tolerance: .float32) + } + } + + @Test("diagonal/offset") + func test_diagonal_offset() throws { + try withIntegrationState(seed: 21806) { + let a = MLXRandom.normal([4, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.diagonal(a, offset: 1) + expectSummary( + result, + ArraySummary( + shape: [3], + dtype: .float32, + mean: 1.1711373329162598, + minimum: 0.8744285702705383, + maximum: 1.4721418619155884, + absoluteSum: 3.5134119987487793, + positionChecksum: 2.2448037465413413, + sampleIndices: [0, 1, 2], + samples: [1.1668416261672974, 1.4721418619155884, 0.8744285702705383]), + tolerance: .float32) + } + } + + @Test("putAlong") + func test_putAlong() throws { + try withIntegrationState(seed: 81478) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let indices = MLXRandom.randInt(low: 0, high: 3, [4, 1], type: Int32.self) + let values = MLXRandom.normal([4, 1], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.putAlong(a, indices, values: values, axis: 1) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.2807766795158386, + minimum: -1.0164457559585571, + maximum: 2.2925961017608643, + absoluteSum: 8.769731521606445, + positionChecksum: 4.279629707336426, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + 0.7069619297981262, -0.7591624855995178, -0.08726170659065247, + -0.21965672075748444, 2.0261762142181396, -0.06557861715555191, + ]), + tolerance: .float32) + } + } + + @Test("argSort") + func test_argSort() throws { + try withIntegrationState(seed: 68988) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.argSort(a, axis: -1) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .uint32, + mean: 1.0, + minimum: 0.0, + maximum: 2.0, + absoluteSum: 12.0, + positionChecksum: 6.5, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [2.0, 1.0, 2.0, 0.0, 2.0, 1.0]), + tolerance: .exact) + } + } + + @Test("argSort/flat") + func test_argSort_flat() throws { + try withIntegrationState(seed: 69849) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.argSort(a) + expectSummary( + result, + ArraySummary( + shape: [12], + dtype: .uint32, + mean: 5.5, + minimum: 0.0, + maximum: 11.0, + absoluteSum: 66.0, + positionChecksum: 28.916666666666668, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [8.0, 9.0, 2.0, 11.0, 0.0, 1.0]), + tolerance: .exact) + } + } + + @Test("partitioned") + func test_partitioned() throws { + try withIntegrationState(seed: 30615) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.partitioned(a, kth: 1, axis: -1) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: -0.32990625500679016, + minimum: -1.1124590635299683, + maximum: 0.795062243938446, + absoluteSum: 7.30184268951416, + positionChecksum: 4.664909362792969, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.405587762594223, 0.795062243938446, -0.04451719671487808, + -0.6547178030014038, -1.1124590635299683, -0.9105625152587891, + ]), + tolerance: .float32) + } + } + + @Test("argPartition") + func test_argPartition() throws { + try withIntegrationState(seed: 22252) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.argPartition(a, kth: 1, axis: -1) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .uint32, + mean: 1.0, + minimum: 0.0, + maximum: 2.0, + absoluteSum: 12.0, + positionChecksum: 6.666666666666667, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [0.0, 1.0, 1.0, 0.0, 2.0, 0.0]), + tolerance: .exact) + } + } + + @Test("diff") + func test_diff() throws { + try withIntegrationState(seed: 47492) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.diff(a) + expectSummary( + result, + ArraySummary( + shape: [4, 2], + dtype: .float32, + mean: -0.19391421973705292, + minimum: -1.7796225547790527, + maximum: 1.2921435832977295, + absoluteSum: 8.09886360168457, + positionChecksum: 4.752481460571289, + sampleIndices: [0, 1, 3, 4, 6, 7], + samples: [ + -0.9081496000289917, 0.06589078903198242, 1.2921435832977295, + 0.8175460696220398, -0.7752882242202759, 1.0981947183609009, + ]), + tolerance: .float32) + } + } + + @Test("diff/n2") + func test_diff_n2() throws { + try withIntegrationState(seed: 36184) { + let a = MLXRandom.normal([4, 5], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.diff(a, n: 2, axis: 1) + expectSummary( + result, + ArraySummary( + shape: [4, 3], + dtype: .float32, + mean: 0.04767768830060959, + minimum: -2.7871458530426025, + maximum: 3.4202797412872314, + absoluteSum: 16.604129791259766, + positionChecksum: 8.50756581624349, + sampleIndices: [0, 2, 4, 7, 9, 11], + samples: [ + -0.49136456847190857, 3.4202797412872314, -0.6970906853675842, + 0.4519670903682709, -0.7846707105636597, 2.357470750808716, + ]), + tolerance: .float32) + } + } + + @Test("asStrided") + func test_asStrided() throws { + try withIntegrationState(seed: 78763) { + let a = MLXRandom.normal([4, 3], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.asStrided(a, [3, 3], strides: [1, 2], offset: 0) + expectSummary( + result, + ArraySummary( + shape: [3, 3], + dtype: .float32, + mean: -0.5163639187812805, + minimum: -2.0306308269500732, + maximum: 0.5134000182151794, + absoluteSum: 6.464855194091797, + positionChecksum: 4.407473246256511, + sampleIndices: [0, 2, 3, 5, 6, 8], + samples: [ + 0.5134000182151794, -1.3329793214797974, -0.3853185772895813, + -0.47415676712989807, 0.1477542519569397, -2.0306308269500732, + ]), + tolerance: .float32) + } + } + + @Test("atLeast1D") + func test_atLeast1D() throws { + try withIntegrationState(seed: 18570) { + let a = MLXRandom.normal([Int](), dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.atLeast1D(a) + expectSummary( + result, + ArraySummary( + shape: [1], + dtype: .float32, + mean: 0.623765230178833, + minimum: 0.623765230178833, + maximum: 0.623765230178833, + absoluteSum: 0.623765230178833, + positionChecksum: 0.623765230178833, + sampleIndices: [0], + samples: [0.623765230178833]), + tolerance: .float32) + } + } + + @Test("atLeast3D") + func test_atLeast3D() throws { + try withIntegrationState(seed: 92904) { + let a = MLXRandom.normal([5], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.atLeast3D(a) + expectSummary( + result, + ArraySummary( + shape: [1, 5, 1], + dtype: .float32, + mean: 0.3729848563671112, + minimum: -1.1092168092727661, + maximum: 1.8132435083389282, + absoluteSum: 6.191127777099609, + positionChecksum: 3.4658935546875, + sampleIndices: [0, 1, 2, 3, 4], + samples: [ + 1.8132435083389282, -1.1092168092727661, 0.9958317875862122, + -1.0538851022720337, 1.2189509868621826, + ]), + tolerance: .float32) + } + } + + @Test("trace") + func test_trace() throws { + try withIntegrationState(seed: 40135) { + let a = MLXRandom.normal([4, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.trace(a) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: 3.576962471008301, + minimum: 3.576962471008301, + maximum: 3.576962471008301, + absoluteSum: 3.576962471008301, + positionChecksum: 3.576962471008301, + sampleIndices: [0], + samples: [3.576962471008301]), + tolerance: .float32) + } + } + + @Test("trace/offset") + func test_trace_offset() throws { + try withIntegrationState(seed: 47702) { + let a = MLXRandom.normal([4, 4], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.trace(a, offset: 1) + expectSummary( + result, + ArraySummary( + shape: [Int](), + dtype: .float32, + mean: -1.7535713911056519, + minimum: -1.7535713911056519, + maximum: -1.7535713911056519, + absoluteSum: 1.7535713911056519, + positionChecksum: 1.7535713911056519, + sampleIndices: [0], + samples: [-1.7535713911056519]), + tolerance: .float32) + } + } + + @Test("hadamardTransform") + func test_hadamardTransform() throws { + try withIntegrationState(seed: 82262) { + let a = MLXRandom.normal([4, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.hadamardTransform(a) + expectSummary( + result, + ArraySummary( + shape: [4, 16], + dtype: .float32, + mean: -0.09357999265193939, + minimum: -2.850245952606201, + maximum: 2.2235944271087646, + absoluteSum: 52.54167175292969, + positionChecksum: 26.58550262451172, + sampleIndices: [0, 13, 25, 38, 50, 63], + samples: [ + -0.9944025278091431, 0.9899498820304871, 1.330984354019165, + -1.6608858108520508, -1.4785182476043701, 0.39248645305633545, + ]), + tolerance: .float32) + } + } + + @Test("hadamardTransform/scale") + func test_hadamardTransform_scale() throws { + try withIntegrationState(seed: 50260) { + let a = MLXRandom.normal([4, 16], dtype: .float32, loc: 0.0, scale: 1.0) + let result = MLX.hadamardTransform(a, scale: 0.25) + expectSummary( + result, + ArraySummary( + shape: [4, 16], + dtype: .float32, + mean: -0.18685631453990936, + minimum: -3.034515380859375, + maximum: 2.81048583984375, + absoluteSum: 58.34088897705078, + positionChecksum: 28.606985092163086, + sampleIndices: [0, 13, 25, 38, 50, 63], + samples: [ + -0.47638797760009766, -0.6138679385185242, -0.5108878016471863, + -1.6855876445770264, -1.0864219665527344, 0.4750545620918274, + ]), + tolerance: .float32) + } + } +} diff --git a/Tests/MLXIntegrationTests/IntegrationSupport.swift b/Tests/MLXIntegrationTests/IntegrationSupport.swift new file mode 100644 index 000000000..65fab9512 --- /dev/null +++ b/Tests/MLXIntegrationTests/IntegrationSupport.swift @@ -0,0 +1,228 @@ +// Copyright © 2025 Apple Inc. + +// Hand written support for the generated integration tests -- NOT generated. +// +// The generated tests (see tools/integration_tests) declare their inputs with a +// seeded random state, evaluate one expression, and compare the result against a +// summary of the value python produced. All of the comparison policy lives here +// so it can be tuned without regenerating. + +import Foundation +import MLX +import MLXNN +import Testing + +#if canImport(Darwin) + import Darwin +#elseif canImport(Glibc) + import Glibc +#endif + +/// Statistics of an `MLXArray`, computed by python `mlx` in +/// `tools/integration_tests/core.py` (`summarize`) and recomputed here. +/// +/// The definitions must match the generator exactly: +/// +/// ``` +/// values = array.astype(float32).reshape(-1) +/// absoluteSum = sum(|values|) +/// positionChecksum = sum(|values| * arange(1, n + 1)) / n +/// samples = values[evenly spaced indices] +/// ``` +/// +/// `positionChecksum` and `samples` are what make this order sensitive: `mean`, +/// `min`, `max` and `absoluteSum` are all invariant under a permutation of the +/// elements, so a wrong axis or a transposed result would otherwise pass. +struct ArraySummary: Sendable { + let shape: [Int] + let dtype: DType + let mean: Double + let minimum: Double + let maximum: Double + let absoluteSum: Double + let positionChecksum: Double + let sampleIndices: [Int] + let samples: [Double] +} + +/// Comparison tolerance: `|actual - expected| <= absolute + relative * |expected|`. +struct Tolerance: Sendable { + let relative: Double + let absolute: Double + + /// integer and bool results must match exactly + static let exact = Tolerance(relative: 0, absolute: 0) + + /// float32 computed on GPU vs python: relative, with a floor for values near zero + static let float32 = Tolerance(relative: 1e-4, absolute: 1e-6) + + /// float16 / bfloat16 + static let float16 = Tolerance(relative: 5e-3, absolute: 1e-3) + + /// for cases where reduction order legitimately differs + static let loose = Tolerance(relative: 1e-2, absolute: 1e-4) +} + +/// Pin float32 matmuls (rather than TF32) for this process, before any MLX call. +/// +/// `MLX_ENABLE_TF32` defaults to *enabled*, and on hardware with neural accelerators +/// (M5 and later) that runs float32 matmuls on the accelerators in TF32. TF32 +/// results differ from float32 by ~1e-4 relative -- far outside the tolerances here +/// -- so every matmul-shaped case (`Linear`, `Conv*`, attention, `GRU`, `Muon`, +/// quantized matmul, `einsum`, ...) would fail unless the machine running the tests +/// matches the machine that generated the values. +/// +/// mlx reads the variable **once**, at its first use, so it has to be set before any +/// MLX work happens. That is why these tests are their own target +/// (`MLXIntegrationTests` in `Package.swift`): nothing else runs in this process, and +/// every generated case enters through `withIntegrationState(seed:)`, so setting it +/// here is enough no matter how the tests are launched. The test plan and CI also +/// set it in the environment; this covers everyone else. +private let matmulPrecisionIsFloat32: Bool = { + setenv("MLX_ENABLE_TF32", "0", 1) + return ProcessInfo.processInfo.environment["MLX_ENABLE_TF32"] == "0" +}() + +/// Run a generated case. +/// +/// - the default device is scoped to the block rather than set globally, so the +/// generated tests can run in parallel with tests that want another device +/// - MLX errors are converted into Swift errors instead of ending the process, +/// so a bad case fails its own test with the mlx message +/// - the random state is task-local and equivalent to python's +/// `mx.random.seed(seed)` +/// - the matmul precision is checked, because the values assume float32 matmuls +func withIntegrationState( + seed: UInt64, sourceLocation: SourceLocation = #_sourceLocation, _ body: () throws -> R +) throws -> R { + #expect( + matmulPrecisionIsFloat32, + """ + could not set MLX_ENABLE_TF32=0: the generated values assume float32 \ + matmuls, not TF32. Set it in the environment, or regenerate with \ + --enable-tf32 on this machine. + """, + sourceLocation: sourceLocation) + + return try Device.withDefaultDevice(.gpu) { + try withError { + try withRandomState(MLXRandom.RandomState(seed: seed), body: body) + } + } +} + +/// Assert that `array` matches the summary python produced. +func expectSummary( + _ array: MLXArray, _ expected: ArraySummary, tolerance: Tolerance = .float32, + sourceLocation: SourceLocation = #_sourceLocation +) { + #expect(array.shape == expected.shape, "shape", sourceLocation: sourceLocation) + #expect(array.dtype == expected.dtype, "dtype", sourceLocation: sourceLocation) + + guard array.shape == expected.shape else { + // the remaining comparisons would be meaningless (and may trap) + return + } + + // `reshaped([-1])` rather than `flattened()`: matches python's `reshape(-1)` + // and also works for a 0-d (scalar) result + let values = array.asType(.float32).reshaped([-1]) + let count = values.size + let absolute = MLX.abs(values) + let weights = MLX.arange(1, count + 1, dtype: .float32) + + expectClose( + Double(values.mean().item(Float.self)), expected.mean, tolerance, "mean", + sourceLocation: sourceLocation) + expectClose( + Double(values.min().item(Float.self)), expected.minimum, tolerance, "minimum", + sourceLocation: sourceLocation) + expectClose( + Double(values.max().item(Float.self)), expected.maximum, tolerance, "maximum", + sourceLocation: sourceLocation) + expectClose( + Double(absolute.sum().item(Float.self)), expected.absoluteSum, tolerance, "absoluteSum", + sourceLocation: sourceLocation) + expectClose( + Double((absolute * weights).sum().item(Float.self)) / Double(count), + expected.positionChecksum, tolerance, "positionChecksum", + sourceLocation: sourceLocation) + + for (index, expectedValue) in zip(expected.sampleIndices, expected.samples) { + expectClose( + Double(values[index].item(Float.self)), expectedValue, tolerance, + "element[\(index)]", sourceLocation: sourceLocation) + } +} + +private func expectClose( + _ actual: Double, _ expected: Double, _ tolerance: Tolerance, _ label: String, + sourceLocation: SourceLocation +) { + if expected.isNaN { + #expect( + actual.isNaN, "\(label): expected nan, got \(actual)", sourceLocation: sourceLocation) + return + } + if expected.isInfinite { + #expect( + actual == expected, "\(label): expected \(expected), got \(actual)", + sourceLocation: sourceLocation) + return + } + + let limit = tolerance.absolute + tolerance.relative * Swift.abs(expected) + let delta = Swift.abs(actual - expected) + #expect( + delta <= limit, + "\(label): expected \(expected), got \(actual) (delta \(delta) > tolerance \(limit))", + sourceLocation: sourceLocation) +} + +// MARK: - Modules + +/// The value a generated module test writes into every parameter. +/// +/// Module tests do **not** rely on python and Swift drawing the same random +/// initialization (they do not: the two implementations draw different numbers of +/// keys in different orders). Instead every parameter is replaced with a +/// deterministic function of its own shape, which both sides can produce exactly. +/// +/// This must match `PARAMETER_PY` in `tools/integration_tests/core.py`: +/// +/// ``` +/// (mx.arange(v.size, dtype=mx.float32).reshape(v.shape) / v.size - 0.5).astype(v.dtype) +/// ``` +func deterministicParameter(_ parameter: MLXArray) -> MLXArray { + let size = parameter.size + let values = MLX.arange(size, dtype: .float32).reshaped(parameter.shape) / Float(size) - 0.5 + return values.asType(parameter.dtype) +} + +/// Assert that a module's parameters are named and shaped exactly as python's. +/// +/// The keys matter beyond this test: mlx-swift uses them to load python +/// checkpoints, so a rename here is a compatibility break (this is why +/// `BatchNorm` spells its keys `running_mean` / `running_var`). The shapes catch +/// a transposed weight, which would otherwise show up as a confusing value +/// mismatch (or a broadcast error) further down. +func expectParameters( + _ module: Module, _ expected: [(String, [Int])], + sourceLocation: SourceLocation = #_sourceLocation +) { + let actual = module.parameters().flattened() + .map { ($0.0, $0.1.shape) } + .sorted { $0.0 < $1.0 } + + #expect( + actual.map { $0.0 } == expected.map { $0.0 }, "parameter keys", + sourceLocation: sourceLocation) + + for (actualParameter, expectedParameter) in zip(actual, expected) + where actualParameter.0 == expectedParameter.0 { + #expect( + actualParameter.1 == expectedParameter.1, + "parameter \(actualParameter.0) shape", + sourceLocation: sourceLocation) + } +} diff --git a/Tests/MLXTests/DistributedNNTests.swift b/Tests/MLXTests/DistributedNNTests.swift index 7cae539e0..b57848678 100644 --- a/Tests/MLXTests/DistributedNNTests.swift +++ b/Tests/MLXTests/DistributedNNTests.swift @@ -592,10 +592,6 @@ private func assertGradientWhole( /// which is why this runs in CI without a launcher. class DistributedNNTests: XCTestCase { - override class func setUp() { - setDefaultDevice() - } - override func setUpWithError() throws { try XCTSkipIf( ProcessInfo.processInfo.environment["MLX_TEST_DISTRIBUTED"] == "1", @@ -651,10 +647,6 @@ class DistributedNNRingTests: XCTestCase { /// four, and a shard of 1024 inputs still holds whole quantization groups. static let rankCount = 4 - override class func setUp() { - setDefaultDevice() - } - func testShardedLayers() throws { try DistributedHarness.run(ranks: Self.rankCount, testName: Self.testName) { group in try shardLinearBody(world: group) diff --git a/Tests/MLXTests/DistributedTests.swift b/Tests/MLXTests/DistributedTests.swift index 050c873de..f9defc304 100644 --- a/Tests/MLXTests/DistributedTests.swift +++ b/Tests/MLXTests/DistributedTests.swift @@ -15,10 +15,6 @@ import XCTest /// program with a hostfile, matching how the Python tests are run. class DistributedTests: XCTestCase { - override class func setUp() { - setDefaultDevice() - } - /// The assertions here describe a group of size one, so they are not valid /// in a multi process run -- see ``DistributedRingTests``. override func setUpWithError() throws { diff --git a/Tests/MLXTests/IntegrationTests.swift b/Tests/MLXTests/IntegrationTests.swift deleted file mode 100644 index a8f4a5399..000000000 --- a/Tests/MLXTests/IntegrationTests.swift +++ /dev/null @@ -1,7076 +0,0 @@ -// Copyright © 2024 Apple Inc. - -import Foundation -import MLX -import MLXNN -import XCTest - -@testable import MLXOptimizers - -/// Integration tests comparing results vs known results from python -/// integration. Generated by `tools/generate_integration_tests.py`. -/// -/// Note: this is not meant to be complete coverage, merely a sanity -/// check that the wrapping of the c++ core matches python (e.g. calls -/// the same functions). -class MLXIntegrationTests: XCTestCase { - - override class func setUp() { - } - - func testRandomSeed() { - MLXRandom.seed(864) - let r = MLXRandom.normal() - XCTAssertEqual( - r.item(Float.self), 1.3235496282577515, - accuracy: 0.001) - - } - - func testAddOp() { - MLXRandom.seed(394) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.17949937283992767, - accuracy: -0.0035899874567985536) - XCTAssertEqual( - a.sum().item(Float.self), -2.1539924144744873, - accuracy: -0.04307984828948975) - let b = MLXRandom.normal([4, 3]) - XCTAssertEqual(b.shape, [4, 3]) - XCTAssertEqual(b.dtype, .float32) - XCTAssertEqual( - b.mean().item(Float.self), 0.12489726394414902, - accuracy: 0.00249794527888298) - XCTAssertEqual( - b.sum().item(Float.self), 1.4987671375274658, - accuracy: 0.029975342750549316) - let result = a + b - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.05460209771990776, - accuracy: -0.0010920419543981553) - XCTAssertEqual( - result.sum().item(Float.self), -0.6552251577377319, - accuracy: -0.01310450315475464) - } - - func testAddOp1() { - MLXRandom.seed(776) - let a = 0.5 - let b = MLXRandom.normal([4, 3]) - XCTAssertEqual(b.shape, [4, 3]) - XCTAssertEqual(b.dtype, .float32) - XCTAssertEqual( - b.mean().item(Float.self), 0.1590166687965393, - accuracy: 0.003180333375930786) - XCTAssertEqual( - b.sum().item(Float.self), 1.9082000255584717, - accuracy: 0.03816400051116944) - let result = a + b - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.6590167284011841, - accuracy: 0.013180334568023682) - XCTAssertEqual( - result.sum().item(Float.self), 7.908200263977051, - accuracy: 0.15816400527954103) - } - - func testAddOp2() { - MLXRandom.seed(911) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.523838222026825, - accuracy: 0.0104767644405365) - XCTAssertEqual( - a.sum().item(Float.self), 6.28605842590332, - accuracy: 0.12572116851806642) - let b = 1.3 - let result = a + b - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 1.8238381147384644, - accuracy: 0.03647676229476929) - XCTAssertEqual( - result.sum().item(Float.self), 21.886056900024414, - accuracy: 0.4377211380004883) - } - - func testSubOp() { - MLXRandom.seed(430) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.05095359683036804, - accuracy: 0.0010190719366073608) - XCTAssertEqual( - a.sum().item(Float.self), 0.6114431619644165, - accuracy: 0.01222886323928833) - let b = MLXRandom.normal([4, 3]) - XCTAssertEqual(b.shape, [4, 3]) - XCTAssertEqual(b.dtype, .float32) - XCTAssertEqual( - b.mean().item(Float.self), 0.13153076171875, - accuracy: 0.002630615234375) - XCTAssertEqual( - b.sum().item(Float.self), 1.5783690214157104, - accuracy: 0.03156738042831421) - let result = a - b - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.08057717978954315, - accuracy: -0.001611543595790863) - XCTAssertEqual( - result.sum().item(Float.self), -0.966926097869873, - accuracy: -0.01933852195739746) - } - - func testSubOp1() { - MLXRandom.seed(41) - let a = 0.5 - let b = MLXRandom.normal([4, 3]) - XCTAssertEqual(b.shape, [4, 3]) - XCTAssertEqual(b.dtype, .float32) - XCTAssertEqual( - b.mean().item(Float.self), -0.27060821652412415, - accuracy: -0.005412164330482483) - XCTAssertEqual( - b.sum().item(Float.self), -3.2472984790802, - accuracy: -0.06494596958160401) - let result = a - b - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.7706083059310913, - accuracy: 0.015412166118621827) - XCTAssertEqual( - result.sum().item(Float.self), 9.247299194335938, - accuracy: 0.18494598388671876) - } - - func testSubOp2() { - MLXRandom.seed(265) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.18188762664794922, - accuracy: 0.0036377525329589844) - XCTAssertEqual( - a.sum().item(Float.self), 2.1826515197753906, - accuracy: 0.043653030395507816) - let b = 1.3 - let result = a - b - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -1.118112325668335, - accuracy: -0.0223622465133667) - XCTAssertEqual( - result.sum().item(Float.self), -13.417346954345703, - accuracy: -0.2683469390869141) - } - - func testMulOp() { - MLXRandom.seed(988) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.03418738767504692, - accuracy: 0.0006837477535009385) - XCTAssertEqual( - a.sum().item(Float.self), 0.41024863719940186, - accuracy: 0.008204972743988037) - let b = MLXRandom.normal([4, 3]) - XCTAssertEqual(b.shape, [4, 3]) - XCTAssertEqual(b.dtype, .float32) - XCTAssertEqual( - b.mean().item(Float.self), 0.841292142868042, - accuracy: 0.01682584285736084) - XCTAssertEqual( - b.sum().item(Float.self), 10.095505714416504, - accuracy: 0.20191011428833008) - let result = a * b - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.44235628843307495, - accuracy: 0.0088471257686615) - XCTAssertEqual( - result.sum().item(Float.self), 5.30827522277832, - accuracy: 0.10616550445556641) - } - - func testMulOp1() { - MLXRandom.seed(523) - let a = 0.5 - let b = MLXRandom.normal([4, 3]) - XCTAssertEqual(b.shape, [4, 3]) - XCTAssertEqual(b.dtype, .float32) - XCTAssertEqual( - b.mean().item(Float.self), 0.112852081656456, - accuracy: 0.00225704163312912) - XCTAssertEqual( - b.sum().item(Float.self), 1.3542249202728271, - accuracy: 0.027084498405456542) - let result = a * b - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.056426040828228, - accuracy: 0.00112852081656456) - XCTAssertEqual( - result.sum().item(Float.self), 0.6771124601364136, - accuracy: 0.013542249202728271) - } - - func testMulOp2() { - MLXRandom.seed(497) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.24078607559204102, - accuracy: -0.004815721511840821) - XCTAssertEqual( - a.sum().item(Float.self), -2.889432907104492, - accuracy: -0.057788658142089847) - let b = 1.3 - let result = a * b - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.3130219280719757, - accuracy: -0.006260438561439514) - XCTAssertEqual( - result.sum().item(Float.self), -3.756263017654419, - accuracy: -0.07512526035308838) - } - - func testDivOp() { - MLXRandom.seed(414) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.12723450362682343, - accuracy: -0.0025446900725364686) - XCTAssertEqual( - a.sum().item(Float.self), -1.5268139839172363, - accuracy: -0.030536279678344727) - let b = MLXRandom.normal([4, 3]) - XCTAssertEqual(b.shape, [4, 3]) - XCTAssertEqual(b.dtype, .float32) - XCTAssertEqual( - b.mean().item(Float.self), -0.47808408737182617, - accuracy: -0.009561681747436523) - XCTAssertEqual( - b.sum().item(Float.self), -5.737009048461914, - accuracy: -0.11474018096923828) - let result = a / b - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 2.7393910884857178, - accuracy: 0.054787821769714355) - XCTAssertEqual( - result.sum().item(Float.self), 32.8726921081543, - accuracy: 0.657453842163086) - } - - func testDivOp1() { - MLXRandom.seed(940) - let a = 0.5 - let b = MLXRandom.normal([4, 3]) - XCTAssertEqual(b.shape, [4, 3]) - XCTAssertEqual(b.dtype, .float32) - XCTAssertEqual( - b.mean().item(Float.self), -0.3181919455528259, - accuracy: -0.0063638389110565186) - XCTAssertEqual( - b.sum().item(Float.self), -3.818303346633911, - accuracy: -0.07636606693267822) - let result = a / b - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.33945927023887634, - accuracy: -0.006789185404777527) - XCTAssertEqual( - result.sum().item(Float.self), -4.073511123657227, - accuracy: -0.08147022247314453) - } - - func testDivOp2() { - MLXRandom.seed(802) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.08001989126205444, - accuracy: 0.001600397825241089) - XCTAssertEqual( - a.sum().item(Float.self), 0.9602386951446533, - accuracy: 0.019204773902893067) - let b = 1.3 - let result = a / b - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.06155376881361008, - accuracy: 0.0012310753762722016) - XCTAssertEqual( - result.sum().item(Float.self), 0.7386451959609985, - accuracy: 0.014772903919219971) - } - - func testModOp() { - MLXRandom.seed(849) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.5959362983703613, - accuracy: -0.011918725967407227) - XCTAssertEqual( - a.sum().item(Float.self), -7.151235103607178, - accuracy: -0.14302470207214354) - let b = MLXRandom.normal([4, 3]) - XCTAssertEqual(b.shape, [4, 3]) - XCTAssertEqual(b.dtype, .float32) - XCTAssertEqual( - b.mean().item(Float.self), 0.42154282331466675, - accuracy: 0.008430856466293336) - XCTAssertEqual( - b.sum().item(Float.self), 5.058513641357422, - accuracy: 0.10117027282714844) - let result = a % b - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.28858500719070435, - accuracy: 0.005771700143814087) - XCTAssertEqual( - result.sum().item(Float.self), 3.463019847869873, - accuracy: 0.06926039695739747) - } - - func testModOp1() { - MLXRandom.seed(310) - let a = 0.5 - let b = MLXRandom.normal([4, 3]) - XCTAssertEqual(b.shape, [4, 3]) - XCTAssertEqual(b.dtype, .float32) - XCTAssertEqual( - b.mean().item(Float.self), 0.15466859936714172, - accuracy: 0.0030933719873428344) - XCTAssertEqual( - b.sum().item(Float.self), 1.8560230731964111, - accuracy: 0.03712046146392822) - let result = a % b - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.029992982745170593, - accuracy: -0.0005998596549034119) - XCTAssertEqual( - result.sum().item(Float.self), -0.3599157929420471, - accuracy: -0.007198315858840942) - } - - func testModOp2() { - MLXRandom.seed(991) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.34309279918670654, - accuracy: -0.006861855983734131) - XCTAssertEqual( - a.sum().item(Float.self), -4.1171135902404785, - accuracy: -0.08234227180480957) - let b = 1.3 - let result = a % b - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.6319072246551514, - accuracy: 0.012638144493103028) - XCTAssertEqual( - result.sum().item(Float.self), 7.582886219024658, - accuracy: 0.15165772438049316) - } - - func testPowOp() { - MLXRandom.seed(488) - let a = MLXRandom.uniform(low: 0.1, high: 2.0, [4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.8261654376983643, - accuracy: 0.016523308753967285) - XCTAssertEqual( - a.sum().item(Float.self), 9.913985252380371, - accuracy: 0.19827970504760742) - let b = MLXRandom.uniform(low: 0.1, high: 2.0, [4, 3]) - XCTAssertEqual(b.shape, [4, 3]) - XCTAssertEqual(b.dtype, .float32) - XCTAssertEqual( - b.mean().item(Float.self), 1.3262829780578613, - accuracy: 0.02652565956115723) - XCTAssertEqual( - b.sum().item(Float.self), 15.915395736694336, - accuracy: 0.3183079147338867) - let result = a ** b - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.8045400381088257, - accuracy: 0.016090800762176515) - XCTAssertEqual( - result.sum().item(Float.self), 9.65447998046875, - accuracy: 0.193089599609375) - } - - func testPowOp1() { - MLXRandom.seed(366) - let a = MLXRandom.uniform(low: 0.1, high: 2.0, [4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.9292365908622742, - accuracy: 0.018584731817245483) - XCTAssertEqual( - a.sum().item(Float.self), 11.150838851928711, - accuracy: 0.22301677703857423) - let b = 1.3 - let result = a ** b - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.9557135105133057, - accuracy: 0.019114270210266113) - XCTAssertEqual( - result.sum().item(Float.self), 11.468562126159668, - accuracy: 0.22937124252319335) - } - - func testEqualOp() { - MLXRandom.seed(597) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.268618643283844, - accuracy: -0.00537237286567688) - XCTAssertEqual( - a.sum().item(Float.self), -3.223423480987549, - accuracy: -0.06446846961975097) - let b = MLXRandom.normal([4, 3]) - XCTAssertEqual(b.shape, [4, 3]) - XCTAssertEqual(b.dtype, .float32) - XCTAssertEqual( - b.mean().item(Float.self), -0.002711281180381775, - accuracy: -5.42256236076355e-05) - XCTAssertEqual( - b.sum().item(Float.self), -0.0325353741645813, - accuracy: -0.000650707483291626) - let result = a .== b - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .bool) - XCTAssertEqual(result.all().item(), false) - XCTAssertEqual(result.any().item(), false) - } - - func testEqualOp1() { - MLXRandom.seed(913) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.040915559977293015, - accuracy: 0.0008183111995458603) - XCTAssertEqual( - a.sum().item(Float.self), 0.490986704826355, - accuracy: 0.0098197340965271) - let b = 1.3 - let result = a .== b - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .bool) - XCTAssertEqual(result.all().item(), false) - XCTAssertEqual(result.any().item(), false) - } - - func testNotEqualOp() { - MLXRandom.seed(929) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.2947184443473816, - accuracy: -0.005894368886947632) - XCTAssertEqual( - a.sum().item(Float.self), -3.536621332168579, - accuracy: -0.07073242664337158) - let b = MLXRandom.normal([4, 3]) - XCTAssertEqual(b.shape, [4, 3]) - XCTAssertEqual(b.dtype, .float32) - XCTAssertEqual( - b.mean().item(Float.self), -0.009182363748550415, - accuracy: -0.0001836472749710083) - XCTAssertEqual( - b.sum().item(Float.self), -0.11018836498260498, - accuracy: -0.0022037672996520997) - let result = a .!= b - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .bool) - XCTAssertEqual(result.all().item(), true) - XCTAssertEqual(result.any().item(), true) - } - - func testNotEqualOp1() { - MLXRandom.seed(223) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.009602058678865433, - accuracy: -0.00019204117357730867) - XCTAssertEqual( - a.sum().item(Float.self), -0.11522470414638519, - accuracy: -0.002304494082927704) - let b = 1.3 - let result = a .!= b - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .bool) - XCTAssertEqual(result.all().item(), true) - XCTAssertEqual(result.any().item(), true) - } - - func testLessThanOp() { - MLXRandom.seed(516) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.671758770942688, - accuracy: 0.01343517541885376) - XCTAssertEqual( - a.sum().item(Float.self), 8.061104774475098, - accuracy: 0.16122209548950195) - let b = MLXRandom.normal([4, 3]) - XCTAssertEqual(b.shape, [4, 3]) - XCTAssertEqual(b.dtype, .float32) - XCTAssertEqual( - b.mean().item(Float.self), -0.04753169044852257, - accuracy: -0.0009506338089704514) - XCTAssertEqual( - b.sum().item(Float.self), -0.5703802704811096, - accuracy: -0.011407605409622193) - let result = a .< b - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .bool) - XCTAssertEqual(result.all().item(), false) - XCTAssertEqual(result.any().item(), true) - } - - func testLessThanOp1() { - MLXRandom.seed(142) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.21625718474388123, - accuracy: 0.004325143694877624) - XCTAssertEqual( - a.sum().item(Float.self), 2.595086097717285, - accuracy: 0.051901721954345705) - let b = 1.3 - let result = a .< b - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .bool) - XCTAssertEqual(result.all().item(), false) - XCTAssertEqual(result.any().item(), true) - } - - func testLessThanEqualOp() { - MLXRandom.seed(288) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.19039803743362427, - accuracy: -0.0038079607486724855) - XCTAssertEqual( - a.sum().item(Float.self), -2.284776449203491, - accuracy: -0.045695528984069825) - let b = MLXRandom.normal([4, 3]) - XCTAssertEqual(b.shape, [4, 3]) - XCTAssertEqual(b.dtype, .float32) - XCTAssertEqual( - b.mean().item(Float.self), 0.10338123887777328, - accuracy: 0.002067624777555466) - XCTAssertEqual( - b.sum().item(Float.self), 1.240574836730957, - accuracy: 0.024811496734619142) - let result = a .<= b - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .bool) - XCTAssertEqual(result.all().item(), false) - XCTAssertEqual(result.any().item(), true) - } - - func testLessThanEqualOp1() { - MLXRandom.seed(143) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.025643646717071533, - accuracy: -0.0005128729343414307) - XCTAssertEqual( - a.sum().item(Float.self), -0.3077237606048584, - accuracy: -0.006154475212097168) - let b = 1.3 - let result = a .<= b - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .bool) - XCTAssertEqual(result.all().item(), true) - XCTAssertEqual(result.any().item(), true) - } - - func testGreaterThanOp() { - MLXRandom.seed(773) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.18657740950584412, - accuracy: -0.003731548190116882) - XCTAssertEqual( - a.sum().item(Float.self), -2.23892879486084, - accuracy: -0.044778575897216795) - let b = MLXRandom.normal([4, 3]) - XCTAssertEqual(b.shape, [4, 3]) - XCTAssertEqual(b.dtype, .float32) - XCTAssertEqual( - b.mean().item(Float.self), 0.0327937975525856, - accuracy: 0.0006558759510517121) - XCTAssertEqual( - b.sum().item(Float.self), 0.39352554082870483, - accuracy: 0.007870510816574097) - let result = a .> b - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .bool) - XCTAssertEqual(result.all().item(), false) - XCTAssertEqual(result.any().item(), true) - } - - func testGreaterThanOp1() { - MLXRandom.seed(97) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.04777970910072327, - accuracy: -0.0009555941820144654) - XCTAssertEqual( - a.sum().item(Float.self), -0.5733565092086792, - accuracy: -0.011467130184173583) - let b = 1.3 - let result = a .> b - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .bool) - XCTAssertEqual(result.all().item(), false) - XCTAssertEqual(result.any().item(), false) - } - - func testGreaterThanEqualOp() { - MLXRandom.seed(633) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.17345473170280457, - accuracy: 0.0034690946340560913) - XCTAssertEqual( - a.sum().item(Float.self), 2.0814566612243652, - accuracy: 0.0416291332244873) - let b = MLXRandom.normal([4, 3]) - XCTAssertEqual(b.shape, [4, 3]) - XCTAssertEqual(b.dtype, .float32) - XCTAssertEqual( - b.mean().item(Float.self), 0.2481841742992401, - accuracy: 0.004963683485984803) - XCTAssertEqual( - b.sum().item(Float.self), 2.978209972381592, - accuracy: 0.05956419944763184) - let result = a .>= b - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .bool) - XCTAssertEqual(result.all().item(), false) - XCTAssertEqual(result.any().item(), true) - } - - func testGreaterThanEqualOp1() { - MLXRandom.seed(818) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.01743374764919281, - accuracy: -0.0003486749529838562) - XCTAssertEqual( - a.sum().item(Float.self), -0.20920497179031372, - accuracy: -0.004184099435806275) - let b = 1.3 - let result = a .>= b - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .bool) - XCTAssertEqual(result.all().item(), false) - XCTAssertEqual(result.any().item(), false) - } - - func testAbs() { - MLXRandom.seed(256) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.15743832290172577, - accuracy: -0.0031487664580345156) - XCTAssertEqual( - a.sum().item(Float.self), -1.8892598152160645, - accuracy: -0.03778519630432129) - let result = a.abs() - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.8130466341972351, - accuracy: 0.0162609326839447) - XCTAssertEqual( - result.sum().item(Float.self), 9.756559371948242, - accuracy: 0.19513118743896485) - } - - func testAbs1() { - MLXRandom.seed(931) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.22450898587703705, - accuracy: -0.004490179717540741) - XCTAssertEqual( - a.sum().item(Float.self), -2.6941077709198, - accuracy: -0.053882155418396) - let result = abs(a) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.6981315016746521, - accuracy: 0.013962630033493042) - XCTAssertEqual( - result.sum().item(Float.self), 8.377577781677246, - accuracy: 0.16755155563354493) - } - - func testAll() { - MLXRandom.seed(545) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.05837365239858627, - accuracy: 0.0011674730479717256) - XCTAssertEqual( - a.sum().item(Float.self), 0.7004837989807129, - accuracy: 0.014009675979614259) - let result = a.all() - XCTAssertEqual(result.shape, []) - XCTAssertEqual(result.dtype, .bool) - XCTAssertEqual(result.all().item(), true) - XCTAssertEqual(result.any().item(), true) - } - - func testAll1() { - MLXRandom.seed(722) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.01978858932852745, - accuracy: -0.00039577178657054904) - XCTAssertEqual( - a.sum().item(Float.self), -0.2374630719423294, - accuracy: -0.004749261438846588) - let result = all(a) - XCTAssertEqual(result.shape, []) - XCTAssertEqual(result.dtype, .bool) - XCTAssertEqual(result.all().item(), true) - XCTAssertEqual(result.any().item(), true) - } - - func testAll2() { - MLXRandom.seed(829) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.03525291010737419, - accuracy: 0.0007050582021474838) - XCTAssertEqual( - a.sum().item(Float.self), 0.4230349063873291, - accuracy: 0.008460698127746582) - let result = a.all(axis: -1) - XCTAssertEqual(result.shape, [4]) - XCTAssertEqual(result.dtype, .bool) - XCTAssertEqual(result.all().item(), true) - XCTAssertEqual(result.any().item(), true) - } - - func testAll3() { - MLXRandom.seed(616) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.1122930571436882, - accuracy: 0.002245861142873764) - XCTAssertEqual( - a.sum().item(Float.self), 1.347516655921936, - accuracy: 0.02695033311843872) - let result = all(a, axis: -1) - XCTAssertEqual(result.shape, [4]) - XCTAssertEqual(result.dtype, .bool) - XCTAssertEqual(result.all().item(), true) - XCTAssertEqual(result.any().item(), true) - } - - func testAll4() { - MLXRandom.seed(923) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 3, 4, 3]) - XCTAssertEqual(a.shape, [2, 3, 4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.4720103442668915, - accuracy: 0.00944020688533783) - XCTAssertEqual( - a.sum().item(Float.self), 33.984745025634766, - accuracy: 0.6796949005126953) - let result = a.all(axes: [0, -1]) - XCTAssertEqual(result.shape, [3, 4]) - XCTAssertEqual(result.dtype, .bool) - XCTAssertEqual(result.all().item(), true) - XCTAssertEqual(result.any().item(), true) - } - - func testAll5() { - MLXRandom.seed(150) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 3, 4, 3]) - XCTAssertEqual(a.shape, [2, 3, 4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.4883745014667511, - accuracy: 0.009767490029335022) - XCTAssertEqual( - a.sum().item(Float.self), 35.1629638671875, - accuracy: 0.70325927734375) - let result = all(a, axes: [0, -1]) - XCTAssertEqual(result.shape, [3, 4]) - XCTAssertEqual(result.dtype, .bool) - XCTAssertEqual(result.all().item(), true) - XCTAssertEqual(result.any().item(), true) - } - - func testAny() { - MLXRandom.seed(317) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.1289883852005005, - accuracy: -0.0025797677040100097) - XCTAssertEqual( - a.sum().item(Float.self), -1.5478605031967163, - accuracy: -0.030957210063934325) - let result = a.any() - XCTAssertEqual(result.shape, []) - XCTAssertEqual(result.dtype, .bool) - XCTAssertEqual(result.all().item(), true) - XCTAssertEqual(result.any().item(), true) - } - - func testAny1() { - MLXRandom.seed(101) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.08150328695774078, - accuracy: -0.0016300657391548157) - XCTAssertEqual( - a.sum().item(Float.self), -0.9780394434928894, - accuracy: -0.01956078886985779) - let result = any(a) - XCTAssertEqual(result.shape, []) - XCTAssertEqual(result.dtype, .bool) - XCTAssertEqual(result.all().item(), true) - XCTAssertEqual(result.any().item(), true) - } - - func testAny2() { - MLXRandom.seed(747) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.23745165765285492, - accuracy: 0.004749033153057099) - XCTAssertEqual( - a.sum().item(Float.self), 2.8494198322296143, - accuracy: 0.056988396644592286) - let result = a.any(axis: -1) - XCTAssertEqual(result.shape, [4]) - XCTAssertEqual(result.dtype, .bool) - XCTAssertEqual(result.all().item(), true) - XCTAssertEqual(result.any().item(), true) - } - - func testAny3() { - MLXRandom.seed(75) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.3039236068725586, - accuracy: -0.006078472137451172) - XCTAssertEqual( - a.sum().item(Float.self), -3.647083282470703, - accuracy: -0.07294166564941407) - let result = any(a, axis: -1) - XCTAssertEqual(result.shape, [4]) - XCTAssertEqual(result.dtype, .bool) - XCTAssertEqual(result.all().item(), true) - XCTAssertEqual(result.any().item(), true) - } - - func testAny4() { - MLXRandom.seed(920) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 3, 4, 3]) - XCTAssertEqual(a.shape, [2, 3, 4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.41749411821365356, - accuracy: 0.008349882364273071) - XCTAssertEqual( - a.sum().item(Float.self), 30.0595760345459, - accuracy: 0.601191520690918) - let result = a.any(axes: [0, -1]) - XCTAssertEqual(result.shape, [3, 4]) - XCTAssertEqual(result.dtype, .bool) - XCTAssertEqual(result.all().item(), true) - XCTAssertEqual(result.any().item(), true) - } - - func testAny5() { - MLXRandom.seed(870) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 3, 4, 3]) - XCTAssertEqual(a.shape, [2, 3, 4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.5518582463264465, - accuracy: 0.01103716492652893) - XCTAssertEqual( - a.sum().item(Float.self), 39.73379135131836, - accuracy: 0.7946758270263672) - let result = any(a, axes: [0, -1]) - XCTAssertEqual(result.shape, [3, 4]) - XCTAssertEqual(result.dtype, .bool) - XCTAssertEqual(result.all().item(), true) - XCTAssertEqual(result.any().item(), true) - } - - func testArgMax() { - MLXRandom.seed(700) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.20903855562210083, - accuracy: -0.004180771112442016) - XCTAssertEqual( - a.sum().item(Float.self), -2.50846266746521, - accuracy: -0.0501692533493042) - let result = a.argMax() - XCTAssertEqual(result.shape, []) - XCTAssertEqual(result.dtype, .uint32) - XCTAssertEqual( - result.mean().item(Float.self), 8.0, - accuracy: 0.16) - XCTAssertEqual( - result.sum().item(Float.self), 8, - accuracy: 0.16) - } - - func testArgMax1() { - MLXRandom.seed(338) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.07185819745063782, - accuracy: 0.0014371639490127564) - XCTAssertEqual( - a.sum().item(Float.self), 0.8622983694076538, - accuracy: 0.017245967388153077) - let result = argMax(a) - XCTAssertEqual(result.shape, []) - XCTAssertEqual(result.dtype, .uint32) - XCTAssertEqual( - result.mean().item(Float.self), 7.0, - accuracy: 0.14) - XCTAssertEqual( - result.sum().item(Float.self), 7, - accuracy: 0.14) - } - - func testArgMax2() { - MLXRandom.seed(483) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.0843842476606369, - accuracy: 0.001687684953212738) - XCTAssertEqual( - a.sum().item(Float.self), 1.012610912322998, - accuracy: 0.020252218246459962) - let result = a.argMax(axis: -1) - XCTAssertEqual(result.shape, [4]) - XCTAssertEqual(result.dtype, .uint32) - XCTAssertEqual( - result.mean().item(Float.self), 0.75, - accuracy: 0.015) - XCTAssertEqual( - result.sum().item(Float.self), 3, - accuracy: 0.06) - } - - func testArgMax3() { - MLXRandom.seed(573) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.5405035018920898, - accuracy: 0.010810070037841797) - XCTAssertEqual( - a.sum().item(Float.self), 6.48604154586792, - accuracy: 0.1297208309173584) - let result = argMax(a, axis: -1) - XCTAssertEqual(result.shape, [4]) - XCTAssertEqual(result.dtype, .uint32) - XCTAssertEqual( - result.mean().item(Float.self), 1.25, - accuracy: 0.025) - XCTAssertEqual( - result.sum().item(Float.self), 5, - accuracy: 0.1) - } - - func testArgMin() { - MLXRandom.seed(103) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.2983229458332062, - accuracy: 0.005966458916664124) - XCTAssertEqual( - a.sum().item(Float.self), 3.5798752307891846, - accuracy: 0.0715975046157837) - let result = a.argMin() - XCTAssertEqual(result.shape, []) - XCTAssertEqual(result.dtype, .uint32) - XCTAssertEqual( - result.mean().item(Float.self), 11.0, - accuracy: 0.22) - XCTAssertEqual( - result.sum().item(Float.self), 11, - accuracy: 0.22) - } - - func testArgMin1() { - MLXRandom.seed(362) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.057320237159729004, - accuracy: -0.00114640474319458) - XCTAssertEqual( - a.sum().item(Float.self), -0.687842845916748, - accuracy: -0.013756856918334961) - let result = argMin(a) - XCTAssertEqual(result.shape, []) - XCTAssertEqual(result.dtype, .uint32) - XCTAssertEqual( - result.mean().item(Float.self), 8.0, - accuracy: 0.16) - XCTAssertEqual( - result.sum().item(Float.self), 8, - accuracy: 0.16) - } - - func testArgMin2() { - MLXRandom.seed(444) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.22760924696922302, - accuracy: 0.004552184939384461) - XCTAssertEqual( - a.sum().item(Float.self), 2.7313108444213867, - accuracy: 0.05462621688842773) - let result = a.argMin(axis: -1) - XCTAssertEqual(result.shape, [4]) - XCTAssertEqual(result.dtype, .uint32) - XCTAssertEqual( - result.mean().item(Float.self), 0.25, - accuracy: 0.005) - XCTAssertEqual( - result.sum().item(Float.self), 1, - accuracy: 0.02) - } - - func testArgMin3() { - MLXRandom.seed(323) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.4485395550727844, - accuracy: -0.008970791101455688) - XCTAssertEqual( - a.sum().item(Float.self), -5.382474422454834, - accuracy: -0.10764948844909668) - let result = argMin(a, axis: -1) - XCTAssertEqual(result.shape, [4]) - XCTAssertEqual(result.dtype, .uint32) - XCTAssertEqual( - result.mean().item(Float.self), 0.5, - accuracy: 0.01) - XCTAssertEqual( - result.sum().item(Float.self), 2, - accuracy: 0.04) - } - - func testCummax() { - MLXRandom.seed(625) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.2803424596786499, - accuracy: -0.0056068491935729985) - XCTAssertEqual( - a.sum().item(Float.self), -3.3641092777252197, - accuracy: -0.0672821855545044) - let result = a.cummax() - XCTAssertEqual(result.shape, [12]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.3141178786754608, - accuracy: 0.006282357573509216) - XCTAssertEqual( - result.sum().item(Float.self), 3.7694144248962402, - accuracy: 0.07538828849792481) - } - - func testCummax1() { - MLXRandom.seed(655) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.37031763792037964, - accuracy: 0.0074063527584075925) - XCTAssertEqual( - a.sum().item(Float.self), 4.443811416625977, - accuracy: 0.08887622833251953) - let result = cummax(a) - XCTAssertEqual(result.shape, [12]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 1.3793113231658936, - accuracy: 0.027586226463317872) - XCTAssertEqual( - result.sum().item(Float.self), 16.551734924316406, - accuracy: 0.3310346984863281) - } - - func testCummax2() { - MLXRandom.seed(934) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.05676526948809624, - accuracy: -0.0011353053897619249) - XCTAssertEqual( - a.sum().item(Float.self), -0.6811832189559937, - accuracy: -0.013623664379119873) - let result = a.cummax(axis: -1) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.018919458612799644, - accuracy: 0.0003783891722559929) - XCTAssertEqual( - result.sum().item(Float.self), 0.22703349590301514, - accuracy: 0.004540669918060303) - } - - func testCummax3() { - MLXRandom.seed(209) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.07204179465770721, - accuracy: 0.0014408358931541443) - XCTAssertEqual( - a.sum().item(Float.self), 0.8645014762878418, - accuracy: 0.017290029525756836) - let result = cummax(a, axis: -1) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 1.061327576637268, - accuracy: 0.02122655153274536) - XCTAssertEqual( - result.sum().item(Float.self), 12.735930442810059, - accuracy: 0.2547186088562012) - } - - func testCummin() { - MLXRandom.seed(989) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.16920256614685059, - accuracy: -0.0033840513229370117) - XCTAssertEqual( - a.sum().item(Float.self), -2.030430793762207, - accuracy: -0.04060861587524414) - let result = a.cummin() - XCTAssertEqual(result.shape, [12]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.9566025733947754, - accuracy: -0.019132051467895508) - XCTAssertEqual( - result.sum().item(Float.self), -11.479230880737305, - accuracy: -0.2295846176147461) - } - - func testCummin1() { - MLXRandom.seed(565) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.644216775894165, - accuracy: -0.0128843355178833) - XCTAssertEqual( - a.sum().item(Float.self), -7.7306013107299805, - accuracy: -0.1546120262145996) - let result = cummin(a) - XCTAssertEqual(result.shape, [12]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -1.7292345762252808, - accuracy: -0.03458469152450561) - XCTAssertEqual( - result.sum().item(Float.self), -20.75081443786621, - accuracy: -0.4150162887573242) - } - - func testCummin2() { - MLXRandom.seed(488) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.398279070854187, - accuracy: -0.00796558141708374) - XCTAssertEqual( - a.sum().item(Float.self), -4.779348850250244, - accuracy: -0.09558697700500489) - let result = a.cummin(axis: -1) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.9816936254501343, - accuracy: -0.019633872509002687) - XCTAssertEqual( - result.sum().item(Float.self), -11.780323028564453, - accuracy: -0.23560646057128906) - } - - func testCummin3() { - MLXRandom.seed(453) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.0004353722033556551, - accuracy: 8.707444067113101e-06) - XCTAssertEqual( - a.sum().item(Float.self), 0.005224466323852539, - accuracy: 0.00010448932647705078) - let result = cummin(a, axis: -1) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.4506932497024536, - accuracy: -0.009013864994049072) - XCTAssertEqual( - result.sum().item(Float.self), -5.408318996429443, - accuracy: -0.10816637992858887) - } - - func testCumprod() { - MLXRandom.seed(886) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.0051885247230529785, - accuracy: -0.00010377049446105957) - XCTAssertEqual( - a.sum().item(Float.self), -0.06226229667663574, - accuracy: -0.001245245933532715) - let result = a.cumprod() - XCTAssertEqual(result.shape, [12]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.14253053069114685, - accuracy: 0.002850610613822937) - XCTAssertEqual( - result.sum().item(Float.self), 1.7103662490844727, - accuracy: 0.034207324981689456) - } - - func testCumprod1() { - MLXRandom.seed(533) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.12974336743354797, - accuracy: -0.0025948673486709596) - XCTAssertEqual( - a.sum().item(Float.self), -1.5569202899932861, - accuracy: -0.031138405799865723) - let result = cumprod(a) - XCTAssertEqual(result.shape, [12]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.020257528871297836, - accuracy: 0.0004051505774259567) - XCTAssertEqual( - result.sum().item(Float.self), 0.24309033155441284, - accuracy: 0.004861806631088257) - } - - func testCumprod2() { - MLXRandom.seed(266) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.20902606844902039, - accuracy: 0.004180521368980407) - XCTAssertEqual( - a.sum().item(Float.self), 2.508312702178955, - accuracy: 0.050166254043579106) - let result = a.cumprod(axis: -1) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.7459201812744141, - accuracy: 0.014918403625488281) - XCTAssertEqual( - result.sum().item(Float.self), 8.951042175292969, - accuracy: 0.17902084350585937) - } - - func testCumprod3() { - MLXRandom.seed(63) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.20442306995391846, - accuracy: 0.004088461399078369) - XCTAssertEqual( - a.sum().item(Float.self), 2.4530768394470215, - accuracy: 0.04906153678894043) - let result = cumprod(a, axis: -1) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.20277641713619232, - accuracy: 0.004055528342723847) - XCTAssertEqual( - result.sum().item(Float.self), 2.433316946029663, - accuracy: 0.048666338920593265) - } - - func testCumsum() { - MLXRandom.seed(824) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.07800808548927307, - accuracy: 0.0015601617097854615) - XCTAssertEqual( - a.sum().item(Float.self), 0.9360969662666321, - accuracy: 0.018721939325332643) - let result = a.cumsum() - XCTAssertEqual(result.shape, [12]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.16304373741149902, - accuracy: -0.0032608747482299806) - XCTAssertEqual( - result.sum().item(Float.self), -1.9565248489379883, - accuracy: -0.03913049697875977) - } - - func testCumsum1() { - MLXRandom.seed(940) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.3181919455528259, - accuracy: -0.0063638389110565186) - XCTAssertEqual( - a.sum().item(Float.self), -3.818303346633911, - accuracy: -0.07636606693267822) - let result = cumsum(a) - XCTAssertEqual(result.shape, [12]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -3.2638518810272217, - accuracy: -0.06527703762054443) - XCTAssertEqual( - result.sum().item(Float.self), -39.166221618652344, - accuracy: -0.7833244323730469) - } - - func testCumsum2() { - MLXRandom.seed(561) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.19832412898540497, - accuracy: 0.0039664825797081) - XCTAssertEqual( - a.sum().item(Float.self), 2.379889488220215, - accuracy: 0.0475977897644043) - let result = a.cumsum(axis: -1) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.7345349788665771, - accuracy: 0.014690699577331543) - XCTAssertEqual( - result.sum().item(Float.self), 8.814419746398926, - accuracy: 0.17628839492797851) - } - - func testCumsum3() { - MLXRandom.seed(937) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.38423240184783936, - accuracy: -0.007684648036956787) - XCTAssertEqual( - a.sum().item(Float.self), -4.610788822174072, - accuracy: -0.09221577644348145) - let result = cumsum(a, axis: -1) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.6559594869613647, - accuracy: -0.013119189739227296) - XCTAssertEqual( - result.sum().item(Float.self), -7.871513366699219, - accuracy: -0.15743026733398438) - } - - func testExpandedDimensions() { - MLXRandom.seed(14) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.20228910446166992, - accuracy: -0.004045782089233399) - XCTAssertEqual( - a.sum().item(Float.self), -2.427469253540039, - accuracy: -0.04854938507080078) - let result = expandedDimensions(a, axis: -1) - XCTAssertEqual(result.shape, [4, 3, 1]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.20228910446166992, - accuracy: -0.004045782089233399) - XCTAssertEqual( - result.sum().item(Float.self), -2.427469253540039, - accuracy: -0.04854938507080078) - } - - func testExpandedDimensions1() { - MLXRandom.seed(95) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 3, 4, 3]) - XCTAssertEqual(a.shape, [2, 3, 4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.5254887342453003, - accuracy: 0.010509774684906006) - XCTAssertEqual( - a.sum().item(Float.self), 37.83518981933594, - accuracy: 0.7567037963867188) - let result = expandedDimensions(a, axes: [0, -1]) - XCTAssertEqual(result.shape, [1, 2, 3, 4, 3, 1]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.5254887342453003, - accuracy: 0.010509774684906006) - XCTAssertEqual( - result.sum().item(Float.self), 37.83518981933594, - accuracy: 0.7567037963867188) - } - - func testFloor() { - MLXRandom.seed(736) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.035255465656518936, - accuracy: -0.0007051093131303787) - XCTAssertEqual( - a.sum().item(Float.self), -0.42306557297706604, - accuracy: -0.00846131145954132) - let result = floor(a) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.4166666865348816, - accuracy: -0.008333333730697633) - XCTAssertEqual( - result.sum().item(Float.self), -5.0, - accuracy: -0.1) - } - - func testLog() { - MLXRandom.seed(860) - let a = MLXRandom.uniform(low: 0.1, high: 2.0, [4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.9451824426651001, - accuracy: 0.018903648853302004) - XCTAssertEqual( - a.sum().item(Float.self), 11.342188835144043, - accuracy: 0.22684377670288086) - let result = a.log() - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.26455360651016235, - accuracy: -0.0052910721302032475) - XCTAssertEqual( - result.sum().item(Float.self), -3.174643039703369, - accuracy: -0.06349286079406738) - } - - func testLog1() { - MLXRandom.seed(408) - let a = MLXRandom.uniform(low: 0.1, high: 2.0, [4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 1.1787240505218506, - accuracy: 0.023574481010437014) - XCTAssertEqual( - a.sum().item(Float.self), 14.144688606262207, - accuracy: 0.28289377212524414) - let result = log(a) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.08071314543485641, - accuracy: 0.0016142629086971284) - XCTAssertEqual( - result.sum().item(Float.self), 0.9685577154159546, - accuracy: 0.01937115430831909) - } - - func testLog2() { - MLXRandom.seed(727) - let a = MLXRandom.uniform(low: 0.1, high: 2.0, [4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 1.0156004428863525, - accuracy: 0.020312008857727052) - XCTAssertEqual( - a.sum().item(Float.self), 12.18720531463623, - accuracy: 0.2437441062927246) - let result = a.log2() - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.2411154955625534, - accuracy: -0.004822309911251068) - XCTAssertEqual( - result.sum().item(Float.self), -2.893385887145996, - accuracy: -0.057867717742919926) - } - - func testLog21() { - MLXRandom.seed(844) - let a = MLXRandom.uniform(low: 0.1, high: 2.0, [4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 1.0754908323287964, - accuracy: 0.02150981664657593) - XCTAssertEqual( - a.sum().item(Float.self), 12.905889511108398, - accuracy: 0.25811779022216796) - let result = log2(a) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.04823940992355347, - accuracy: -0.0009647881984710694) - XCTAssertEqual( - result.sum().item(Float.self), -0.5788729190826416, - accuracy: -0.011577458381652831) - } - - func testLog10() { - MLXRandom.seed(803) - let a = MLXRandom.uniform(low: 0.1, high: 2.0, [4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 1.1359423398971558, - accuracy: 0.022718846797943115) - XCTAssertEqual( - a.sum().item(Float.self), 13.631307601928711, - accuracy: 0.2726261520385742) - let result = a.log10() - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.07397204637527466, - accuracy: -0.0014794409275054933) - XCTAssertEqual( - result.sum().item(Float.self), -0.8876644968986511, - accuracy: -0.017753289937973024) - } - - func testLog101() { - MLXRandom.seed(684) - let a = MLXRandom.uniform(low: 0.1, high: 2.0, [4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 1.3902796506881714, - accuracy: 0.027805593013763428) - XCTAssertEqual( - a.sum().item(Float.self), 16.6833553314209, - accuracy: 0.333667106628418) - let result = log10(a) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.11499349772930145, - accuracy: 0.002299869954586029) - XCTAssertEqual( - result.sum().item(Float.self), 1.3799219131469727, - accuracy: 0.027598438262939454) - } - - func testLog1p() { - MLXRandom.seed(640) - let a = MLXRandom.uniform(low: 0.1, high: 2.0, [4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 1.2758265733718872, - accuracy: 0.025516531467437743) - XCTAssertEqual( - a.sum().item(Float.self), 15.309918403625488, - accuracy: 0.3061983680725098) - let result = a.log1p() - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.7979601621627808, - accuracy: 0.015959203243255615) - XCTAssertEqual( - result.sum().item(Float.self), 9.575521469116211, - accuracy: 0.1915104293823242) - } - - func testLog1p1() { - MLXRandom.seed(1) - let a = MLXRandom.uniform(low: 0.1, high: 2.0, [4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 1.045818567276001, - accuracy: 0.02091637134552002) - XCTAssertEqual( - a.sum().item(Float.self), 12.549821853637695, - accuracy: 0.25099643707275393) - let result = log1p(a) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.686416506767273, - accuracy: 0.01372833013534546) - XCTAssertEqual( - result.sum().item(Float.self), 8.236997604370117, - accuracy: 0.16473995208740236) - } - - func testLogSumExp() { - MLXRandom.seed(626) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.09981077909469604, - accuracy: -0.001996215581893921) - XCTAssertEqual( - a.sum().item(Float.self), -1.1977293491363525, - accuracy: -0.023954586982727052) - let result = a.logSumExp() - XCTAssertEqual(result.shape, []) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 2.6805179119110107, - accuracy: 0.05361035823822022) - XCTAssertEqual( - result.sum().item(Float.self), 2.6805179119110107, - accuracy: 0.05361035823822022) - } - - func testLogSumExp1() { - MLXRandom.seed(505) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.16046583652496338, - accuracy: 0.0032093167304992677) - XCTAssertEqual( - a.sum().item(Float.self), 1.9255900382995605, - accuracy: 0.038511800765991214) - let result = logSumExp(a) - XCTAssertEqual(result.shape, []) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 3.126417636871338, - accuracy: 0.06252835273742675) - XCTAssertEqual( - result.sum().item(Float.self), 3.126417636871338, - accuracy: 0.06252835273742675) - } - - func testLogSumExp2() { - MLXRandom.seed(847) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.2282017469406128, - accuracy: -0.004564034938812256) - XCTAssertEqual( - a.sum().item(Float.self), -2.7384209632873535, - accuracy: -0.05476841926574707) - let result = a.logSumExp(axis: -1) - XCTAssertEqual(result.shape, [4]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 1.163524866104126, - accuracy: 0.02327049732208252) - XCTAssertEqual( - result.sum().item(Float.self), 4.654099464416504, - accuracy: 0.09308198928833009) - } - - func testLogSumExp3() { - MLXRandom.seed(888) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.05087994784116745, - accuracy: 0.001017598956823349) - XCTAssertEqual( - a.sum().item(Float.self), 0.610559344291687, - accuracy: 0.01221118688583374) - let result = logSumExp(a, axis: -1) - XCTAssertEqual(result.shape, [4]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 1.5178632736206055, - accuracy: 0.03035726547241211) - XCTAssertEqual( - result.sum().item(Float.self), 6.071453094482422, - accuracy: 0.12142906188964844) - } - - func testLogSumExp4() { - MLXRandom.seed(341) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 3, 4, 3]) - XCTAssertEqual(a.shape, [2, 3, 4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.5185534358024597, - accuracy: 0.010371068716049195) - XCTAssertEqual( - a.sum().item(Float.self), 37.335845947265625, - accuracy: 0.7467169189453126) - let result = a.logSumExp(axes: [0, -1]) - XCTAssertEqual(result.shape, [3, 4]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 2.3395986557006836, - accuracy: 0.046791973114013674) - XCTAssertEqual( - result.sum().item(Float.self), 28.07518196105957, - accuracy: 0.5615036392211914) - } - - func testLogSumExp5() { - MLXRandom.seed(249) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 3, 4, 3]) - XCTAssertEqual(a.shape, [2, 3, 4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.48312702775001526, - accuracy: 0.009662540555000305) - XCTAssertEqual( - a.sum().item(Float.self), 34.7851448059082, - accuracy: 0.695702896118164) - let result = logSumExp(a, axes: [0, -1]) - XCTAssertEqual(result.shape, [3, 4]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 2.314382553100586, - accuracy: 0.04628765106201172) - XCTAssertEqual( - result.sum().item(Float.self), 27.7725887298584, - accuracy: 0.555451774597168) - } - - func testMax() { - MLXRandom.seed(747) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.23745165765285492, - accuracy: 0.004749033153057099) - XCTAssertEqual( - a.sum().item(Float.self), 2.8494198322296143, - accuracy: 0.056988396644592286) - let result = a.max() - XCTAssertEqual(result.shape, []) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 1.8713812828063965, - accuracy: 0.03742762565612793) - XCTAssertEqual( - result.sum().item(Float.self), 1.8713812828063965, - accuracy: 0.03742762565612793) - } - - func testMax1() { - MLXRandom.seed(333) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.08884426206350327, - accuracy: 0.0017768852412700653) - XCTAssertEqual( - a.sum().item(Float.self), 1.0661311149597168, - accuracy: 0.021322622299194335) - let result = max(a) - XCTAssertEqual(result.shape, []) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 1.4511317014694214, - accuracy: 0.029022634029388428) - XCTAssertEqual( - result.sum().item(Float.self), 1.4511317014694214, - accuracy: 0.029022634029388428) - } - - func testMax2() { - MLXRandom.seed(720) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.03598502278327942, - accuracy: -0.0007197004556655884) - XCTAssertEqual( - a.sum().item(Float.self), -0.431820273399353, - accuracy: -0.00863640546798706) - let result = a.max(axis: -1) - XCTAssertEqual(result.shape, [4]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 1.0814929008483887, - accuracy: 0.021629858016967773) - XCTAssertEqual( - result.sum().item(Float.self), 4.325971603393555, - accuracy: 0.08651943206787109) - } - - func testMax3() { - MLXRandom.seed(891) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.2537704408168793, - accuracy: -0.005075408816337585) - XCTAssertEqual( - a.sum().item(Float.self), -3.0452451705932617, - accuracy: -0.06090490341186523) - let result = max(a, axis: -1) - XCTAssertEqual(result.shape, [4]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.30163291096687317, - accuracy: 0.006032658219337464) - XCTAssertEqual( - result.sum().item(Float.self), 1.2065316438674927, - accuracy: 0.024130632877349855) - } - - func testMax4() { - MLXRandom.seed(64) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 3, 4, 3]) - XCTAssertEqual(a.shape, [2, 3, 4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.5222265124320984, - accuracy: 0.010444530248641969) - XCTAssertEqual( - a.sum().item(Float.self), 37.60030746459961, - accuracy: 0.7520061492919922) - let result = a.max(axes: [0, -1]) - XCTAssertEqual(result.shape, [3, 4]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.8983621597290039, - accuracy: 0.01796724319458008) - XCTAssertEqual( - result.sum().item(Float.self), 10.780345916748047, - accuracy: 0.21560691833496093) - } - - func testMax5() { - MLXRandom.seed(195) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 3, 4, 3]) - XCTAssertEqual(a.shape, [2, 3, 4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.4561886787414551, - accuracy: 0.009123773574829101) - XCTAssertEqual( - a.sum().item(Float.self), 32.845584869384766, - accuracy: 0.6569116973876953) - let result = max(a, axes: [0, -1]) - XCTAssertEqual(result.shape, [3, 4]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.7991632223129272, - accuracy: 0.015983264446258545) - XCTAssertEqual( - result.sum().item(Float.self), 9.589958190917969, - accuracy: 0.19179916381835938) - } - - func testMean() { - MLXRandom.seed(939) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.1672295778989792, - accuracy: -0.0033445915579795836) - XCTAssertEqual( - a.sum().item(Float.self), -2.0067548751831055, - accuracy: -0.040135097503662114) - let result = a.mean() - XCTAssertEqual(result.shape, []) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.1672295778989792, - accuracy: -0.0033445915579795836) - XCTAssertEqual( - result.sum().item(Float.self), -0.1672295778989792, - accuracy: -0.0033445915579795836) - } - - func testMean1() { - MLXRandom.seed(581) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.3291994333267212, - accuracy: -0.006583988666534424) - XCTAssertEqual( - a.sum().item(Float.self), -3.950392961502075, - accuracy: -0.0790078592300415) - let result = mean(a) - XCTAssertEqual(result.shape, []) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.3291994333267212, - accuracy: -0.006583988666534424) - XCTAssertEqual( - result.sum().item(Float.self), -0.3291994333267212, - accuracy: -0.006583988666534424) - } - - func testMean2() { - MLXRandom.seed(227) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.28339889645576477, - accuracy: -0.005667977929115295) - XCTAssertEqual( - a.sum().item(Float.self), -3.4007866382598877, - accuracy: -0.06801573276519776) - let result = a.mean(axis: -1) - XCTAssertEqual(result.shape, [4]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.2833988666534424, - accuracy: -0.005667977333068848) - XCTAssertEqual( - result.sum().item(Float.self), -1.1335954666137695, - accuracy: -0.022671909332275392) - } - - func testMean3() { - MLXRandom.seed(244) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.5779480934143066, - accuracy: 0.011558961868286134) - XCTAssertEqual( - a.sum().item(Float.self), 6.93537712097168, - accuracy: 0.1387075424194336) - let result = mean(a, axis: -1) - XCTAssertEqual(result.shape, [4]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.5779481530189514, - accuracy: 0.011558963060379028) - XCTAssertEqual( - result.sum().item(Float.self), 2.3117926120758057, - accuracy: 0.046235852241516114) - } - - func testMean4() { - MLXRandom.seed(822) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 3, 4, 3]) - XCTAssertEqual(a.shape, [2, 3, 4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.4916161596775055, - accuracy: 0.00983232319355011) - XCTAssertEqual( - a.sum().item(Float.self), 35.3963623046875, - accuracy: 0.70792724609375) - let result = a.mean(axes: [0, -1]) - XCTAssertEqual(result.shape, [3, 4]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.4916161894798279, - accuracy: 0.009832323789596557) - XCTAssertEqual( - result.sum().item(Float.self), 5.8993940353393555, - accuracy: 0.11798788070678712) - } - - func testMean5() { - MLXRandom.seed(990) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 3, 4, 3]) - XCTAssertEqual(a.shape, [2, 3, 4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.4603706896305084, - accuracy: 0.009207413792610168) - XCTAssertEqual( - a.sum().item(Float.self), 33.146690368652344, - accuracy: 0.6629338073730469) - let result = mean(a, axes: [0, -1]) - XCTAssertEqual(result.shape, [3, 4]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.4603707790374756, - accuracy: 0.009207415580749513) - XCTAssertEqual( - result.sum().item(Float.self), 5.524449348449707, - accuracy: 0.11048898696899415) - } - - func testMin() { - MLXRandom.seed(145) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.16449235379695892, - accuracy: -0.0032898470759391784) - XCTAssertEqual( - a.sum().item(Float.self), -1.9739081859588623, - accuracy: -0.03947816371917725) - let result = a.min() - XCTAssertEqual(result.shape, []) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -1.3645747900009155, - accuracy: -0.02729149580001831) - XCTAssertEqual( - result.sum().item(Float.self), -1.3645747900009155, - accuracy: -0.02729149580001831) - } - - func testMin1() { - MLXRandom.seed(822) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.025983989238739014, - accuracy: 0.0005196797847747803) - XCTAssertEqual( - a.sum().item(Float.self), 0.31180787086486816, - accuracy: 0.006236157417297363) - let result = min(a) - XCTAssertEqual(result.shape, []) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -1.8071335554122925, - accuracy: -0.03614267110824585) - XCTAssertEqual( - result.sum().item(Float.self), -1.8071335554122925, - accuracy: -0.03614267110824585) - } - - func testMin2() { - MLXRandom.seed(556) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.47528061270713806, - accuracy: 0.00950561225414276) - XCTAssertEqual( - a.sum().item(Float.self), 5.703367233276367, - accuracy: 0.11406734466552734) - let result = a.min(axis: -1) - XCTAssertEqual(result.shape, [4]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.6413009166717529, - accuracy: -0.012826018333435059) - XCTAssertEqual( - result.sum().item(Float.self), -2.5652036666870117, - accuracy: -0.051304073333740235) - } - - func testMin3() { - MLXRandom.seed(458) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.1867717206478119, - accuracy: 0.003735434412956238) - XCTAssertEqual( - a.sum().item(Float.self), 2.241260528564453, - accuracy: 0.04482521057128906) - let result = min(a, axis: -1) - XCTAssertEqual(result.shape, [4]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.8374209403991699, - accuracy: -0.0167484188079834) - XCTAssertEqual( - result.sum().item(Float.self), -3.3496837615966797, - accuracy: -0.0669936752319336) - } - - func testMin4() { - MLXRandom.seed(93) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 3, 4, 3]) - XCTAssertEqual(a.shape, [2, 3, 4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.4830058515071869, - accuracy: 0.009660117030143738) - XCTAssertEqual( - a.sum().item(Float.self), 34.77642059326172, - accuracy: 0.6955284118652344) - let result = a.min(axes: [0, -1]) - XCTAssertEqual(result.shape, [3, 4]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.11048994958400726, - accuracy: 0.0022097989916801454) - XCTAssertEqual( - result.sum().item(Float.self), 1.3258793354034424, - accuracy: 0.02651758670806885) - } - - func testMin5() { - MLXRandom.seed(82) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 3, 4, 3]) - XCTAssertEqual(a.shape, [2, 3, 4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.5077146887779236, - accuracy: 0.010154293775558472) - XCTAssertEqual( - a.sum().item(Float.self), 36.555458068847656, - accuracy: 0.7311091613769531) - let result = min(a, axes: [0, -1]) - XCTAssertEqual(result.shape, [3, 4]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.12779314815998077, - accuracy: 0.0025558629631996154) - XCTAssertEqual( - result.sum().item(Float.self), 1.5335177183151245, - accuracy: 0.030670354366302492) - } - - func testProduct() { - MLXRandom.seed(327) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.296897828578949, - accuracy: -0.00593795657157898) - XCTAssertEqual( - a.sum().item(Float.self), -3.5627739429473877, - accuracy: -0.07125547885894776) - let result = a.product() - XCTAssertEqual(result.shape, []) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 3.555632019924815e-06, - accuracy: 7.11126403984963e-08) - XCTAssertEqual( - result.sum().item(Float.self), 3.555632019924815e-06, - accuracy: 7.11126403984963e-08) - } - - func testProduct1() { - MLXRandom.seed(896) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.05823778361082077, - accuracy: -0.0011647556722164155) - XCTAssertEqual( - a.sum().item(Float.self), -0.6988533735275269, - accuracy: -0.013977067470550537) - let result = product(a) - XCTAssertEqual(result.shape, []) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.0047314404509961605, - accuracy: 9.462880901992321e-05) - XCTAssertEqual( - result.sum().item(Float.self), 0.0047314404509961605, - accuracy: 9.462880901992321e-05) - } - - func testProduct2() { - MLXRandom.seed(520) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.19801780581474304, - accuracy: 0.003960356116294861) - XCTAssertEqual( - a.sum().item(Float.self), 2.376213550567627, - accuracy: 0.04752427101135254) - let result = a.product(axis: -1) - XCTAssertEqual(result.shape, [4]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.0764155238866806, - accuracy: 0.0015283104777336122) - XCTAssertEqual( - result.sum().item(Float.self), 0.3056620955467224, - accuracy: 0.006113241910934449) - } - - func testProduct3() { - MLXRandom.seed(955) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.36796003580093384, - accuracy: 0.007359200716018677) - XCTAssertEqual( - a.sum().item(Float.self), 4.415520191192627, - accuracy: 0.08831040382385254) - let result = product(a, axis: -1) - XCTAssertEqual(result.shape, [4]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.0716419368982315, - accuracy: 0.0014328387379646302) - XCTAssertEqual( - result.sum().item(Float.self), 0.286567747592926, - accuracy: 0.005731354951858521) - } - - func testProduct4() { - MLXRandom.seed(501) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 3, 4, 3]) - XCTAssertEqual(a.shape, [2, 3, 4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.5341216921806335, - accuracy: 0.010682433843612671) - XCTAssertEqual( - a.sum().item(Float.self), 38.45676040649414, - accuracy: 0.7691352081298828) - let result = a.product(axes: [0, -1]) - XCTAssertEqual(result.shape, [3, 4]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.02406497858464718, - accuracy: 0.0004812995716929436) - XCTAssertEqual( - result.sum().item(Float.self), 0.28877973556518555, - accuracy: 0.0057755947113037115) - } - - func testProduct5() { - MLXRandom.seed(111) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 3, 4, 3]) - XCTAssertEqual(a.shape, [2, 3, 4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.4650188088417053, - accuracy: 0.009300376176834107) - XCTAssertEqual( - a.sum().item(Float.self), 33.481353759765625, - accuracy: 0.6696270751953125) - let result = product(a, axes: [0, -1]) - XCTAssertEqual(result.shape, [3, 4]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.007884619757533073, - accuracy: 0.00015769239515066148) - XCTAssertEqual( - result.sum().item(Float.self), 0.09461543709039688, - accuracy: 0.0018923087418079377) - } - - func testReciprocal() { - MLXRandom.seed(308) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.12339428067207336, - accuracy: -0.0024678856134414673) - XCTAssertEqual( - a.sum().item(Float.self), -1.4807313680648804, - accuracy: -0.029614627361297607) - let result = a.reciprocal() - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 1.3978341817855835, - accuracy: 0.02795668363571167) - XCTAssertEqual( - result.sum().item(Float.self), 16.774009704589844, - accuracy: 0.33548019409179686) - } - - func testReciprocal1() { - MLXRandom.seed(564) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.20050053298473358, - accuracy: -0.004010010659694672) - XCTAssertEqual( - a.sum().item(Float.self), -2.406006336212158, - accuracy: -0.048120126724243165) - let result = reciprocal(a) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.753860592842102, - accuracy: -0.015077211856842042) - XCTAssertEqual( - result.sum().item(Float.self), -9.046326637268066, - accuracy: -0.18092653274536133) - } - - func testRound() { - MLXRandom.seed(298) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.23528259992599487, - accuracy: 0.004705651998519898) - XCTAssertEqual( - a.sum().item(Float.self), 2.8233911991119385, - accuracy: 0.05646782398223877) - let result = a.round() - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.25, - accuracy: 0.005) - XCTAssertEqual( - result.sum().item(Float.self), 3.0, - accuracy: 0.06) - } - - func testRound1() { - MLXRandom.seed(723) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.16675782203674316, - accuracy: 0.0033351564407348632) - XCTAssertEqual( - a.sum().item(Float.self), 2.001093864440918, - accuracy: 0.04002187728881836) - let result = round(a) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.25, - accuracy: 0.005) - XCTAssertEqual( - result.sum().item(Float.self), 3.0, - accuracy: 0.06) - } - - func testSin() { - MLXRandom.seed(127) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.18461060523986816, - accuracy: -0.0036922121047973633) - XCTAssertEqual( - a.sum().item(Float.self), -2.215327262878418, - accuracy: -0.04430654525756836) - let result = a.sin() - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.15534619987010956, - accuracy: -0.0031069239974021914) - XCTAssertEqual( - result.sum().item(Float.self), -1.86415433883667, - accuracy: -0.0372830867767334) - } - - func testSin1() { - MLXRandom.seed(560) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.15535306930541992, - accuracy: -0.0031070613861083987) - XCTAssertEqual( - a.sum().item(Float.self), -1.864236831665039, - accuracy: -0.03728473663330078) - let result = sin(a) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.13987457752227783, - accuracy: -0.002797491550445557) - XCTAssertEqual( - result.sum().item(Float.self), -1.6784948110580444, - accuracy: -0.03356989622116089) - } - - func testCos() { - MLXRandom.seed(340) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.18920934200286865, - accuracy: 0.003784186840057373) - XCTAssertEqual( - a.sum().item(Float.self), 2.270512104034424, - accuracy: 0.04541024208068848) - let result = a.cos() - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.7257080078125, - accuracy: 0.01451416015625) - XCTAssertEqual( - result.sum().item(Float.self), 8.70849609375, - accuracy: 0.174169921875) - } - - func testCos1() { - MLXRandom.seed(834) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.18475177884101868, - accuracy: 0.0036950355768203737) - XCTAssertEqual( - a.sum().item(Float.self), 2.2170212268829346, - accuracy: 0.04434042453765869) - let result = cos(a) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.4963728189468384, - accuracy: 0.009927456378936769) - XCTAssertEqual( - result.sum().item(Float.self), 5.9564738273620605, - accuracy: 0.11912947654724121) - } - - func testSqrt() { - MLXRandom.seed(944) - let a = MLXRandom.uniform(low: 0.1, high: 2.0, [4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 1.1813615560531616, - accuracy: 0.023627231121063234) - XCTAssertEqual( - a.sum().item(Float.self), 14.176338195800781, - accuracy: 0.2835267639160156) - let result = a.sqrt() - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 1.040137767791748, - accuracy: 0.02080275535583496) - XCTAssertEqual( - result.sum().item(Float.self), 12.481653213500977, - accuracy: 0.24963306427001955) - } - - func testSqrt1() { - MLXRandom.seed(553) - let a = MLXRandom.uniform(low: 0.1, high: 2.0, [4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 1.2980049848556519, - accuracy: 0.025960099697113038) - XCTAssertEqual( - a.sum().item(Float.self), 15.576059341430664, - accuracy: 0.3115211868286133) - let result = sqrt(a) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 1.118424415588379, - accuracy: 0.02236848831176758) - XCTAssertEqual( - result.sum().item(Float.self), 13.421092987060547, - accuracy: 0.26842185974121097) - } - - func testSum() { - MLXRandom.seed(208) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.013726795092225075, - accuracy: 0.0002745359018445015) - XCTAssertEqual( - a.sum().item(Float.self), 0.1647215336561203, - accuracy: 0.003294430673122406) - let result = a.sum() - XCTAssertEqual(result.shape, []) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.1647215336561203, - accuracy: 0.003294430673122406) - XCTAssertEqual( - result.sum().item(Float.self), 0.1647215336561203, - accuracy: 0.003294430673122406) - } - - func testSum1() { - MLXRandom.seed(986) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.08680947870016098, - accuracy: 0.0017361895740032197) - XCTAssertEqual( - a.sum().item(Float.self), 1.0417137145996094, - accuracy: 0.020834274291992187) - let result = sum(a) - XCTAssertEqual(result.shape, []) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 1.0417137145996094, - accuracy: 0.020834274291992187) - XCTAssertEqual( - result.sum().item(Float.self), 1.0417137145996094, - accuracy: 0.020834274291992187) - } - - func testSum2() { - MLXRandom.seed(818) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.01743374764919281, - accuracy: -0.0003486749529838562) - XCTAssertEqual( - a.sum().item(Float.self), -0.20920497179031372, - accuracy: -0.004184099435806275) - let result = a.sum(axis: -1) - XCTAssertEqual(result.shape, [4]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.05230128765106201, - accuracy: -0.0010460257530212403) - XCTAssertEqual( - result.sum().item(Float.self), -0.20920515060424805, - accuracy: -0.004184103012084961) - } - - func testSum3() { - MLXRandom.seed(617) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.2992609739303589, - accuracy: -0.005985219478607178) - XCTAssertEqual( - a.sum().item(Float.self), -3.5911316871643066, - accuracy: -0.07182263374328614) - let result = sum(a, axis: -1) - XCTAssertEqual(result.shape, [4]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.8977828621864319, - accuracy: -0.017955657243728638) - XCTAssertEqual( - result.sum().item(Float.self), -3.5911314487457275, - accuracy: -0.07182262897491455) - } - - func testSum4() { - MLXRandom.seed(560) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 3, 4, 3]) - XCTAssertEqual(a.shape, [2, 3, 4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.4772148132324219, - accuracy: 0.009544296264648438) - XCTAssertEqual( - a.sum().item(Float.self), 34.359466552734375, - accuracy: 0.6871893310546875) - let result = a.sum(axes: [0, -1]) - XCTAssertEqual(result.shape, [3, 4]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 2.8632893562316895, - accuracy: 0.05726578712463379) - XCTAssertEqual( - result.sum().item(Float.self), 34.35947036743164, - accuracy: 0.6871894073486329) - } - - func testSum5() { - MLXRandom.seed(601) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 3, 4, 3]) - XCTAssertEqual(a.shape, [2, 3, 4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.5380033850669861, - accuracy: 0.010760067701339722) - XCTAssertEqual( - a.sum().item(Float.self), 38.736244201660156, - accuracy: 0.7747248840332032) - let result = sum(a, axes: [0, -1]) - XCTAssertEqual(result.shape, [3, 4]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 3.228020429611206, - accuracy: 0.06456040859222412) - XCTAssertEqual( - result.sum().item(Float.self), 38.736244201660156, - accuracy: 0.7747248840332032) - } - - func testVariance() { - MLXRandom.seed(294) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.047887593507766724, - accuracy: 0.0009577518701553345) - XCTAssertEqual( - a.sum().item(Float.self), 0.5746511220932007, - accuracy: 0.011493022441864014) - let result = a.variance() - XCTAssertEqual(result.shape, []) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.6768754124641418, - accuracy: 0.013537508249282838) - XCTAssertEqual( - result.sum().item(Float.self), 0.6768754124641418, - accuracy: 0.013537508249282838) - } - - func testVariance1() { - MLXRandom.seed(455) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.07983233034610748, - accuracy: 0.0015966466069221496) - XCTAssertEqual( - a.sum().item(Float.self), 0.957987904548645, - accuracy: 0.019159758090972902) - let result = variance(a) - XCTAssertEqual(result.shape, []) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.5957978367805481, - accuracy: 0.011915956735610963) - XCTAssertEqual( - result.sum().item(Float.self), 0.5957978367805481, - accuracy: 0.011915956735610963) - } - - func testVariance2() { - MLXRandom.seed(93) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.028267383575439453, - accuracy: -0.0005653476715087891) - XCTAssertEqual( - a.sum().item(Float.self), -0.33920860290527344, - accuracy: -0.006784172058105469) - let result = a.variance(axis: -1) - XCTAssertEqual(result.shape, [4]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 1.0322922468185425, - accuracy: 0.02064584493637085) - XCTAssertEqual( - result.sum().item(Float.self), 4.12916898727417, - accuracy: 0.0825833797454834) - } - - func testVariance3() { - MLXRandom.seed(610) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.01769038289785385, - accuracy: -0.00035380765795707704) - XCTAssertEqual( - a.sum().item(Float.self), -0.21228459477424622, - accuracy: -0.004245691895484924) - let result = variance(a, axis: -1) - XCTAssertEqual(result.shape, [4]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.591558039188385, - accuracy: 0.0118311607837677) - XCTAssertEqual( - result.sum().item(Float.self), 2.36623215675354, - accuracy: 0.0473246431350708) - } - - func testVariance4() { - MLXRandom.seed(817) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 3, 4, 3]) - XCTAssertEqual(a.shape, [2, 3, 4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.4460234045982361, - accuracy: 0.008920468091964721) - XCTAssertEqual( - a.sum().item(Float.self), 32.113685607910156, - accuracy: 0.6422737121582032) - let result = a.variance(axes: [0, -1]) - XCTAssertEqual(result.shape, [3, 4]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.0666377991437912, - accuracy: 0.001332755982875824) - XCTAssertEqual( - result.sum().item(Float.self), 0.7996535301208496, - accuracy: 0.015993070602416993) - } - - func testVariance5() { - MLXRandom.seed(394) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 3, 4, 3]) - XCTAssertEqual(a.shape, [2, 3, 4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.56817227602005, - accuracy: 0.011363445520401002) - XCTAssertEqual( - a.sum().item(Float.self), 40.90840530395508, - accuracy: 0.8181681060791016) - let result = variance(a, axes: [0, -1]) - XCTAssertEqual(result.shape, [3, 4]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.07289446890354156, - accuracy: 0.0014578893780708313) - XCTAssertEqual( - result.sum().item(Float.self), 0.8747336268424988, - accuracy: 0.017494672536849977) - } - - func testAcos() { - MLXRandom.seed(324) - let a = MLXRandom.uniform(low: 0.1, high: 1.0, [4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.5365500450134277, - accuracy: 0.010731000900268555) - XCTAssertEqual( - a.sum().item(Float.self), 6.438600063323975, - accuracy: 0.1287720012664795) - let result = acos(a) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.9425711631774902, - accuracy: 0.018851423263549806) - XCTAssertEqual( - result.sum().item(Float.self), 11.310853958129883, - accuracy: 0.22621707916259767) - } - - func testAcosh() { - MLXRandom.seed(589) - let a = MLXRandom.uniform(low: 1, high: 3, [4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 1.9692752361297607, - accuracy: 0.039385504722595215) - XCTAssertEqual( - a.sum().item(Float.self), 23.631301879882812, - accuracy: 0.47262603759765626) - let result = acosh(a) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 1.259089469909668, - accuracy: 0.02518178939819336) - XCTAssertEqual( - result.sum().item(Float.self), 15.109073638916016, - accuracy: 0.3021814727783203) - } - - func testAsin() { - MLXRandom.seed(247) - let a = MLXRandom.uniform(low: 0.1, high: 1.0, [4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.4699937105178833, - accuracy: 0.009399874210357666) - XCTAssertEqual( - a.sum().item(Float.self), 5.6399245262146, - accuracy: 0.112798490524292) - let result = asin(a) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.5069084763526917, - accuracy: 0.010138169527053834) - XCTAssertEqual( - result.sum().item(Float.self), 6.082901477813721, - accuracy: 0.12165802955627442) - } - - func testAsinh() { - MLXRandom.seed(297) - let a = MLXRandom.uniform(low: 1, high: 3, [4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 1.7906322479248047, - accuracy: 0.03581264495849609) - XCTAssertEqual( - a.sum().item(Float.self), 21.487586975097656, - accuracy: 0.4297517395019531) - let result = asinh(a) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 1.3248964548110962, - accuracy: 0.026497929096221923) - XCTAssertEqual( - result.sum().item(Float.self), 15.898756980895996, - accuracy: 0.31797513961791996) - } - - func testAtan() { - MLXRandom.seed(188) - let a = MLXRandom.uniform(low: 0.1, high: 1.0, [4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.43250617384910583, - accuracy: 0.008650123476982116) - XCTAssertEqual( - a.sum().item(Float.self), 5.1900739669799805, - accuracy: 0.10380147933959961) - let result = atan(a) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.3987109661102295, - accuracy: 0.00797421932220459) - XCTAssertEqual( - result.sum().item(Float.self), 4.784531593322754, - accuracy: 0.09569063186645509) - } - - func testAtanh() { - MLXRandom.seed(193) - let a = MLXRandom.uniform(low: 0.1, high: 0.9, [4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.593014657497406, - accuracy: 0.01186029314994812) - XCTAssertEqual( - a.sum().item(Float.self), 7.116175651550293, - accuracy: 0.14232351303100585) - let result = atanh(a) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.759239137172699, - accuracy: 0.01518478274345398) - XCTAssertEqual( - result.sum().item(Float.self), 9.110869407653809, - accuracy: 0.18221738815307617) - } - - func testCeil() { - MLXRandom.seed(841) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.3292633295059204, - accuracy: -0.006585266590118408) - XCTAssertEqual( - a.sum().item(Float.self), -3.951159715652466, - accuracy: -0.07902319431304931) - let result = ceil(a) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.3333333432674408, - accuracy: 0.006666666865348816) - XCTAssertEqual( - result.sum().item(Float.self), 4.0, - accuracy: 0.08) - } - - func testCosh() { - MLXRandom.seed(191) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.5465152263641357, - accuracy: -0.010930304527282716) - XCTAssertEqual( - a.sum().item(Float.self), -6.558182239532471, - accuracy: -0.1311636447906494) - let result = cosh(a) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 1.735574722290039, - accuracy: 0.03471149444580078) - XCTAssertEqual( - result.sum().item(Float.self), 20.82689666748047, - accuracy: 0.41653793334960937) - } - - func testErf() { - MLXRandom.seed(33) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.15268580615520477, - accuracy: 0.0030537161231040956) - XCTAssertEqual( - a.sum().item(Float.self), 1.8322296142578125, - accuracy: 0.03664459228515625) - let result = erf(a) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.16305577754974365, - accuracy: 0.003261115550994873) - XCTAssertEqual( - result.sum().item(Float.self), 1.9566693305969238, - accuracy: 0.03913338661193848) - } - - func testErfInverse() { - MLXRandom.seed(627) - let a = MLXRandom.uniform(low: 0.1, high: 0.9, [4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.4853326678276062, - accuracy: 0.009706653356552124) - XCTAssertEqual( - a.sum().item(Float.self), 5.823991775512695, - accuracy: 0.11647983551025391) - let result = erfInverse(a) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.5043282508850098, - accuracy: 0.010086565017700196) - XCTAssertEqual( - result.sum().item(Float.self), 6.051938533782959, - accuracy: 0.12103877067565919) - } - - func testLogicalNot() { - MLXRandom.seed(672) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.16328130662441254, - accuracy: 0.0032656261324882506) - XCTAssertEqual( - a.sum().item(Float.self), 1.9593756198883057, - accuracy: 0.03918751239776611) - let result = logicalNot(a) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .bool) - XCTAssertEqual(result.all().item(), false) - XCTAssertEqual(result.any().item(), false) - } - - func testNegative() { - MLXRandom.seed(266) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.20902606844902039, - accuracy: 0.004180521368980407) - XCTAssertEqual( - a.sum().item(Float.self), 2.508312702178955, - accuracy: 0.050166254043579106) - let result = negative(a) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.20902606844902039, - accuracy: -0.004180521368980407) - XCTAssertEqual( - result.sum().item(Float.self), -2.508312702178955, - accuracy: -0.050166254043579106) - } - - func testSigmoid() { - MLXRandom.seed(487) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.4737564027309418, - accuracy: -0.009475128054618835) - XCTAssertEqual( - a.sum().item(Float.self), -5.685076713562012, - accuracy: -0.11370153427124023) - let result = sigmoid(a) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.40112757682800293, - accuracy: 0.00802255153656006) - XCTAssertEqual( - result.sum().item(Float.self), 4.813530921936035, - accuracy: 0.0962706184387207) - } - - func testSign() { - MLXRandom.seed(70) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.19362865388393402, - accuracy: -0.0038725730776786806) - XCTAssertEqual( - a.sum().item(Float.self), -2.3235437870025635, - accuracy: -0.04647087574005127) - let result = sign(a) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.3333333432674408, - accuracy: -0.006666666865348816) - XCTAssertEqual( - result.sum().item(Float.self), -4.0, - accuracy: -0.08) - } - - func testSinh() { - MLXRandom.seed(91) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.37823864817619324, - accuracy: 0.007564772963523865) - XCTAssertEqual( - a.sum().item(Float.self), 4.538863658905029, - accuracy: 0.09077727317810058) - let result = sinh(a) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.8675357103347778, - accuracy: 0.017350714206695556) - XCTAssertEqual( - result.sum().item(Float.self), 10.410428047180176, - accuracy: 0.20820856094360352) - } - - func testSoftMax() { - MLXRandom.seed(695) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.14673520624637604, - accuracy: -0.002934704124927521) - XCTAssertEqual( - a.sum().item(Float.self), -1.7608224153518677, - accuracy: -0.03521644830703735) - let result = softMax(a) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.0833333283662796, - accuracy: 0.0016666665673255921) - XCTAssertEqual( - result.sum().item(Float.self), 0.9999998807907104, - accuracy: 0.01999999761581421) - } - - func testSoftMax1() { - MLXRandom.seed(775) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.27563148736953735, - accuracy: 0.005512629747390747) - XCTAssertEqual( - a.sum().item(Float.self), 3.307577610015869, - accuracy: 0.06615155220031739) - let result = softMax(a, axis: -1) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.3333333432674408, - accuracy: 0.006666666865348816) - XCTAssertEqual( - result.sum().item(Float.self), 4.0, - accuracy: 0.08) - } - - func testSoftMax2() { - MLXRandom.seed(133) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 3, 4, 3]) - XCTAssertEqual(a.shape, [2, 3, 4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.5304720401763916, - accuracy: 0.010609440803527832) - XCTAssertEqual( - a.sum().item(Float.self), 38.19398498535156, - accuracy: 0.7638796997070313) - let result = softMax(a, axes: [0, -1]) - XCTAssertEqual(result.shape, [2, 3, 4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.1666666716337204, - accuracy: 0.003333333432674408) - XCTAssertEqual( - result.sum().item(Float.self), 12.0, - accuracy: 0.24) - } - - func testTan() { - MLXRandom.seed(897) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.3048655390739441, - accuracy: -0.006097310781478882) - XCTAssertEqual( - a.sum().item(Float.self), -3.65838623046875, - accuracy: -0.073167724609375) - let result = tan(a) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.0694185122847557, - accuracy: 0.001388370245695114) - XCTAssertEqual( - result.sum().item(Float.self), 0.8330221176147461, - accuracy: 0.016660442352294923) - } - - func testTanh() { - MLXRandom.seed(153) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.45169347524642944, - accuracy: -0.009033869504928588) - XCTAssertEqual( - a.sum().item(Float.self), -5.420321464538574, - accuracy: -0.10840642929077149) - let result = tanh(a) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.3471390902996063, - accuracy: -0.006942781805992127) - XCTAssertEqual( - result.sum().item(Float.self), -4.165668964385986, - accuracy: -0.08331337928771973) - } - - func testMLXadd() { - MLXRandom.seed(945) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.3744015097618103, - accuracy: 0.007488030195236206) - XCTAssertEqual( - a.sum().item(Float.self), 4.4928178787231445, - accuracy: 0.08985635757446289) - let b = MLXRandom.normal([4, 3]) - XCTAssertEqual(b.shape, [4, 3]) - XCTAssertEqual(b.dtype, .float32) - XCTAssertEqual( - b.mean().item(Float.self), -0.20527128875255585, - accuracy: -0.004105425775051117) - XCTAssertEqual( - b.sum().item(Float.self), -2.4632554054260254, - accuracy: -0.04926510810852051) - let result = MLX.add(a, b) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.16913020610809326, - accuracy: 0.0033826041221618653) - XCTAssertEqual( - result.sum().item(Float.self), 2.029562473297119, - accuracy: 0.040591249465942385) - } - - func testConv1d() { - MLXRandom.seed(39) - let a = MLXRandom.uniform(0.0 ..< 1.0, [4, 10, 4]) - XCTAssertEqual(a.shape, [4, 10, 4]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.5018469095230103, - accuracy: 0.010036938190460205) - XCTAssertEqual( - a.sum().item(Float.self), 80.29550170898438, - accuracy: 1.6059100341796875) - let b = MLXRandom.uniform(0.0 ..< 1.0, [2, 10, 4]) - XCTAssertEqual(b.shape, [2, 10, 4]) - XCTAssertEqual(b.dtype, .float32) - XCTAssertEqual( - b.mean().item(Float.self), 0.5621036291122437, - accuracy: 0.011242072582244873) - XCTAssertEqual( - b.sum().item(Float.self), 44.96828842163086, - accuracy: 0.8993657684326172) - let result = conv1d(a, b) - XCTAssertEqual(result.shape, [4, 1, 2]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 11.004511833190918, - accuracy: 0.22009023666381836) - XCTAssertEqual( - result.sum().item(Float.self), 88.03609466552734, - accuracy: 1.7607218933105468) - } - - func testConv2d() { - MLXRandom.seed(862) - let a = MLXRandom.uniform(0.0 ..< 1.0, [4, 10, 12, 4]) - XCTAssertEqual(a.shape, [4, 10, 12, 4]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.5062730312347412, - accuracy: 0.010125460624694825) - XCTAssertEqual( - a.sum().item(Float.self), 972.044189453125, - accuracy: 19.4408837890625) - let b = MLXRandom.uniform(0.0 ..< 1.0, [2, 10, 12, 4]) - XCTAssertEqual(b.shape, [2, 10, 12, 4]) - XCTAssertEqual(b.dtype, .float32) - XCTAssertEqual( - b.mean().item(Float.self), 0.4991826117038727, - accuracy: 0.009983652234077454) - XCTAssertEqual( - b.sum().item(Float.self), 479.21527099609375, - accuracy: 9.584305419921876) - let result = conv2d(a, b) - XCTAssertEqual(result.shape, [4, 1, 1, 2]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 121.8356704711914, - accuracy: 2.4367134094238283) - XCTAssertEqual( - result.sum().item(Float.self), 974.6853637695312, - accuracy: 19.493707275390626) - } - - func testConvolve() { - MLXRandom.seed(82) - let a = MLXRandom.uniform(0.0 ..< 1.0, [20]) - XCTAssertEqual(a.shape, [20]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.4547460675239563, - accuracy: 0.009094921350479125) - XCTAssertEqual( - a.sum().item(Float.self), 9.094921112060547, - accuracy: 0.18189842224121094) - let b = MLXRandom.uniform(0.0 ..< 1.0, [4]) - XCTAssertEqual(b.shape, [4]) - XCTAssertEqual(b.dtype, .float32) - XCTAssertEqual( - b.mean().item(Float.self), 0.36941707134246826, - accuracy: 0.007388341426849365) - XCTAssertEqual( - b.sum().item(Float.self), 1.477668285369873, - accuracy: 0.02955336570739746) - let result = convolve(a, b) - XCTAssertEqual(result.shape, [23]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.5843163132667542, - accuracy: 0.011686326265335084) - XCTAssertEqual( - result.sum().item(Float.self), 13.439274787902832, - accuracy: 0.26878549575805666) - } - - func testDivide() { - MLXRandom.seed(919) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.2972916066646576, - accuracy: -0.005945832133293152) - XCTAssertEqual( - a.sum().item(Float.self), -3.5674991607666016, - accuracy: -0.07134998321533204) - let b = MLXRandom.normal([4, 3]) - XCTAssertEqual(b.shape, [4, 3]) - XCTAssertEqual(b.dtype, .float32) - XCTAssertEqual( - b.mean().item(Float.self), 0.04941798374056816, - accuracy: 0.0009883596748113633) - XCTAssertEqual( - b.sum().item(Float.self), 0.5930157899856567, - accuracy: 0.011860315799713136) - let result = divide(a, b) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -3.7239060401916504, - accuracy: -0.07447812080383301) - XCTAssertEqual( - result.sum().item(Float.self), -44.68687057495117, - accuracy: -0.8937374114990234) - } - - func testEqual() { - MLXRandom.seed(716) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.3898356854915619, - accuracy: 0.007796713709831238) - XCTAssertEqual( - a.sum().item(Float.self), 4.678028106689453, - accuracy: 0.09356056213378906) - let b = MLXRandom.normal([4, 3]) - XCTAssertEqual(b.shape, [4, 3]) - XCTAssertEqual(b.dtype, .float32) - XCTAssertEqual( - b.mean().item(Float.self), -0.30894142389297485, - accuracy: -0.0061788284778594976) - XCTAssertEqual( - b.sum().item(Float.self), -3.707296848297119, - accuracy: -0.07414593696594239) - let result = equal(a, b) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .bool) - XCTAssertEqual(result.all().item(), false) - XCTAssertEqual(result.any().item(), false) - } - - func testGreater() { - MLXRandom.seed(945) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.3744015097618103, - accuracy: 0.007488030195236206) - XCTAssertEqual( - a.sum().item(Float.self), 4.4928178787231445, - accuracy: 0.08985635757446289) - let b = MLXRandom.normal([4, 3]) - XCTAssertEqual(b.shape, [4, 3]) - XCTAssertEqual(b.dtype, .float32) - XCTAssertEqual( - b.mean().item(Float.self), -0.20527128875255585, - accuracy: -0.004105425775051117) - XCTAssertEqual( - b.sum().item(Float.self), -2.4632554054260254, - accuracy: -0.04926510810852051) - let result = greater(a, b) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .bool) - XCTAssertEqual(result.all().item(), false) - XCTAssertEqual(result.any().item(), true) - } - - func testGreaterEqual() { - MLXRandom.seed(849) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.5959362983703613, - accuracy: -0.011918725967407227) - XCTAssertEqual( - a.sum().item(Float.self), -7.151235103607178, - accuracy: -0.14302470207214354) - let b = MLXRandom.normal([4, 3]) - XCTAssertEqual(b.shape, [4, 3]) - XCTAssertEqual(b.dtype, .float32) - XCTAssertEqual( - b.mean().item(Float.self), 0.42154282331466675, - accuracy: 0.008430856466293336) - XCTAssertEqual( - b.sum().item(Float.self), 5.058513641357422, - accuracy: 0.10117027282714844) - let result = greaterEqual(a, b) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .bool) - XCTAssertEqual(result.all().item(), false) - XCTAssertEqual(result.any().item(), true) - } - - func testLess() { - MLXRandom.seed(553) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.41794854402542114, - accuracy: 0.008358970880508423) - XCTAssertEqual( - a.sum().item(Float.self), 5.015382289886475, - accuracy: 0.1003076457977295) - let b = MLXRandom.normal([4, 3]) - XCTAssertEqual(b.shape, [4, 3]) - XCTAssertEqual(b.dtype, .float32) - XCTAssertEqual( - b.mean().item(Float.self), -0.19319520890712738, - accuracy: -0.0038639041781425476) - XCTAssertEqual( - b.sum().item(Float.self), -2.318342447280884, - accuracy: -0.04636684894561768) - let result = less(a, b) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .bool) - XCTAssertEqual(result.all().item(), false) - XCTAssertEqual(result.any().item(), true) - } - - func testLessEqual() { - MLXRandom.seed(699) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.06386314332485199, - accuracy: 0.0012772628664970398) - XCTAssertEqual( - a.sum().item(Float.self), 0.7663577198982239, - accuracy: 0.015327154397964478) - let b = MLXRandom.normal([4, 3]) - XCTAssertEqual(b.shape, [4, 3]) - XCTAssertEqual(b.dtype, .float32) - XCTAssertEqual( - b.mean().item(Float.self), -0.44176310300827026, - accuracy: -0.008835262060165406) - XCTAssertEqual( - b.sum().item(Float.self), -5.301156997680664, - accuracy: -0.10602313995361329) - let result = lessEqual(a, b) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .bool) - XCTAssertEqual(result.all().item(), false) - XCTAssertEqual(result.any().item(), true) - } - - func testLogAddExp() { - MLXRandom.seed(400) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.19093476235866547, - accuracy: -0.0038186952471733096) - XCTAssertEqual( - a.sum().item(Float.self), -2.291217088699341, - accuracy: -0.04582434177398682) - let b = MLXRandom.normal([4, 3]) - XCTAssertEqual(b.shape, [4, 3]) - XCTAssertEqual(b.dtype, .float32) - XCTAssertEqual( - b.mean().item(Float.self), 0.09774001687765121, - accuracy: 0.001954800337553024) - XCTAssertEqual( - b.sum().item(Float.self), 1.1728801727294922, - accuracy: 0.023457603454589845) - let result = logAddExp(a, b) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.7312957048416138, - accuracy: 0.014625914096832275) - XCTAssertEqual( - result.sum().item(Float.self), 8.775547981262207, - accuracy: 0.17551095962524416) - } - - func testMatmul() { - MLXRandom.seed(857) - let a = MLXRandom.uniform(0.0 ..< 1.0, [10, 8]) - XCTAssertEqual(a.shape, [10, 8]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.517481803894043, - accuracy: 0.01034963607788086) - XCTAssertEqual( - a.sum().item(Float.self), 41.39854431152344, - accuracy: 0.8279708862304688) - let b = MLXRandom.uniform(0.0 ..< 1.0, [8, 13]) - XCTAssertEqual(b.shape, [8, 13]) - XCTAssertEqual(b.dtype, .float32) - XCTAssertEqual( - b.mean().item(Float.self), 0.4867754876613617, - accuracy: 0.009735509753227234) - XCTAssertEqual( - b.sum().item(Float.self), 50.62464904785156, - accuracy: 1.0124929809570313) - let result = matmul(a, b) - XCTAssertEqual(result.shape, [10, 13]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 2.032482862472534, - accuracy: 0.04064965724945068) - XCTAssertEqual( - result.sum().item(Float.self), 264.2227783203125, - accuracy: 5.28445556640625) - } - - func testMaximum() { - MLXRandom.seed(722) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.01978858932852745, - accuracy: -0.00039577178657054904) - XCTAssertEqual( - a.sum().item(Float.self), -0.2374630719423294, - accuracy: -0.004749261438846588) - let b = MLXRandom.normal([4, 3]) - XCTAssertEqual(b.shape, [4, 3]) - XCTAssertEqual(b.dtype, .float32) - XCTAssertEqual( - b.mean().item(Float.self), 0.34027594327926636, - accuracy: 0.006805518865585327) - XCTAssertEqual( - b.sum().item(Float.self), 4.083311080932617, - accuracy: 0.08166622161865235) - let result = maximum(a, b) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.6268421411514282, - accuracy: 0.012536842823028565) - XCTAssertEqual( - result.sum().item(Float.self), 7.522105693817139, - accuracy: 0.15044211387634276) - } - - func testMinimum() { - MLXRandom.seed(537) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.15274646878242493, - accuracy: 0.0030549293756484985) - XCTAssertEqual( - a.sum().item(Float.self), 1.8329575061798096, - accuracy: 0.03665915012359619) - let b = MLXRandom.normal([4, 3]) - XCTAssertEqual(b.shape, [4, 3]) - XCTAssertEqual(b.dtype, .float32) - XCTAssertEqual( - b.mean().item(Float.self), 0.16062960028648376, - accuracy: 0.003212592005729675) - XCTAssertEqual( - b.sum().item(Float.self), 1.9275552034378052, - accuracy: 0.03855110406875611) - let result = minimum(a, b) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.4159923493862152, - accuracy: -0.008319846987724304) - XCTAssertEqual( - result.sum().item(Float.self), -4.991908073425293, - accuracy: -0.09983816146850587) - } - - func testMultiply() { - MLXRandom.seed(282) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.016050420701503754, - accuracy: -0.0003210084140300751) - XCTAssertEqual( - a.sum().item(Float.self), -0.19260503351688385, - accuracy: -0.003852100670337677) - let b = MLXRandom.normal([4, 3]) - XCTAssertEqual(b.shape, [4, 3]) - XCTAssertEqual(b.dtype, .float32) - XCTAssertEqual( - b.mean().item(Float.self), 0.2504279613494873, - accuracy: 0.0050085592269897465) - XCTAssertEqual( - b.sum().item(Float.self), 3.0051355361938477, - accuracy: 0.060102710723876955) - let result = multiply(a, b) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.2712330222129822, - accuracy: -0.005424660444259643) - XCTAssertEqual( - result.sum().item(Float.self), -3.254796266555786, - accuracy: -0.06509592533111572) - } - - func testNotEqual() { - MLXRandom.seed(534) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.5144619941711426, - accuracy: -0.010289239883422853) - XCTAssertEqual( - a.sum().item(Float.self), -6.173543453216553, - accuracy: -0.12347086906433105) - let b = MLXRandom.normal([4, 3]) - XCTAssertEqual(b.shape, [4, 3]) - XCTAssertEqual(b.dtype, .float32) - XCTAssertEqual( - b.mean().item(Float.self), -0.15448352694511414, - accuracy: -0.0030896705389022827) - XCTAssertEqual( - b.sum().item(Float.self), -1.8538023233413696, - accuracy: -0.03707604646682739) - let result = notEqual(a, b) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .bool) - XCTAssertEqual(result.all().item(), true) - XCTAssertEqual(result.any().item(), true) - } - - func testRemainder() { - MLXRandom.seed(831) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.27154868841171265, - accuracy: 0.005430973768234253) - XCTAssertEqual( - a.sum().item(Float.self), 3.2585840225219727, - accuracy: 0.06517168045043946) - let b = MLXRandom.normal([4, 3]) - XCTAssertEqual(b.shape, [4, 3]) - XCTAssertEqual(b.dtype, .float32) - XCTAssertEqual( - b.mean().item(Float.self), 0.1420227438211441, - accuracy: 0.0028404548764228823) - XCTAssertEqual( - b.sum().item(Float.self), 1.7042728662490845, - accuracy: 0.03408545732498169) - let result = remainder(a, b) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.007391604594886303, - accuracy: -0.00014783209189772606) - XCTAssertEqual( - result.sum().item(Float.self), -0.08869925141334534, - accuracy: -0.0017739850282669067) - } - - func testSubtract() { - MLXRandom.seed(241) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.1329968273639679, - accuracy: -0.002659936547279358) - XCTAssertEqual( - a.sum().item(Float.self), -1.5959618091583252, - accuracy: -0.031919236183166506) - let b = MLXRandom.normal([4, 3]) - XCTAssertEqual(b.shape, [4, 3]) - XCTAssertEqual(b.dtype, .float32) - XCTAssertEqual( - b.mean().item(Float.self), 0.03048303723335266, - accuracy: 0.0006096607446670532) - XCTAssertEqual( - b.sum().item(Float.self), 0.36579644680023193, - accuracy: 0.007315928936004639) - let result = subtract(a, b) - XCTAssertEqual(result.shape, [4, 3]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.16347989439964294, - accuracy: -0.003269597887992859) - XCTAssertEqual( - result.sum().item(Float.self), -1.9617586135864258, - accuracy: -0.03923517227172851) - } - - func testQuantize() { - MLXRandom.seed(869) - let w = MLXRandom.uniform(0.0 ..< 1.0, [32, 256]) - let (wq, scales, biases) = quantized(w, bits: 8) - XCTAssertEqual(wq.shape, [32, 64]) - XCTAssertEqual(wq.dtype, .uint32) - XCTAssertEqual( - wq.mean().item(Float.self), 732984.1875, - accuracy: 14659.68375) - XCTAssertEqual( - wq.sum().item(Float.self), 1_501_151_616, - accuracy: 30023031.62) - - XCTAssertEqual(scales.shape, [32, 4]) - XCTAssertEqual(scales.dtype, .float32) - XCTAssertEqual( - scales.mean().item(Float.self), -0.0037985900416970253, - accuracy: -7.597180083394051e-05) - XCTAssertEqual( - scales.sum().item(Float.self), -0.48621952533721924, - accuracy: -0.009724390506744385) - - if let biases { - XCTAssertEqual(biases.shape, [32, 4]) - XCTAssertEqual(biases.dtype, .float32) - XCTAssertEqual( - biases.mean().item(Float.self), 0.9862217307090759, - accuracy: 0.01972443461418152) - XCTAssertEqual( - biases.sum().item(Float.self), 126.23638153076172, - accuracy: 2.5247276306152346) - } else { - XCTFail("biases should not be nil") - } - - } - - func testFft_() { - MLXRandom.seed(220) - let r = MLXRandom.uniform(0.0 ..< 1.0, [100, 100]) - XCTAssertEqual(r.shape, [100, 100]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.49703681468963623, - accuracy: 0.009940736293792725) - XCTAssertEqual( - r.sum().item(Float.self), 4970.3681640625, - accuracy: 99.40736328125) - let i = MLXRandom.uniform(0.0 ..< 1.0, [100, 100]) - XCTAssertEqual(i.shape, [100, 100]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.49617284536361694, - accuracy: 0.00992345690727234) - XCTAssertEqual( - i.sum().item(Float.self), 4961.728515625, - accuracy: 99.2345703125) - let c = r + i.asImaginary() - let result = fft(c, stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [100, 100]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 0.5186547636985779, - accuracy: 0.010373095273971558) - XCTAssertEqual( - resultReal.sum().item(Float.self), 5186.5478515625, - accuracy: 103.73095703125) - XCTAssertEqual(resultImaginary.shape, [100, 100]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), 0.4938635528087616, - accuracy: 0.009877271056175233) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), 4938.6357421875, - accuracy: 98.77271484375001) - } - - func testFft_1() { - MLXRandom.seed(916) - let r = MLXRandom.uniform(0.0 ..< 1.0, [100, 100]) - XCTAssertEqual(r.shape, [100, 100]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.4961331784725189, - accuracy: 0.00992266356945038) - XCTAssertEqual( - r.sum().item(Float.self), 4961.33203125, - accuracy: 99.226640625) - let i = MLXRandom.uniform(0.0 ..< 1.0, [100, 100]) - XCTAssertEqual(i.shape, [100, 100]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.49821561574935913, - accuracy: 0.009964312314987183) - XCTAssertEqual( - i.sum().item(Float.self), 4982.15625, - accuracy: 99.643125) - let c = r + i.asImaginary() - let result = fft(c, n: 80, stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [100, 80]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 0.4616168737411499, - accuracy: 0.009232337474822999) - XCTAssertEqual( - resultReal.sum().item(Float.self), 3692.934814453125, - accuracy: 73.85869628906251) - XCTAssertEqual(resultImaginary.shape, [100, 80]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), 0.4507213234901428, - accuracy: 0.009014426469802857) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), 3605.7705078125, - accuracy: 72.11541015625001) - } - - func testFft_2() { - MLXRandom.seed(695) - let r = MLXRandom.uniform(0.0 ..< 1.0, [100, 100]) - XCTAssertEqual(r.shape, [100, 100]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.497734934091568, - accuracy: 0.00995469868183136) - XCTAssertEqual( - r.sum().item(Float.self), 4977.349609375, - accuracy: 99.5469921875) - let i = MLXRandom.uniform(0.0 ..< 1.0, [100, 100]) - XCTAssertEqual(i.shape, [100, 100]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.5017335414886475, - accuracy: 0.010034670829772949) - XCTAssertEqual( - i.sum().item(Float.self), 5017.33544921875, - accuracy: 100.346708984375) - let c = r + i.asImaginary() - let result = fft(c, n: 120, stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [100, 120]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 0.5616999268531799, - accuracy: 0.011233998537063599) - XCTAssertEqual( - resultReal.sum().item(Float.self), 6740.3994140625, - accuracy: 134.80798828125) - XCTAssertEqual(resultImaginary.shape, [100, 120]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), 0.5033472180366516, - accuracy: 0.010066944360733033) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), 6040.1669921875, - accuracy: 120.80333984375) - } - - func testFft_3() { - MLXRandom.seed(603) - let r = MLXRandom.uniform(0.0 ..< 1.0, [100, 100]) - XCTAssertEqual(r.shape, [100, 100]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.4938477873802185, - accuracy: 0.00987695574760437) - XCTAssertEqual( - r.sum().item(Float.self), 4938.47802734375, - accuracy: 98.769560546875) - let i = MLXRandom.uniform(0.0 ..< 1.0, [100, 100]) - XCTAssertEqual(i.shape, [100, 100]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.5019298791885376, - accuracy: 0.010038597583770752) - XCTAssertEqual( - i.sum().item(Float.self), 5019.298828125, - accuracy: 100.3859765625) - let c = r + i.asImaginary() - let result = fft(c, axis: 0, stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [100, 100]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 0.5286791324615479, - accuracy: 0.010573582649230957) - XCTAssertEqual( - resultReal.sum().item(Float.self), 5286.79150390625, - accuracy: 105.735830078125) - XCTAssertEqual(resultImaginary.shape, [100, 100]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), 0.5054363012313843, - accuracy: 0.010108726024627685) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), 5054.36328125, - accuracy: 101.087265625) - } - - func testIfft_() { - MLXRandom.seed(845) - let r = MLXRandom.uniform(0.0 ..< 1.0, [100]) - XCTAssertEqual(r.shape, [100]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.4823684096336365, - accuracy: 0.00964736819267273) - XCTAssertEqual( - r.sum().item(Float.self), 48.23684310913086, - accuracy: 0.9647368621826172) - let i = MLXRandom.uniform(0.0 ..< 1.0, [100]) - XCTAssertEqual(i.shape, [100]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.47583281993865967, - accuracy: 0.009516656398773193) - XCTAssertEqual( - i.sum().item(Float.self), 47.583282470703125, - accuracy: 0.9516656494140625) - let c = r + i.asImaginary() - let result = ifft(c, stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [100]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 0.006908764597028494, - accuracy: 0.00013817529194056987) - XCTAssertEqual( - resultReal.sum().item(Float.self), 0.6908764839172363, - accuracy: 0.013817529678344726) - XCTAssertEqual(resultImaginary.shape, [100]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), 0.005689740646630526, - accuracy: 0.00011379481293261051) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), 0.5689740777015686, - accuracy: 0.011379481554031371) - } - - func testIfft_1() { - MLXRandom.seed(972) - let r = MLXRandom.uniform(0.0 ..< 1.0, [100]) - XCTAssertEqual(r.shape, [100]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.46860700845718384, - accuracy: 0.009372140169143678) - XCTAssertEqual( - r.sum().item(Float.self), 46.86070251464844, - accuracy: 0.9372140502929688) - let i = MLXRandom.uniform(0.0 ..< 1.0, [100]) - XCTAssertEqual(i.shape, [100]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.4761311113834381, - accuracy: 0.009522622227668762) - XCTAssertEqual( - i.sum().item(Float.self), 47.61311340332031, - accuracy: 0.9522622680664062) - let c = r + i.asImaginary() - let result = ifft(c, n: 80, stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [80]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 0.005004124250262976, - accuracy: 0.00010008248500525951) - XCTAssertEqual( - resultReal.sum().item(Float.self), 0.40032994747161865, - accuracy: 0.008006598949432373) - XCTAssertEqual(resultImaginary.shape, [80]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), 0.0003413451777305454, - accuracy: 6.826903554610908e-06) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), 0.02730761468410492, - accuracy: 0.0005461522936820984) - } - - func testIfft_2() { - MLXRandom.seed(429) - let r = MLXRandom.uniform(0.0 ..< 1.0, [100]) - XCTAssertEqual(r.shape, [100]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.4427134692668915, - accuracy: 0.008854269385337829) - XCTAssertEqual( - r.sum().item(Float.self), 44.27134704589844, - accuracy: 0.8854269409179688) - let i = MLXRandom.uniform(0.0 ..< 1.0, [100]) - XCTAssertEqual(i.shape, [100]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.46247541904449463, - accuracy: 0.009249508380889893) - XCTAssertEqual( - i.sum().item(Float.self), 46.24754333496094, - accuracy: 0.9249508666992188) - let c = r + i.asImaginary() - let result = ifft(c, n: 120, stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [120]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 0.0033431013580411673, - accuracy: 6.686202716082335e-05) - XCTAssertEqual( - resultReal.sum().item(Float.self), 0.40117213129997253, - accuracy: 0.00802344262599945) - XCTAssertEqual(resultImaginary.shape, [120]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), 0.0048831733874976635, - accuracy: 9.766346774995328e-05) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), 0.5859807729721069, - accuracy: 0.01171961545944214) - } - - func testIfft_3() { - MLXRandom.seed(593) - let r = MLXRandom.uniform(0.0 ..< 1.0, [100]) - XCTAssertEqual(r.shape, [100]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.4806605279445648, - accuracy: 0.009613210558891297) - XCTAssertEqual( - r.sum().item(Float.self), 48.06605529785156, - accuracy: 0.9613211059570312) - let i = MLXRandom.uniform(0.0 ..< 1.0, [100]) - XCTAssertEqual(i.shape, [100]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.48936721682548523, - accuracy: 0.009787344336509705) - XCTAssertEqual( - i.sum().item(Float.self), 48.93672180175781, - accuracy: 0.9787344360351563) - let c = r + i.asImaginary() - let result = ifft(c, axis: 0, stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [100]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 0.009695236571133137, - accuracy: 0.00019390473142266273) - XCTAssertEqual( - resultReal.sum().item(Float.self), 0.9695236682891846, - accuracy: 0.019390473365783693) - XCTAssertEqual(resultImaginary.shape, [100]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), 0.008965172804892063, - accuracy: 0.00017930345609784126) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), 0.896517276763916, - accuracy: 0.017930345535278322) - } - - func testRfft_() { - MLXRandom.seed(281) - let r = MLXRandom.uniform(0.0 ..< 1.0, [100]) - XCTAssertEqual(r.shape, [100]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.512880265712738, - accuracy: 0.01025760531425476) - XCTAssertEqual( - r.sum().item(Float.self), 51.28802490234375, - accuracy: 1.025760498046875) - let i = MLXRandom.uniform(0.0 ..< 1.0, [100]) - XCTAssertEqual(i.shape, [100]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.5312841534614563, - accuracy: 0.010625683069229126) - XCTAssertEqual( - i.sum().item(Float.self), 53.12841796875, - accuracy: 1.062568359375) - let c = r + i.asImaginary() - let result = rfft(c, stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [51]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 1.0831178426742554, - accuracy: 0.021662356853485106) - XCTAssertEqual( - resultReal.sum().item(Float.self), 55.23900604248047, - accuracy: 1.1047801208496093) - XCTAssertEqual(resultImaginary.shape, [51]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), 0.31754177808761597, - accuracy: 0.006350835561752319) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), 16.194629669189453, - accuracy: 0.32389259338378906) - } - - func testRfft_1() { - MLXRandom.seed(461) - let r = MLXRandom.uniform(0.0 ..< 1.0, [100]) - XCTAssertEqual(r.shape, [100]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.4627799987792969, - accuracy: 0.009255599975585938) - XCTAssertEqual( - r.sum().item(Float.self), 46.27799987792969, - accuracy: 0.9255599975585938) - let i = MLXRandom.uniform(0.0 ..< 1.0, [100]) - XCTAssertEqual(i.shape, [100]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.5184644460678101, - accuracy: 0.010369288921356202) - XCTAssertEqual( - i.sum().item(Float.self), 51.84644317626953, - accuracy: 1.0369288635253906) - let c = r + i.asImaginary() - let result = rfft(c, n: 80, stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [41]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 1.3340529203414917, - accuracy: 0.026681058406829834) - XCTAssertEqual( - resultReal.sum().item(Float.self), 54.696170806884766, - accuracy: 1.0939234161376954) - XCTAssertEqual(resultImaginary.shape, [41]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), 0.09739136695861816, - accuracy: 0.0019478273391723632) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), 3.993046283721924, - accuracy: 0.07986092567443848) - } - - func testRfft_2() { - MLXRandom.seed(504) - let r = MLXRandom.uniform(0.0 ..< 1.0, [100]) - XCTAssertEqual(r.shape, [100]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.47210997343063354, - accuracy: 0.009442199468612671) - XCTAssertEqual( - r.sum().item(Float.self), 47.21099853515625, - accuracy: 0.944219970703125) - let i = MLXRandom.uniform(0.0 ..< 1.0, [100]) - XCTAssertEqual(i.shape, [100]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.4813127815723419, - accuracy: 0.009626255631446838) - XCTAssertEqual( - i.sum().item(Float.self), 48.13127899169922, - accuracy: 0.9626255798339844) - let c = r + i.asImaginary() - let result = rfft(c, n: 120, stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [61]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 0.421972393989563, - accuracy: 0.00843944787979126) - XCTAssertEqual( - resultReal.sum().item(Float.self), 25.740318298339844, - accuracy: 0.5148063659667969) - XCTAssertEqual(resultImaginary.shape, [61]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), -0.24379387497901917, - accuracy: -0.004875877499580384) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), -14.871427536010742, - accuracy: -0.29742855072021485) - } - - func testRfft_3() { - MLXRandom.seed(676) - let r = MLXRandom.uniform(0.0 ..< 1.0, [100]) - XCTAssertEqual(r.shape, [100]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.4711613357067108, - accuracy: 0.009423226714134217) - XCTAssertEqual( - r.sum().item(Float.self), 47.11613464355469, - accuracy: 0.9423226928710937) - let i = MLXRandom.uniform(0.0 ..< 1.0, [100]) - XCTAssertEqual(i.shape, [100]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.47868087887763977, - accuracy: 0.009573617577552795) - XCTAssertEqual( - i.sum().item(Float.self), 47.86808776855469, - accuracy: 0.9573617553710938) - let c = r + i.asImaginary() - let result = rfft(c, axis: 0, stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [51]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 0.5287001729011536, - accuracy: 0.010574003458023071) - XCTAssertEqual( - resultReal.sum().item(Float.self), 26.963706970214844, - accuracy: 0.5392741394042969) - XCTAssertEqual(resultImaginary.shape, [51]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), 0.25333887338638306, - accuracy: 0.005066777467727661) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), 12.920282363891602, - accuracy: 0.25840564727783205) - } - - func testIrfft_() { - MLXRandom.seed(656) - let r = MLXRandom.uniform(0.0 ..< 1.0, [100]) - XCTAssertEqual(r.shape, [100]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.476201593875885, - accuracy: 0.0095240318775177) - XCTAssertEqual( - r.sum().item(Float.self), 47.62015914916992, - accuracy: 0.9524031829833984) - let i = MLXRandom.uniform(0.0 ..< 1.0, [100]) - XCTAssertEqual(i.shape, [100]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.5490368008613586, - accuracy: 0.010980736017227172) - XCTAssertEqual( - i.sum().item(Float.self), 54.90367889404297, - accuracy: 1.0980735778808595) - let c = r + i.asImaginary() - let result = irfft(c, stream: .cpu) - XCTAssertEqual(result.shape, [198]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.0028988695703446865, - accuracy: 5.797739140689373e-05) - XCTAssertEqual( - result.sum().item(Float.self), 0.5739761590957642, - accuracy: 0.011479523181915283) - } - - func testIrfft_1() { - MLXRandom.seed(717) - let r = MLXRandom.uniform(0.0 ..< 1.0, [100]) - XCTAssertEqual(r.shape, [100]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.5386900305747986, - accuracy: 0.010773800611495972) - XCTAssertEqual( - r.sum().item(Float.self), 53.86900329589844, - accuracy: 1.0773800659179689) - let i = MLXRandom.uniform(0.0 ..< 1.0, [100]) - XCTAssertEqual(i.shape, [100]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.4603566527366638, - accuracy: 0.009207133054733276) - XCTAssertEqual( - i.sum().item(Float.self), 46.035667419433594, - accuracy: 0.9207133483886719) - let c = r + i.asImaginary() - let result = irfft(c, n: 80, stream: .cpu) - XCTAssertEqual(result.shape, [80]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.0038530919700860977, - accuracy: 7.706183940172196e-05) - XCTAssertEqual( - result.sum().item(Float.self), 0.3082473576068878, - accuracy: 0.006164947152137757) - } - - func testIrfft_2() { - MLXRandom.seed(938) - let r = MLXRandom.uniform(0.0 ..< 1.0, [100]) - XCTAssertEqual(r.shape, [100]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.5083585381507874, - accuracy: 0.010167170763015747) - XCTAssertEqual( - r.sum().item(Float.self), 50.835853576660156, - accuracy: 1.0167170715332032) - let i = MLXRandom.uniform(0.0 ..< 1.0, [100]) - XCTAssertEqual(i.shape, [100]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.4999184012413025, - accuracy: 0.00999836802482605) - XCTAssertEqual( - i.sum().item(Float.self), 49.99184036254883, - accuracy: 0.9998368072509766) - let c = r + i.asImaginary() - let result = irfft(c, n: 120, stream: .cpu) - XCTAssertEqual(result.shape, [120]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.0034326675813645124, - accuracy: 6.865335162729025e-05) - XCTAssertEqual( - result.sum().item(Float.self), 0.41192010045051575, - accuracy: 0.008238402009010316) - } - - func testIrfft_3() { - MLXRandom.seed(812) - let r = MLXRandom.uniform(0.0 ..< 1.0, [100]) - XCTAssertEqual(r.shape, [100]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.4437080919742584, - accuracy: 0.008874161839485169) - XCTAssertEqual( - r.sum().item(Float.self), 44.370811462402344, - accuracy: 0.8874162292480469) - let i = MLXRandom.uniform(0.0 ..< 1.0, [100]) - XCTAssertEqual(i.shape, [100]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.5141764879226685, - accuracy: 0.01028352975845337) - XCTAssertEqual( - i.sum().item(Float.self), 51.41764831542969, - accuracy: 1.0283529663085937) - let c = r + i.asImaginary() - let result = irfft(c, axis: 0, stream: .cpu) - XCTAssertEqual(result.shape, [198]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.004781102295964956, - accuracy: 9.562204591929913e-05) - XCTAssertEqual( - result.sum().item(Float.self), 0.9466582536697388, - accuracy: 0.018933165073394775) - } - - func testFft2_() { - MLXRandom.seed(365) - let r = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(r.shape, [8, 8, 8]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.5127536058425903, - accuracy: 0.010255072116851807) - XCTAssertEqual( - r.sum().item(Float.self), 262.52984619140625, - accuracy: 5.250596923828125) - let i = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(i.shape, [8, 8, 8]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.5092339515686035, - accuracy: 0.01018467903137207) - XCTAssertEqual( - i.sum().item(Float.self), 260.727783203125, - accuracy: 5.2145556640625) - let c = r + i.asImaginary() - let result = fft2(c, stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [8, 8, 8]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 0.3378489017486572, - accuracy: 0.0067569780349731445) - XCTAssertEqual( - resultReal.sum().item(Float.self), 172.9786376953125, - accuracy: 3.45957275390625) - XCTAssertEqual(resultImaginary.shape, [8, 8, 8]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), 0.4062993824481964, - accuracy: 0.008125987648963929) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), 208.02528381347656, - accuracy: 4.160505676269532) - } - - func testFft2_1() { - MLXRandom.seed(84) - let r = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(r.shape, [8, 8, 8]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.5123671293258667, - accuracy: 0.010247342586517334) - XCTAssertEqual( - r.sum().item(Float.self), 262.33197021484375, - accuracy: 5.246639404296875) - let i = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(i.shape, [8, 8, 8]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.502106249332428, - accuracy: 0.01004212498664856) - XCTAssertEqual( - i.sum().item(Float.self), 257.0783996582031, - accuracy: 5.141567993164062) - let c = r + i.asImaginary() - let result = fft2(c, s: [3, 4], stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [8, 3, 4]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 0.5198299884796143, - accuracy: 0.010396599769592285) - XCTAssertEqual( - resultReal.sum().item(Float.self), 49.9036750793457, - accuracy: 0.9980735015869141) - XCTAssertEqual(resultImaginary.shape, [8, 3, 4]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), 0.5818836688995361, - accuracy: 0.011637673377990723) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), 55.8608283996582, - accuracy: 1.117216567993164) - } - - func testFft2_2() { - MLXRandom.seed(332) - let r = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(r.shape, [8, 8, 8]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.4954897463321686, - accuracy: 0.009909794926643371) - XCTAssertEqual( - r.sum().item(Float.self), 253.6907501220703, - accuracy: 5.073815002441406) - let i = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(i.shape, [8, 8, 8]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.5005261898040771, - accuracy: 0.010010523796081543) - XCTAssertEqual( - i.sum().item(Float.self), 256.2694091796875, - accuracy: 5.12538818359375) - let c = r + i.asImaginary() - let result = fft2(c, axes: [0, 2], stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [8, 8, 8]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 0.3796372711658478, - accuracy: 0.007592745423316956) - XCTAssertEqual( - resultReal.sum().item(Float.self), 194.37428283691406, - accuracy: 3.8874856567382814) - XCTAssertEqual(resultImaginary.shape, [8, 8, 8]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), 0.6848205327987671, - accuracy: 0.013696410655975343) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), 350.62811279296875, - accuracy: 7.012562255859375) - } - - func testFft2_3() { - MLXRandom.seed(627) - let r = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(r.shape, [8, 8, 8]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.48098403215408325, - accuracy: 0.009619680643081665) - XCTAssertEqual( - r.sum().item(Float.self), 246.26382446289062, - accuracy: 4.925276489257812) - let i = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(i.shape, [8, 8, 8]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.49524593353271484, - accuracy: 0.009904918670654296) - XCTAssertEqual( - i.sum().item(Float.self), 253.56591796875, - accuracy: 5.071318359375) - let c = r + i.asImaginary() - let result = fft2(c, s: [10, 5], axes: [2, 1], stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [8, 5, 10]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 0.39667367935180664, - accuracy: 0.007933473587036133) - XCTAssertEqual( - resultReal.sum().item(Float.self), 158.6694793701172, - accuracy: 3.173389587402344) - XCTAssertEqual(resultImaginary.shape, [8, 5, 10]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), 0.46095579862594604, - accuracy: 0.009219115972518921) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), 184.38232421875, - accuracy: 3.687646484375) - } - - func testIfft2_() { - MLXRandom.seed(118) - let r = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(r.shape, [8, 8, 8]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.4939275085926056, - accuracy: 0.009878550171852112) - XCTAssertEqual( - r.sum().item(Float.self), 252.89088439941406, - accuracy: 5.0578176879882815) - let i = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(i.shape, [8, 8, 8]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.5031070709228516, - accuracy: 0.010062141418457031) - XCTAssertEqual( - i.sum().item(Float.self), 257.5908203125, - accuracy: 5.15181640625) - let c = r + i.asImaginary() - let result = ifft2(c, stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [8, 8, 8]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 0.009310859255492687, - accuracy: 0.00018621718510985374) - XCTAssertEqual( - resultReal.sum().item(Float.self), 4.767159938812256, - accuracy: 0.09534319877624511) - XCTAssertEqual(resultImaginary.shape, [8, 8, 8]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), 0.006652233190834522, - accuracy: 0.00013304466381669044) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), 3.4059433937072754, - accuracy: 0.0681188678741455) - } - - func testIfft2_1() { - MLXRandom.seed(498) - let r = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(r.shape, [8, 8, 8]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.5009815096855164, - accuracy: 0.010019630193710327) - XCTAssertEqual( - r.sum().item(Float.self), 256.5025329589844, - accuracy: 5.1300506591796875) - let i = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(i.shape, [8, 8, 8]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.49243441224098206, - accuracy: 0.009848688244819642) - XCTAssertEqual( - i.sum().item(Float.self), 252.1264190673828, - accuracy: 5.0425283813476565) - let c = r + i.asImaginary() - let result = ifft2(c, s: [3, 4], stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [8, 3, 4]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 0.055086322128772736, - accuracy: 0.0011017264425754547) - XCTAssertEqual( - resultReal.sum().item(Float.self), 5.2882866859436035, - accuracy: 0.10576573371887207) - XCTAssertEqual(resultImaginary.shape, [8, 3, 4]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), 0.052069030702114105, - accuracy: 0.001041380614042282) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), 4.998626708984375, - accuracy: 0.0999725341796875) - } - - func testIfft2_2() { - MLXRandom.seed(601) - let r = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(r.shape, [8, 8, 8]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.5043255686759949, - accuracy: 0.010086511373519898) - XCTAssertEqual( - r.sum().item(Float.self), 258.2146911621094, - accuracy: 5.164293823242188) - let i = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(i.shape, [8, 8, 8]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.5282406806945801, - accuracy: 0.010564813613891602) - XCTAssertEqual( - i.sum().item(Float.self), 270.459228515625, - accuracy: 5.4091845703125) - let c = r + i.asImaginary() - let result = ifft2(c, axes: [0, 2], stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [8, 8, 8]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 0.005900253541767597, - accuracy: 0.00011800507083535194) - XCTAssertEqual( - resultReal.sum().item(Float.self), 3.0209298133850098, - accuracy: 0.06041859626770019) - XCTAssertEqual(resultImaginary.shape, [8, 8, 8]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), 0.009091435000300407, - accuracy: 0.00018182870000600816) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), 4.654814720153809, - accuracy: 0.09309629440307618) - } - - func testIfft2_3() { - MLXRandom.seed(645) - let r = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(r.shape, [8, 8, 8]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.5046442151069641, - accuracy: 0.010092884302139282) - XCTAssertEqual( - r.sum().item(Float.self), 258.3778381347656, - accuracy: 5.1675567626953125) - let i = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(i.shape, [8, 8, 8]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.4897768199443817, - accuracy: 0.009795536398887635) - XCTAssertEqual( - i.sum().item(Float.self), 250.76573181152344, - accuracy: 5.015314636230469) - let c = r + i.asImaginary() - let result = ifft2(c, s: [10, 5], axes: [2, 1], stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [8, 5, 10]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 0.011821601539850235, - accuracy: 0.0002364320307970047) - XCTAssertEqual( - resultReal.sum().item(Float.self), 4.728640556335449, - accuracy: 0.09457281112670898) - XCTAssertEqual(resultImaginary.shape, [8, 5, 10]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), 0.007806990761309862, - accuracy: 0.00015613981522619725) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), 3.1227962970733643, - accuracy: 0.06245592594146729) - } - - func testFftn_() { - MLXRandom.seed(343) - let r = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(r.shape, [8, 8, 8]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.4879269003868103, - accuracy: 0.009758538007736206) - XCTAssertEqual( - r.sum().item(Float.self), 249.81857299804688, - accuracy: 4.996371459960938) - let i = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(i.shape, [8, 8, 8]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.49864500761032104, - accuracy: 0.009972900152206421) - XCTAssertEqual( - i.sum().item(Float.self), 255.30624389648438, - accuracy: 5.1061248779296875) - let c = r + i.asImaginary() - let result = fftn(c, stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [8, 8, 8]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 0.8467151522636414, - accuracy: 0.01693430304527283) - XCTAssertEqual( - resultReal.sum().item(Float.self), 433.5181579589844, - accuracy: 8.670363159179688) - XCTAssertEqual(resultImaginary.shape, [8, 8, 8]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), 0.28999343514442444, - accuracy: 0.005799868702888489) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), 148.4766387939453, - accuracy: 2.9695327758789065) - } - - func testFftn_1() { - MLXRandom.seed(865) - let r = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(r.shape, [8, 8, 8]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.506011962890625, - accuracy: 0.0101202392578125) - XCTAssertEqual( - r.sum().item(Float.self), 259.078125, - accuracy: 5.1815625) - let i = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(i.shape, [8, 8, 8]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.4955616891384125, - accuracy: 0.00991123378276825) - XCTAssertEqual( - i.sum().item(Float.self), 253.7275848388672, - accuracy: 5.074551696777344) - let c = r + i.asImaginary() - let result = fftn(c, s: [3, 4], stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [8, 3, 4]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 0.3270348012447357, - accuracy: 0.006540696024894714) - XCTAssertEqual( - resultReal.sum().item(Float.self), 31.395339965820312, - accuracy: 0.6279067993164062) - XCTAssertEqual(resultImaginary.shape, [8, 3, 4]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), 0.4940086305141449, - accuracy: 0.009880172610282898) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), 47.424827575683594, - accuracy: 0.9484965515136718) - } - - func testFftn_2() { - MLXRandom.seed(194) - let r = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(r.shape, [8, 8, 8]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.48457071185112, - accuracy: 0.0096914142370224) - XCTAssertEqual( - r.sum().item(Float.self), 248.10020446777344, - accuracy: 4.962004089355469) - let i = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(i.shape, [8, 8, 8]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.4813538193702698, - accuracy: 0.009627076387405396) - XCTAssertEqual( - i.sum().item(Float.self), 246.45315551757812, - accuracy: 4.929063110351563) - let c = r + i.asImaginary() - let result = fftn(c, axes: [0, 2], stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [8, 8, 8]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 0.6415871381759644, - accuracy: 0.012831742763519288) - XCTAssertEqual( - resultReal.sum().item(Float.self), 328.49261474609375, - accuracy: 6.569852294921875) - XCTAssertEqual(resultImaginary.shape, [8, 8, 8]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), 0.5985268950462341, - accuracy: 0.011970537900924684) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), 306.4457702636719, - accuracy: 6.128915405273438) - } - - func testFftn_3() { - MLXRandom.seed(248) - let r = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(r.shape, [8, 8, 8]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.48458826541900635, - accuracy: 0.009691765308380127) - XCTAssertEqual( - r.sum().item(Float.self), 248.10919189453125, - accuracy: 4.962183837890625) - let i = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(i.shape, [8, 8, 8]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.5036082863807678, - accuracy: 0.010072165727615356) - XCTAssertEqual( - i.sum().item(Float.self), 257.8474426269531, - accuracy: 5.156948852539062) - let c = r + i.asImaginary() - let result = fftn(c, s: [10, 5], axes: [2, 1], stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [8, 5, 10]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 0.44054391980171204, - accuracy: 0.008810878396034241) - XCTAssertEqual( - resultReal.sum().item(Float.self), 176.2175750732422, - accuracy: 3.5243515014648437) - XCTAssertEqual(resultImaginary.shape, [8, 5, 10]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), 0.5384612679481506, - accuracy: 0.010769225358963012) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), 215.384521484375, - accuracy: 4.3076904296875) - } - - func testIfftn_() { - MLXRandom.seed(16) - let r = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(r.shape, [8, 8, 8]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.4910537600517273, - accuracy: 0.009821075201034547) - XCTAssertEqual( - r.sum().item(Float.self), 251.41952514648438, - accuracy: 5.028390502929688) - let i = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(i.shape, [8, 8, 8]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.4993511438369751, - accuracy: 0.009987022876739502) - XCTAssertEqual( - i.sum().item(Float.self), 255.66778564453125, - accuracy: 5.113355712890625) - let c = r + i.asImaginary() - let result = ifftn(c, stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [8, 8, 8]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 0.000745800556614995, - accuracy: 1.49160111322999e-05) - XCTAssertEqual( - resultReal.sum().item(Float.self), 0.38184988498687744, - accuracy: 0.007636997699737549) - XCTAssertEqual(resultImaginary.shape, [8, 8, 8]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), 0.00013762176968157291, - accuracy: 2.7524353936314585e-06) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), 0.07046234607696533, - accuracy: 0.0014092469215393067) - } - - func testIfftn_1() { - MLXRandom.seed(749) - let r = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(r.shape, [8, 8, 8]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.5013963580131531, - accuracy: 0.010027927160263062) - XCTAssertEqual( - r.sum().item(Float.self), 256.7149353027344, - accuracy: 5.134298706054688) - let i = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(i.shape, [8, 8, 8]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.49108564853668213, - accuracy: 0.009821712970733643) - XCTAssertEqual( - i.sum().item(Float.self), 251.43585205078125, - accuracy: 5.028717041015625) - let c = r + i.asImaginary() - let result = ifftn(c, s: [3, 4], stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [8, 3, 4]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 0.029997922480106354, - accuracy: 0.0005999584496021271) - XCTAssertEqual( - resultReal.sum().item(Float.self), 2.87980055809021, - accuracy: 0.0575960111618042) - XCTAssertEqual(resultImaginary.shape, [8, 3, 4]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), 0.04324401170015335, - accuracy: 0.0008648802340030671) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), 4.151424884796143, - accuracy: 0.08302849769592285) - } - - func testIfftn_2() { - MLXRandom.seed(277) - let r = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(r.shape, [8, 8, 8]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.5160338878631592, - accuracy: 0.010320677757263183) - XCTAssertEqual( - r.sum().item(Float.self), 264.2093505859375, - accuracy: 5.28418701171875) - let i = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(i.shape, [8, 8, 8]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.48932191729545593, - accuracy: 0.009786438345909119) - XCTAssertEqual( - i.sum().item(Float.self), 250.53282165527344, - accuracy: 5.010656433105469) - let c = r + i.asImaginary() - let result = ifftn(c, axes: [0, 2], stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [8, 8, 8]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 0.007633390370756388, - accuracy: 0.00015266780741512775) - XCTAssertEqual( - resultReal.sum().item(Float.self), 3.9082958698272705, - accuracy: 0.07816591739654541) - XCTAssertEqual(resultImaginary.shape, [8, 8, 8]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), 0.010958677157759666, - accuracy: 0.00021917354315519335) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), 5.610842704772949, - accuracy: 0.11221685409545899) - } - - func testIfftn_3() { - MLXRandom.seed(119) - let r = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(r.shape, [8, 8, 8]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.5159510970115662, - accuracy: 0.010319021940231323) - XCTAssertEqual( - r.sum().item(Float.self), 264.1669616699219, - accuracy: 5.283339233398437) - let i = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(i.shape, [8, 8, 8]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.5054599046707153, - accuracy: 0.010109198093414307) - XCTAssertEqual( - i.sum().item(Float.self), 258.79547119140625, - accuracy: 5.175909423828125) - let c = r + i.asImaginary() - let result = ifftn(c, s: [10, 5], axes: [2, 1], stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [8, 5, 10]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 0.008541149087250233, - accuracy: 0.00017082298174500465) - XCTAssertEqual( - resultReal.sum().item(Float.self), 3.416459560394287, - accuracy: 0.06832919120788575) - XCTAssertEqual(resultImaginary.shape, [8, 5, 10]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), 0.013507889583706856, - accuracy: 0.00027015779167413714) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), 5.40315580368042, - accuracy: 0.1080631160736084) - } - - func testRfft2_() { - MLXRandom.seed(722) - let r = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(r.shape, [8, 8, 8]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.4949501156806946, - accuracy: 0.009899002313613892) - XCTAssertEqual( - r.sum().item(Float.self), 253.41445922851562, - accuracy: 5.068289184570313) - let i = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(i.shape, [8, 8, 8]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.4984145760536194, - accuracy: 0.009968291521072387) - XCTAssertEqual( - i.sum().item(Float.self), 255.18826293945312, - accuracy: 5.103765258789062) - let c = r + i.asImaginary() - let result = rfft2(c, stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [8, 8, 5]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 0.8082265853881836, - accuracy: 0.016164531707763673) - XCTAssertEqual( - resultReal.sum().item(Float.self), 258.63250732421875, - accuracy: 5.172650146484375) - XCTAssertEqual(resultImaginary.shape, [8, 8, 5]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), 0.03572067618370056, - accuracy: 0.0007144135236740113) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), 11.43061637878418, - accuracy: 0.2286123275756836) - } - - func testRfft2_1() { - MLXRandom.seed(225) - let r = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(r.shape, [8, 8, 8]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.5058765411376953, - accuracy: 0.010117530822753906) - XCTAssertEqual( - r.sum().item(Float.self), 259.0087890625, - accuracy: 5.18017578125) - let i = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(i.shape, [8, 8, 8]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.4878406226634979, - accuracy: 0.009756812453269958) - XCTAssertEqual( - i.sum().item(Float.self), 249.77439880371094, - accuracy: 4.995487976074219) - let c = r + i.asImaginary() - let result = rfft2(c, s: [3, 4], stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [8, 3, 3]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 0.7214547991752625, - accuracy: 0.014429095983505249) - XCTAssertEqual( - resultReal.sum().item(Float.self), 51.94474411010742, - accuracy: 1.0388948822021484) - XCTAssertEqual(resultImaginary.shape, [8, 3, 3]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), -0.03948277235031128, - accuracy: -0.0007896554470062257) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), -2.842759609222412, - accuracy: -0.05685519218444824) - } - - func testRfft2_2() { - MLXRandom.seed(380) - let r = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(r.shape, [8, 8, 8]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.4951004981994629, - accuracy: 0.009902009963989258) - XCTAssertEqual( - r.sum().item(Float.self), 253.491455078125, - accuracy: 5.0698291015625) - let i = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(i.shape, [8, 8, 8]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.4919356107711792, - accuracy: 0.009838712215423585) - XCTAssertEqual( - i.sum().item(Float.self), 251.87103271484375, - accuracy: 5.037420654296875) - let c = r + i.asImaginary() - let result = rfft2(c, axes: [0, 2], stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [8, 8, 5]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 0.86641925573349, - accuracy: 0.0173283851146698) - XCTAssertEqual( - resultReal.sum().item(Float.self), 277.254150390625, - accuracy: 5.5450830078125) - XCTAssertEqual(resultImaginary.shape, [8, 8, 5]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), 0.006711974740028381, - accuracy: 0.00013423949480056762) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), 2.147831916809082, - accuracy: 0.04295663833618164) - } - - func testRfft2_3() { - MLXRandom.seed(813) - let r = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(r.shape, [8, 8, 8]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.4799543023109436, - accuracy: 0.009599086046218872) - XCTAssertEqual( - r.sum().item(Float.self), 245.73660278320312, - accuracy: 4.914732055664063) - let i = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(i.shape, [8, 8, 8]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.4987284541130066, - accuracy: 0.009974569082260132) - XCTAssertEqual( - i.sum().item(Float.self), 255.34896850585938, - accuracy: 5.1069793701171875) - let c = r + i.asImaginary() - let result = rfft2(c, s: [10, 5], axes: [2, 1], stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [8, 3, 10]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 0.7862035036087036, - accuracy: 0.015724070072174072) - XCTAssertEqual( - resultReal.sum().item(Float.self), 188.68882751464844, - accuracy: 3.7737765502929688) - XCTAssertEqual(resultImaginary.shape, [8, 3, 10]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), 0.11010777205228806, - accuracy: 0.002202155441045761) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), 26.42586326599121, - accuracy: 0.5285172653198242) - } - - func testIrfft2_() { - MLXRandom.seed(174) - let r = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(r.shape, [8, 8, 8]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.514319121837616, - accuracy: 0.01028638243675232) - XCTAssertEqual( - r.sum().item(Float.self), 263.3313903808594, - accuracy: 5.266627807617188) - let i = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(i.shape, [8, 8, 8]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.5133156776428223, - accuracy: 0.010266313552856446) - XCTAssertEqual( - i.sum().item(Float.self), 262.817626953125, - accuracy: 5.2563525390625) - let c = r + i.asImaginary() - let result = irfft2(c, stream: .cpu) - XCTAssertEqual(result.shape, [8, 8, 14]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.005966233555227518, - accuracy: 0.00011932467110455036) - XCTAssertEqual( - result.sum().item(Float.self), 5.345745086669922, - accuracy: 0.10691490173339845) - } - - func testIrfft2_1() { - MLXRandom.seed(340) - let r = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(r.shape, [8, 8, 8]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.5179769992828369, - accuracy: 0.010359539985656738) - XCTAssertEqual( - r.sum().item(Float.self), 265.2042236328125, - accuracy: 5.30408447265625) - let i = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(i.shape, [8, 8, 8]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.5126595497131348, - accuracy: 0.010253190994262695) - XCTAssertEqual( - i.sum().item(Float.self), 262.481689453125, - accuracy: 5.2496337890625) - let c = r + i.asImaginary() - let result = irfft2(c, s: [3, 4], stream: .cpu) - XCTAssertEqual(result.shape, [8, 3, 4]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.047055989503860474, - accuracy: 0.0009411197900772095) - XCTAssertEqual( - result.sum().item(Float.self), 4.5173749923706055, - accuracy: 0.09034749984741211) - } - - func testIrfft2_2() { - MLXRandom.seed(436) - let r = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(r.shape, [8, 8, 8]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.5037882328033447, - accuracy: 0.010075764656066894) - XCTAssertEqual( - r.sum().item(Float.self), 257.9395751953125, - accuracy: 5.15879150390625) - let i = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(i.shape, [8, 8, 8]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.5085217356681824, - accuracy: 0.010170434713363648) - XCTAssertEqual( - i.sum().item(Float.self), 260.3631286621094, - accuracy: 5.207262573242188) - let c = r + i.asImaginary() - let result = irfft2(c, axes: [0, 2], stream: .cpu) - XCTAssertEqual(result.shape, [8, 8, 14]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.004838266409933567, - accuracy: 9.676532819867134e-05) - XCTAssertEqual( - result.sum().item(Float.self), 4.335086345672607, - accuracy: 0.08670172691345215) - } - - func testIrfft2_3() { - MLXRandom.seed(835) - let r = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(r.shape, [8, 8, 8]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.507753849029541, - accuracy: 0.010155076980590821) - XCTAssertEqual( - r.sum().item(Float.self), 259.969970703125, - accuracy: 5.1993994140625) - let i = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(i.shape, [8, 8, 8]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.49466848373413086, - accuracy: 0.009893369674682618) - XCTAssertEqual( - i.sum().item(Float.self), 253.270263671875, - accuracy: 5.0654052734375) - let c = r + i.asImaginary() - let result = irfft2(c, s: [10, 5], axes: [2, 1], stream: .cpu) - XCTAssertEqual(result.shape, [8, 5, 10]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.013026298955082893, - accuracy: 0.0002605259791016579) - XCTAssertEqual( - result.sum().item(Float.self), 5.210519790649414, - accuracy: 0.10421039581298829) - } - - func testRfftn_() { - MLXRandom.seed(63) - let r = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(r.shape, [8, 8, 8]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.5128694772720337, - accuracy: 0.010257389545440674) - XCTAssertEqual( - r.sum().item(Float.self), 262.58917236328125, - accuracy: 5.251783447265625) - let i = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(i.shape, [8, 8, 8]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.4861868619918823, - accuracy: 0.009723737239837646) - XCTAssertEqual( - i.sum().item(Float.self), 248.92767333984375, - accuracy: 4.978553466796875) - let c = r + i.asImaginary() - let result = rfftn(c, stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [8, 8, 5]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 0.5379745960235596, - accuracy: 0.010759491920471192) - XCTAssertEqual( - resultReal.sum().item(Float.self), 172.15187072753906, - accuracy: 3.443037414550781) - XCTAssertEqual(resultImaginary.shape, [8, 8, 5]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), 0.12080971151590347, - accuracy: 0.0024161942303180697) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), 38.65910720825195, - accuracy: 0.773182144165039) - } - - func testRfftn_1() { - MLXRandom.seed(103) - let r = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(r.shape, [8, 8, 8]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.5262448787689209, - accuracy: 0.010524897575378419) - XCTAssertEqual( - r.sum().item(Float.self), 269.4373779296875, - accuracy: 5.38874755859375) - let i = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(i.shape, [8, 8, 8]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.4829203188419342, - accuracy: 0.009658406376838685) - XCTAssertEqual( - i.sum().item(Float.self), 247.2552032470703, - accuracy: 4.945104064941407) - let c = r + i.asImaginary() - let result = rfftn(c, s: [3, 4], stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [8, 3, 3]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 0.5502752065658569, - accuracy: 0.011005504131317139) - XCTAssertEqual( - resultReal.sum().item(Float.self), 39.619815826416016, - accuracy: 0.7923963165283203) - XCTAssertEqual(resultImaginary.shape, [8, 3, 3]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), 0.025733524933457375, - accuracy: 0.0005146704986691475) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), 1.852813720703125, - accuracy: 0.0370562744140625) - } - - func testRfftn_2() { - MLXRandom.seed(801) - let r = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(r.shape, [8, 8, 8]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.50965815782547, - accuracy: 0.010193163156509399) - XCTAssertEqual( - r.sum().item(Float.self), 260.9449768066406, - accuracy: 5.218899536132812) - let i = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(i.shape, [8, 8, 8]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.5044752955436707, - accuracy: 0.010089505910873414) - XCTAssertEqual( - i.sum().item(Float.self), 258.2913513183594, - accuracy: 5.165827026367188) - let c = r + i.asImaginary() - let result = rfftn(c, axes: [0, 2], stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [8, 8, 5]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 0.8939148187637329, - accuracy: 0.017878296375274657) - XCTAssertEqual( - resultReal.sum().item(Float.self), 286.052734375, - accuracy: 5.7210546875) - XCTAssertEqual(resultImaginary.shape, [8, 8, 5]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), -0.05308229848742485, - accuracy: -0.001061645969748497) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), -16.98633575439453, - accuracy: -0.33972671508789065) - } - - func testRfftn_3() { - MLXRandom.seed(149) - let r = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(r.shape, [8, 8, 8]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.5064899921417236, - accuracy: 0.010129799842834472) - XCTAssertEqual( - r.sum().item(Float.self), 259.3228759765625, - accuracy: 5.18645751953125) - let i = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(i.shape, [8, 8, 8]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.4949530363082886, - accuracy: 0.009899060726165771) - XCTAssertEqual( - i.sum().item(Float.self), 253.41595458984375, - accuracy: 5.068319091796875) - let c = r + i.asImaginary() - let result = rfftn(c, s: [10, 5], axes: [2, 1], stream: .cpu) - let resultReal = result.realPart() - let resultImaginary = result.imaginaryPart() - XCTAssertEqual(resultReal.shape, [8, 3, 10]) - XCTAssertEqual(resultReal.dtype, .float32) - XCTAssertEqual( - resultReal.mean().item(Float.self), 0.7501934170722961, - accuracy: 0.015003868341445924) - XCTAssertEqual( - resultReal.sum().item(Float.self), 180.04641723632812, - accuracy: 3.6009283447265625) - XCTAssertEqual(resultImaginary.shape, [8, 3, 10]) - XCTAssertEqual(resultImaginary.dtype, .float32) - XCTAssertEqual( - resultImaginary.mean().item(Float.self), -0.039641283452510834, - accuracy: -0.0007928256690502167) - XCTAssertEqual( - resultImaginary.sum().item(Float.self), -9.513907432556152, - accuracy: -0.19027814865112305) - } - - func testIrfftn_() { - MLXRandom.seed(875) - let r = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(r.shape, [8, 8, 8]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.5080322623252869, - accuracy: 0.010160645246505737) - XCTAssertEqual( - r.sum().item(Float.self), 260.1125183105469, - accuracy: 5.202250366210937) - let i = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(i.shape, [8, 8, 8]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.5002621412277222, - accuracy: 0.010005242824554443) - XCTAssertEqual( - i.sum().item(Float.self), 256.13421630859375, - accuracy: 5.122684326171875) - let c = r + i.asImaginary() - let result = irfftn(c, stream: .cpu) - XCTAssertEqual(result.shape, [8, 8, 14]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 9.236017649527639e-05, - accuracy: 1.8472035299055278e-06) - XCTAssertEqual( - result.sum().item(Float.self), 0.0827547162771225, - accuracy: 0.00165509432554245) - } - - func testIrfftn_1() { - MLXRandom.seed(714) - let r = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(r.shape, [8, 8, 8]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.49393296241760254, - accuracy: 0.00987865924835205) - XCTAssertEqual( - r.sum().item(Float.self), 252.8936767578125, - accuracy: 5.05787353515625) - let i = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(i.shape, [8, 8, 8]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.5138311386108398, - accuracy: 0.010276622772216797) - XCTAssertEqual( - i.sum().item(Float.self), 263.08154296875, - accuracy: 5.261630859375) - let c = r + i.asImaginary() - let result = irfftn(c, s: [3, 4], stream: .cpu) - XCTAssertEqual(result.shape, [8, 3, 4]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.041039690375328064, - accuracy: 0.0008207938075065613) - XCTAssertEqual( - result.sum().item(Float.self), 3.939810276031494, - accuracy: 0.07879620552062988) - } - - func testIrfftn_2() { - MLXRandom.seed(224) - let r = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(r.shape, [8, 8, 8]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.5111634731292725, - accuracy: 0.01022326946258545) - XCTAssertEqual( - r.sum().item(Float.self), 261.7156982421875, - accuracy: 5.23431396484375) - let i = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(i.shape, [8, 8, 8]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.49477437138557434, - accuracy: 0.009895487427711487) - XCTAssertEqual( - i.sum().item(Float.self), 253.32447814941406, - accuracy: 5.066489562988282) - let c = r + i.asImaginary() - let result = irfftn(c, axes: [0, 2], stream: .cpu) - XCTAssertEqual(result.shape, [8, 8, 14]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.00553929852321744, - accuracy: 0.00011078597046434879) - XCTAssertEqual( - result.sum().item(Float.self), 4.9632110595703125, - accuracy: 0.09926422119140625) - } - - func testIrfftn_3() { - MLXRandom.seed(46) - let r = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(r.shape, [8, 8, 8]) - XCTAssertEqual(r.dtype, .float32) - XCTAssertEqual( - r.mean().item(Float.self), 0.5066219568252563, - accuracy: 0.010132439136505127) - XCTAssertEqual( - r.sum().item(Float.self), 259.39044189453125, - accuracy: 5.187808837890625) - let i = MLXRandom.uniform(0.0 ..< 1.0, [8, 8, 8]) - XCTAssertEqual(i.shape, [8, 8, 8]) - XCTAssertEqual(i.dtype, .float32) - XCTAssertEqual( - i.mean().item(Float.self), 0.5006622076034546, - accuracy: 0.010013244152069093) - XCTAssertEqual( - i.sum().item(Float.self), 256.33905029296875, - accuracy: 5.1267810058593755) - let c = r + i.asImaginary() - let result = irfftn(c, s: [10, 5], axes: [2, 1], stream: .cpu) - XCTAssertEqual(result.shape, [8, 5, 10]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.010128638707101345, - accuracy: 0.00020257277414202692) - XCTAssertEqual( - result.sum().item(Float.self), 4.051455497741699, - accuracy: 0.08102910995483399) - } - - func testSGD() { - MLXRandom.seed(836) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.07495629787445068, - accuracy: 0.0014991259574890137) - XCTAssertEqual( - a.sum().item(Float.self), 0.8994755148887634, - accuracy: 0.01798951029777527) - let aGrad = MLXRandom.normal([4, 3]) - XCTAssertEqual(aGrad.shape, [4, 3]) - XCTAssertEqual(aGrad.dtype, .float32) - XCTAssertEqual( - aGrad.mean().item(Float.self), 0.09326401352882385, - accuracy: 0.001865280270576477) - XCTAssertEqual( - aGrad.sum().item(Float.self), 1.1191681623458862, - accuracy: 0.022383363246917726) - let aModel = ModuleParameters(values: ["a": .value(a)]) - let aGradParams = ModuleParameters(values: ["a": .value(aGrad)]) - let result = SGD(learningRate: 0.1).apply(gradients: aGradParams, modelParameters: aModel) - XCTAssertEqual(result[unwrapping: "a"]!.shape, [4, 3]) - XCTAssertEqual(result[unwrapping: "a"]!.dtype, .float32) - XCTAssertEqual( - result[unwrapping: "a"]!.mean().item(Float.self), 0.06562991440296173, - accuracy: 0.0013125982880592346) - XCTAssertEqual( - result[unwrapping: "a"]!.sum().item(Float.self), 0.787558913230896, - accuracy: 0.01575117826461792) - } - - func testSGD1() { - MLXRandom.seed(587) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.10747849941253662, - accuracy: -0.0021495699882507326) - XCTAssertEqual( - a.sum().item(Float.self), -1.2897419929504395, - accuracy: -0.02579483985900879) - let aGrad = MLXRandom.normal([4, 3]) - XCTAssertEqual(aGrad.shape, [4, 3]) - XCTAssertEqual(aGrad.dtype, .float32) - XCTAssertEqual( - aGrad.mean().item(Float.self), -0.32190611958503723, - accuracy: -0.006438122391700745) - XCTAssertEqual( - aGrad.sum().item(Float.self), -3.8628733158111572, - accuracy: -0.07725746631622314) - let aModel = ModuleParameters(values: ["a": .value(a)]) - let aGradParams = ModuleParameters(values: ["a": .value(aGrad)]) - let result = SGD(learningRate: 0.1, momentum: 0.1).apply( - gradients: aGradParams, modelParameters: aModel) - XCTAssertEqual(result[unwrapping: "a"]!.shape, [4, 3]) - XCTAssertEqual(result[unwrapping: "a"]!.dtype, .float32) - XCTAssertEqual( - result[unwrapping: "a"]!.mean().item(Float.self), -0.07528789341449738, - accuracy: -0.0015057578682899475) - XCTAssertEqual( - result[unwrapping: "a"]!.sum().item(Float.self), -0.9034546613693237, - accuracy: -0.018069093227386476) - } - - func testSGD2() { - MLXRandom.seed(649) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.7766839861869812, - accuracy: 0.015533679723739624) - XCTAssertEqual( - a.sum().item(Float.self), 9.320207595825195, - accuracy: 0.18640415191650392) - let aGrad = MLXRandom.normal([4, 3]) - XCTAssertEqual(aGrad.shape, [4, 3]) - XCTAssertEqual(aGrad.dtype, .float32) - XCTAssertEqual( - aGrad.mean().item(Float.self), 0.06737084686756134, - accuracy: 0.0013474169373512267) - XCTAssertEqual( - aGrad.sum().item(Float.self), 0.8084501028060913, - accuracy: 0.016169002056121828) - let aModel = ModuleParameters(values: ["a": .value(a)]) - let aGradParams = ModuleParameters(values: ["a": .value(aGrad)]) - let result = SGD(learningRate: 0.1, momentum: 0.1, dampening: 0.1).apply( - gradients: aGradParams, modelParameters: aModel) - XCTAssertEqual(result[unwrapping: "a"]!.shape, [4, 3]) - XCTAssertEqual(result[unwrapping: "a"]!.dtype, .float32) - XCTAssertEqual( - result[unwrapping: "a"]!.mean().item(Float.self), 0.7706205248832703, - accuracy: 0.015412410497665405) - XCTAssertEqual( - result[unwrapping: "a"]!.sum().item(Float.self), 9.247446060180664, - accuracy: 0.18494892120361328) - } - - func testRMSprop() { - MLXRandom.seed(931) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.22450898587703705, - accuracy: -0.004490179717540741) - XCTAssertEqual( - a.sum().item(Float.self), -2.6941077709198, - accuracy: -0.053882155418396) - let aGrad = MLXRandom.normal([4, 3]) - XCTAssertEqual(aGrad.shape, [4, 3]) - XCTAssertEqual(aGrad.dtype, .float32) - XCTAssertEqual( - aGrad.mean().item(Float.self), 0.24865484237670898, - accuracy: 0.00497309684753418) - XCTAssertEqual( - aGrad.sum().item(Float.self), 2.983858108520508, - accuracy: 0.05967716217041016) - let aModel = ModuleParameters(values: ["a": .value(a)]) - let aGradParams = ModuleParameters(values: ["a": .value(aGrad)]) - let result = RMSprop(learningRate: 0.1).apply( - gradients: aGradParams, modelParameters: aModel) - XCTAssertEqual(result[unwrapping: "a"]!.shape, [4, 3]) - XCTAssertEqual(result[unwrapping: "a"]!.dtype, .float32) - XCTAssertEqual( - result[unwrapping: "a"]!.mean().item(Float.self), -0.22450922429561615, - accuracy: -0.004490184485912323) - XCTAssertEqual( - result[unwrapping: "a"]!.sum().item(Float.self), -2.694110631942749, - accuracy: -0.05388221263885498) - } - - func testAdagrad() { - MLXRandom.seed(958) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.04584333300590515, - accuracy: -0.0009168666601181031) - XCTAssertEqual( - a.sum().item(Float.self), -0.5501199960708618, - accuracy: -0.011002399921417237) - let aGrad = MLXRandom.normal([4, 3]) - XCTAssertEqual(aGrad.shape, [4, 3]) - XCTAssertEqual(aGrad.dtype, .float32) - XCTAssertEqual( - aGrad.mean().item(Float.self), 0.23250393569469452, - accuracy: 0.00465007871389389) - XCTAssertEqual( - aGrad.sum().item(Float.self), 2.7900471687316895, - accuracy: 0.05580094337463379) - let aModel = ModuleParameters(values: ["a": .value(a)]) - let aGradParams = ModuleParameters(values: ["a": .value(aGrad)]) - let result = AdaGrad(learningRate: 0.1).apply( - gradients: aGradParams, modelParameters: aModel) - XCTAssertEqual(result[unwrapping: "a"]!.shape, [4, 3]) - XCTAssertEqual(result[unwrapping: "a"]!.dtype, .float32) - XCTAssertEqual( - result[unwrapping: "a"]!.mean().item(Float.self), -0.06250998377799988, - accuracy: -0.0012501996755599975) - XCTAssertEqual( - result[unwrapping: "a"]!.sum().item(Float.self), -0.7501198053359985, - accuracy: -0.01500239610671997) - } - - func testAdaDelta() { - MLXRandom.seed(547) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.3483370244503021, - accuracy: -0.006966740489006043) - XCTAssertEqual( - a.sum().item(Float.self), -4.180044174194336, - accuracy: -0.08360088348388672) - let aGrad = MLXRandom.normal([4, 3]) - XCTAssertEqual(aGrad.shape, [4, 3]) - XCTAssertEqual(aGrad.dtype, .float32) - XCTAssertEqual( - aGrad.mean().item(Float.self), 0.5226783752441406, - accuracy: 0.010453567504882813) - XCTAssertEqual( - aGrad.sum().item(Float.self), 6.272140026092529, - accuracy: 0.12544280052185058) - let aModel = ModuleParameters(values: ["a": .value(a)]) - let aGradParams = ModuleParameters(values: ["a": .value(aGrad)]) - let result = AdaDelta(learningRate: 0.1).apply( - gradients: aGradParams, modelParameters: aModel) - XCTAssertEqual(result[unwrapping: "a"]!.shape, [4, 3]) - XCTAssertEqual(result[unwrapping: "a"]!.dtype, .float32) - XCTAssertEqual( - result[unwrapping: "a"]!.mean().item(Float.self), -0.348442405462265, - accuracy: -0.0069688481092453) - XCTAssertEqual( - result[unwrapping: "a"]!.sum().item(Float.self), -4.181308746337891, - accuracy: -0.08362617492675782) - } - - func testAdam() { - MLXRandom.seed(616) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.1122930571436882, - accuracy: 0.002245861142873764) - XCTAssertEqual( - a.sum().item(Float.self), 1.347516655921936, - accuracy: 0.02695033311843872) - let aGrad = MLXRandom.normal([4, 3]) - XCTAssertEqual(aGrad.shape, [4, 3]) - XCTAssertEqual(aGrad.dtype, .float32) - XCTAssertEqual( - aGrad.mean().item(Float.self), 0.305597722530365, - accuracy: 0.0061119544506073) - XCTAssertEqual( - aGrad.sum().item(Float.self), 3.66717267036438, - accuracy: 0.0733434534072876) - let aModel = ModuleParameters(values: ["a": .value(a)]) - let aGradParams = ModuleParameters(values: ["a": .value(aGrad)]) - let result = Adam(learningRate: 0.1).apply(gradients: aGradParams, modelParameters: aModel) - XCTAssertEqual(result[unwrapping: "a"]!.shape, [4, 3]) - XCTAssertEqual(result[unwrapping: "a"]!.dtype, .float32) - XCTAssertEqual( - result[unwrapping: "a"]!.mean().item(Float.self), 0.11229278147220612, - accuracy: 0.0022458556294441224) - XCTAssertEqual( - result[unwrapping: "a"]!.sum().item(Float.self), 1.3475133180618286, - accuracy: 0.026950266361236572) - } - - func testAdamW() { - MLXRandom.seed(696) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.3633918762207031, - accuracy: -0.007267837524414063) - XCTAssertEqual( - a.sum().item(Float.self), -4.3607025146484375, - accuracy: -0.08721405029296875) - let aGrad = MLXRandom.normal([4, 3]) - XCTAssertEqual(aGrad.shape, [4, 3]) - XCTAssertEqual(aGrad.dtype, .float32) - XCTAssertEqual( - aGrad.mean().item(Float.self), 0.22175447642803192, - accuracy: 0.0044350895285606385) - XCTAssertEqual( - aGrad.sum().item(Float.self), 2.6610536575317383, - accuracy: 0.05322107315063477) - let aModel = ModuleParameters(values: ["a": .value(a)]) - let aGradParams = ModuleParameters(values: ["a": .value(aGrad)]) - let result = AdamW(learningRate: 0.1).apply(gradients: aGradParams, modelParameters: aModel) - XCTAssertEqual(result[unwrapping: "a"]!.shape, [4, 3]) - XCTAssertEqual(result[unwrapping: "a"]!.dtype, .float32) - XCTAssertEqual( - result[unwrapping: "a"]!.mean().item(Float.self), -0.4684376120567322, - accuracy: -0.009368752241134645) - XCTAssertEqual( - result[unwrapping: "a"]!.sum().item(Float.self), -5.621251106262207, - accuracy: -0.11242502212524415) - } - - func testAdamax() { - MLXRandom.seed(75) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.3039236068725586, - accuracy: -0.006078472137451172) - XCTAssertEqual( - a.sum().item(Float.self), -3.647083282470703, - accuracy: -0.07294166564941407) - let aGrad = MLXRandom.normal([4, 3]) - XCTAssertEqual(aGrad.shape, [4, 3]) - XCTAssertEqual(aGrad.dtype, .float32) - XCTAssertEqual( - aGrad.mean().item(Float.self), -0.24271723628044128, - accuracy: -0.004854344725608826) - XCTAssertEqual( - aGrad.sum().item(Float.self), -2.912606716156006, - accuracy: -0.05825213432312012) - let aModel = ModuleParameters(values: ["a": .value(a)]) - let aGradParams = ModuleParameters(values: ["a": .value(aGrad)]) - let result = Adamax(learningRate: 0.1).apply( - gradients: aGradParams, modelParameters: aModel) - XCTAssertEqual(result[unwrapping: "a"]!.shape, [4, 3]) - XCTAssertEqual(result[unwrapping: "a"]!.dtype, .float32) - XCTAssertEqual( - result[unwrapping: "a"]!.mean().item(Float.self), -0.3039236068725586, - accuracy: -0.006078472137451172) - XCTAssertEqual( - result[unwrapping: "a"]!.sum().item(Float.self), -3.647083282470703, - accuracy: -0.07294166564941407) - } - - func testLion() { - MLXRandom.seed(27) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.1776922345161438, - accuracy: 0.003553844690322876) - XCTAssertEqual( - a.sum().item(Float.self), 2.1323068141937256, - accuracy: 0.042646136283874515) - let aGrad = MLXRandom.normal([4, 3]) - XCTAssertEqual(aGrad.shape, [4, 3]) - XCTAssertEqual(aGrad.dtype, .float32) - XCTAssertEqual( - aGrad.mean().item(Float.self), -0.02118723653256893, - accuracy: -0.00042374473065137863) - XCTAssertEqual( - aGrad.sum().item(Float.self), -0.2542468309402466, - accuracy: -0.005084936618804932) - let aModel = ModuleParameters(values: ["a": .value(a)]) - let aGradParams = ModuleParameters(values: ["a": .value(aGrad)]) - let result = Lion(learningRate: 0.1).apply(gradients: aGradParams, modelParameters: aModel) - XCTAssertEqual(result[unwrapping: "a"]!.shape, [4, 3]) - XCTAssertEqual(result[unwrapping: "a"]!.dtype, .float32) - XCTAssertEqual( - result[unwrapping: "a"]!.mean().item(Float.self), 0.21102556586265564, - accuracy: 0.004220511317253113) - XCTAssertEqual( - result[unwrapping: "a"]!.sum().item(Float.self), 2.532306671142578, - accuracy: 0.05064613342285156) - } - - func testLion1() { - MLXRandom.seed(127) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.18461060523986816, - accuracy: -0.0036922121047973633) - XCTAssertEqual( - a.sum().item(Float.self), -2.215327262878418, - accuracy: -0.04430654525756836) - let aGrad = MLXRandom.normal([4, 3]) - XCTAssertEqual(aGrad.shape, [4, 3]) - XCTAssertEqual(aGrad.dtype, .float32) - XCTAssertEqual( - aGrad.mean().item(Float.self), -0.03600400686264038, - accuracy: -0.0007200801372528076) - XCTAssertEqual( - aGrad.sum().item(Float.self), -0.43204808235168457, - accuracy: -0.008640961647033691) - let aModel = ModuleParameters(values: ["a": .value(a)]) - let aGradParams = ModuleParameters(values: ["a": .value(aGrad)]) - let result = Lion(learningRate: 0.1, weightDecay: 0.1).apply( - gradients: aGradParams, modelParameters: aModel) - XCTAssertEqual(result[unwrapping: "a"]!.shape, [4, 3]) - XCTAssertEqual(result[unwrapping: "a"]!.dtype, .float32) - XCTAssertEqual( - result[unwrapping: "a"]!.mean().item(Float.self), -0.18276450037956238, - accuracy: -0.003655290007591248) - XCTAssertEqual( - result[unwrapping: "a"]!.sum().item(Float.self), -2.193173885345459, - accuracy: -0.04386347770690918) - } - - func testAdafactor() { - MLXRandom.seed(650) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), -0.5207136869430542, - accuracy: -0.010414273738861083) - XCTAssertEqual( - a.sum().item(Float.self), -6.248563766479492, - accuracy: -0.12497127532958985) - let aGrad = MLXRandom.normal([4, 3]) - XCTAssertEqual(aGrad.shape, [4, 3]) - XCTAssertEqual(aGrad.dtype, .float32) - XCTAssertEqual( - aGrad.mean().item(Float.self), 0.4333036541938782, - accuracy: 0.008666073083877564) - XCTAssertEqual( - aGrad.sum().item(Float.self), 5.199643611907959, - accuracy: 0.10399287223815919) - let aModel = ModuleParameters(values: ["a": .value(a)]) - let aGradParams = ModuleParameters(values: ["a": .value(aGrad)]) - let result = Adafactor(learningRate: 0.1).apply( - gradients: aGradParams, modelParameters: aModel) - XCTAssertEqual(result[unwrapping: "a"]!.shape, [4, 3]) - XCTAssertEqual(result[unwrapping: "a"]!.dtype, .float32) - XCTAssertEqual( - result[unwrapping: "a"]!.mean().item(Float.self), -0.5268284678459167, - accuracy: -0.010536569356918336) - XCTAssertEqual( - result[unwrapping: "a"]!.sum().item(Float.self), -6.321941375732422, - accuracy: -0.12643882751464844) - } - - func testAdafactor1() { - MLXRandom.seed(193) - let a = MLXRandom.normal([4, 3]) - XCTAssertEqual(a.shape, [4, 3]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.4008181691169739, - accuracy: 0.008016363382339478) - XCTAssertEqual( - a.sum().item(Float.self), 4.809817790985107, - accuracy: 0.09619635581970215) - let aGrad = MLXRandom.normal([4, 3]) - XCTAssertEqual(aGrad.shape, [4, 3]) - XCTAssertEqual(aGrad.dtype, .float32) - XCTAssertEqual( - aGrad.mean().item(Float.self), 0.21447472274303436, - accuracy: 0.004289494454860688) - XCTAssertEqual( - aGrad.sum().item(Float.self), 2.5736966133117676, - accuracy: 0.05147393226623535) - let aModel = ModuleParameters(values: ["a": .value(a)]) - let aGradParams = ModuleParameters(values: ["a": .value(aGrad)]) - let result = Adafactor(learningRate: 0.1, beta1: 0.1).apply( - gradients: aGradParams, modelParameters: aModel) - XCTAssertEqual(result[unwrapping: "a"]!.shape, [4, 3]) - XCTAssertEqual(result[unwrapping: "a"]!.dtype, .float32) - XCTAssertEqual( - result[unwrapping: "a"]!.mean().item(Float.self), 0.39943069219589233, - accuracy: 0.007988613843917847) - XCTAssertEqual( - result[unwrapping: "a"]!.sum().item(Float.self), 4.793168067932129, - accuracy: 0.09586336135864258) - } - - func testAdafactor2() { - MLXRandom.seed(620) - let a = MLXRandom.uniform(0.0 ..< 1.0, [10]) - XCTAssertEqual(a.shape, [10]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.4890245497226715, - accuracy: 0.00978049099445343) - XCTAssertEqual( - a.sum().item(Float.self), 4.89024543762207, - accuracy: 0.09780490875244141) - let aGrad = MLXRandom.uniform(0.0 ..< 1.0, [10]) - XCTAssertEqual(aGrad.shape, [10]) - XCTAssertEqual(aGrad.dtype, .float32) - XCTAssertEqual( - aGrad.mean().item(Float.self), 0.6818901896476746, - accuracy: 0.013637803792953491) - XCTAssertEqual( - aGrad.sum().item(Float.self), 6.818902015686035, - accuracy: 0.1363780403137207) - let aModel = ModuleParameters(values: ["a": .value(a)]) - let aGradParams = ModuleParameters(values: ["a": .value(aGrad)]) - let result = Adafactor(learningRate: 0.1).apply( - gradients: aGradParams, modelParameters: aModel) - XCTAssertEqual(result[unwrapping: "a"]!.shape, [10]) - XCTAssertEqual(result[unwrapping: "a"]!.dtype, .float32) - XCTAssertEqual( - result[unwrapping: "a"]!.mean().item(Float.self), 0.4835330545902252, - accuracy: 0.009670661091804504) - XCTAssertEqual( - result[unwrapping: "a"]!.sum().item(Float.self), 4.835330486297607, - accuracy: 0.09670660972595214) - } - - func testGLU() { - MLXRandom.seed(850) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 8, 16]) - XCTAssertEqual(a.shape, [2, 8, 16]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.5472526550292969, - accuracy: 0.010945053100585937) - XCTAssertEqual( - a.sum().item(Float.self), 140.0966796875, - accuracy: 2.80193359375) - let result = GLU()(a) - XCTAssertEqual(result.shape, [2, 8, 8]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.33327674865722656, - accuracy: 0.006665534973144531) - XCTAssertEqual( - result.sum().item(Float.self), 42.659423828125, - accuracy: 0.8531884765625) - } - - func testSigmoid1() { - MLXRandom.seed(589) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 8, 16]) - XCTAssertEqual(a.shape, [2, 8, 16]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.5296978950500488, - accuracy: 0.010593957901000976) - XCTAssertEqual( - a.sum().item(Float.self), 135.6026611328125, - accuracy: 2.71205322265625) - let result = Sigmoid()(a) - XCTAssertEqual(result.shape, [2, 8, 16]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.6270139813423157, - accuracy: 0.012540279626846314) - XCTAssertEqual( - result.sum().item(Float.self), 160.5155792236328, - accuracy: 3.2103115844726564) - } - - func testMish() { - MLXRandom.seed(122) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 8, 16]) - XCTAssertEqual(a.shape, [2, 8, 16]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.5017197728157043, - accuracy: 0.010034395456314087) - XCTAssertEqual( - a.sum().item(Float.self), 128.4402618408203, - accuracy: 2.5688052368164063) - let result = Mish()(a) - XCTAssertEqual(result.shape, [2, 8, 16]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.39537572860717773, - accuracy: 0.007907514572143556) - XCTAssertEqual( - result.sum().item(Float.self), 101.2161865234375, - accuracy: 2.0243237304687502) - } - - func testReLU() { - MLXRandom.seed(400) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 8, 16]) - XCTAssertEqual(a.shape, [2, 8, 16]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.47832274436950684, - accuracy: 0.009566454887390137) - XCTAssertEqual( - a.sum().item(Float.self), 122.45062255859375, - accuracy: 2.449012451171875) - let result = ReLU()(a) - XCTAssertEqual(result.shape, [2, 8, 16]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.47832274436950684, - accuracy: 0.009566454887390137) - XCTAssertEqual( - result.sum().item(Float.self), 122.45062255859375, - accuracy: 2.449012451171875) - } - - func testLeakyReLU() { - MLXRandom.seed(93) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 8, 16]) - XCTAssertEqual(a.shape, [2, 8, 16]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.4999306797981262, - accuracy: 0.009998613595962524) - XCTAssertEqual( - a.sum().item(Float.self), 127.98225402832031, - accuracy: 2.559645080566406) - let result = LeakyReLU()(a) - XCTAssertEqual(result.shape, [2, 8, 16]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.4999306797981262, - accuracy: 0.009998613595962524) - XCTAssertEqual( - result.sum().item(Float.self), 127.98225402832031, - accuracy: 2.559645080566406) - } - - func testReLU6() { - MLXRandom.seed(379) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 8, 16]) - XCTAssertEqual(a.shape, [2, 8, 16]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.49325865507125854, - accuracy: 0.009865173101425172) - XCTAssertEqual( - a.sum().item(Float.self), 126.27421569824219, - accuracy: 2.525484313964844) - let result = ReLU6()(a) - XCTAssertEqual(result.shape, [2, 8, 16]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.49325865507125854, - accuracy: 0.009865173101425172) - XCTAssertEqual( - result.sum().item(Float.self), 126.27421569824219, - accuracy: 2.525484313964844) - } - - func testSoftmax() { - MLXRandom.seed(853) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 8, 16]) - XCTAssertEqual(a.shape, [2, 8, 16]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.5143963098526001, - accuracy: 0.010287926197052003) - XCTAssertEqual( - a.sum().item(Float.self), 131.68545532226562, - accuracy: 2.6337091064453126) - let result = Softmax()(a) - XCTAssertEqual(result.shape, [2, 8, 16]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.0624999962747097, - accuracy: 0.001249999925494194) - XCTAssertEqual( - result.sum().item(Float.self), 15.999999046325684, - accuracy: 0.31999998092651366) - } - - func testSoftplus() { - MLXRandom.seed(118) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 8, 16]) - XCTAssertEqual(a.shape, [2, 8, 16]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.49898141622543335, - accuracy: 0.009979628324508667) - XCTAssertEqual( - a.sum().item(Float.self), 127.73924255371094, - accuracy: 2.554784851074219) - let result = SoftPlus()(a) - XCTAssertEqual(result.shape, [2, 8, 16]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.9828577637672424, - accuracy: 0.01965715527534485) - XCTAssertEqual( - result.sum().item(Float.self), 251.61158752441406, - accuracy: 5.032231750488282) - } - - func testSoftsign() { - MLXRandom.seed(37) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 8, 16]) - XCTAssertEqual(a.shape, [2, 8, 16]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.5065512657165527, - accuracy: 0.010131025314331054) - XCTAssertEqual( - a.sum().item(Float.self), 129.6771240234375, - accuracy: 2.59354248046875) - let result = SoftSign()(a) - XCTAssertEqual(result.shape, [2, 8, 16]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.314089834690094, - accuracy: 0.00628179669380188) - XCTAssertEqual( - result.sum().item(Float.self), 80.40699768066406, - accuracy: 1.6081399536132812) - } - - func testCELU() { - MLXRandom.seed(620) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 8, 16]) - XCTAssertEqual(a.shape, [2, 8, 16]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.4667481780052185, - accuracy: 0.00933496356010437) - XCTAssertEqual( - a.sum().item(Float.self), 119.48753356933594, - accuracy: 2.389750671386719) - let result = CELU()(a) - XCTAssertEqual(result.shape, [2, 8, 16]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.4667481780052185, - accuracy: 0.00933496356010437) - XCTAssertEqual( - result.sum().item(Float.self), 119.48753356933594, - accuracy: 2.389750671386719) - } - - func testSiLU() { - MLXRandom.seed(22) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 8, 16]) - XCTAssertEqual(a.shape, [2, 8, 16]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.5029705762863159, - accuracy: 0.010059411525726319) - XCTAssertEqual( - a.sum().item(Float.self), 128.76046752929688, - accuracy: 2.5752093505859377) - let result = SiLU()(a) - XCTAssertEqual(result.shape, [2, 8, 16]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.3319709300994873, - accuracy: 0.0066394186019897465) - XCTAssertEqual( - result.sum().item(Float.self), 84.98455810546875, - accuracy: 1.699691162109375) - } - - func testLogSoftmax() { - MLXRandom.seed(199) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 8, 16]) - XCTAssertEqual(a.shape, [2, 8, 16]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.527843713760376, - accuracy: 0.010556874275207519) - XCTAssertEqual( - a.sum().item(Float.self), 135.12799072265625, - accuracy: 2.702559814453125) - let result = LogSoftMax()(a) - XCTAssertEqual(result.shape, [2, 8, 16]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -2.8109545707702637, - accuracy: -0.05621909141540527) - XCTAssertEqual( - result.sum().item(Float.self), -719.6043701171875, - accuracy: -14.39208740234375) - } - - func testLogSigmoid() { - MLXRandom.seed(984) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 8, 16]) - XCTAssertEqual(a.shape, [2, 8, 16]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.5109776854515076, - accuracy: 0.010219553709030152) - XCTAssertEqual( - a.sum().item(Float.self), 130.81028747558594, - accuracy: 2.616205749511719) - let result = LogSigmoid()(a) - XCTAssertEqual(result.shape, [2, 8, 16]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.4795985519886017, - accuracy: -0.009591971039772034) - XCTAssertEqual( - result.sum().item(Float.self), -122.77722930908203, - accuracy: -2.4555445861816407) - } - - func testPReLU() { - MLXRandom.seed(993) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 8, 16]) - XCTAssertEqual(a.shape, [2, 8, 16]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.49665144085884094, - accuracy: 0.00993302881717682) - XCTAssertEqual( - a.sum().item(Float.self), 127.14276885986328, - accuracy: 2.5428553771972657) - let result = PReLU()(a) - XCTAssertEqual(result.shape, [2, 8, 16]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.49665144085884094, - accuracy: 0.00993302881717682) - XCTAssertEqual( - result.sum().item(Float.self), 127.14276885986328, - accuracy: 2.5428553771972657) - } - - func testGELU() { - MLXRandom.seed(189) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 8, 16]) - XCTAssertEqual(a.shape, [2, 8, 16]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.49295032024383545, - accuracy: 0.009859006404876709) - XCTAssertEqual( - a.sum().item(Float.self), 126.19528198242188, - accuracy: 2.5239056396484374) - let result = GELU()(a) - XCTAssertEqual(result.shape, [2, 8, 16]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.3656383752822876, - accuracy: 0.007312767505645752) - XCTAssertEqual( - result.sum().item(Float.self), 93.60342407226562, - accuracy: 1.8720684814453126) - } - - func testTanh1() { - MLXRandom.seed(735) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 8, 16]) - XCTAssertEqual(a.shape, [2, 8, 16]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.47412270307540894, - accuracy: 0.009482454061508178) - XCTAssertEqual( - a.sum().item(Float.self), 121.37541198730469, - accuracy: 2.4275082397460936) - let result = Tanh()(a) - XCTAssertEqual(result.shape, [2, 8, 16]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.4130796790122986, - accuracy: 0.008261593580245972) - XCTAssertEqual( - result.sum().item(Float.self), 105.74839782714844, - accuracy: 2.114967956542969) - } - - func testHardswish() { - MLXRandom.seed(126) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 8, 16]) - XCTAssertEqual(a.shape, [2, 8, 16]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.4918924570083618, - accuracy: 0.009837849140167236) - XCTAssertEqual( - a.sum().item(Float.self), 125.92446899414062, - accuracy: 2.5184893798828125) - let result = HardSwish()(a) - XCTAssertEqual(result.shape, [2, 8, 16]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.2996022403240204, - accuracy: 0.005992044806480408) - XCTAssertEqual( - result.sum().item(Float.self), 76.69817352294922, - accuracy: 1.5339634704589844) - } - - func testStep() { - MLXRandom.seed(490) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 8, 16]) - XCTAssertEqual(a.shape, [2, 8, 16]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.4793606400489807, - accuracy: 0.009587212800979614) - XCTAssertEqual( - a.sum().item(Float.self), 122.71632385253906, - accuracy: 2.454326477050781) - let result = Step()(a) - XCTAssertEqual(result.shape, [2, 8, 16]) - XCTAssertEqual(result.dtype, .int32) - XCTAssertEqual( - result.mean().item(Float.self), 1.0, - accuracy: 0.02) - XCTAssertEqual( - result.sum().item(Float.self), 256, - accuracy: 5.12) - } - - func testSELU() { - MLXRandom.seed(215) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 8, 16]) - XCTAssertEqual(a.shape, [2, 8, 16]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.4930267930030823, - accuracy: 0.009860535860061645) - XCTAssertEqual( - a.sum().item(Float.self), 126.21485900878906, - accuracy: 2.524297180175781) - let result = SELU()(a) - XCTAssertEqual(result.shape, [2, 8, 16]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.5180231928825378, - accuracy: 0.010360463857650756) - XCTAssertEqual( - result.sum().item(Float.self), 132.6139373779297, - accuracy: 2.6522787475585936) - } - - func testLinear() { - MLXRandom.seed(744) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 8, 16]) - XCTAssertEqual(a.shape, [2, 8, 16]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.5086885690689087, - accuracy: 0.010173771381378174) - XCTAssertEqual( - a.sum().item(Float.self), 130.22427368164062, - accuracy: 2.6044854736328125) - let result = Linear(16, 5)(a) - XCTAssertEqual(result.shape, [2, 8, 5]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.10419309139251709, - accuracy: 0.002083861827850342) - XCTAssertEqual( - result.sum().item(Float.self), 8.335447311401367, - accuracy: 0.16670894622802734) - } - - func testConv1d1() { - MLXRandom.seed(819) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 8, 16]) - XCTAssertEqual(a.shape, [2, 8, 16]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.512987494468689, - accuracy: 0.01025974988937378) - XCTAssertEqual( - a.sum().item(Float.self), 131.32479858398438, - accuracy: 2.6264959716796876) - let result = Conv1d(inputChannels: 16, outputChannels: 2, kernelSize: 8)(a) - XCTAssertEqual(result.shape, [2, 1, 2]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.2648651897907257, - accuracy: 0.005297303795814514) - XCTAssertEqual( - result.sum().item(Float.self), 1.0594607591629028, - accuracy: 0.021189215183258055) - } - - func testConv2d1() { - MLXRandom.seed(62) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 8, 8, 4]) - XCTAssertEqual(a.shape, [2, 8, 8, 4]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.5225042700767517, - accuracy: 0.010450085401535034) - XCTAssertEqual( - a.sum().item(Float.self), 267.5221862792969, - accuracy: 5.350443725585937) - let result = Conv2d(inputChannels: 4, outputChannels: 2, kernelSize: 8)(a) - XCTAssertEqual(result.shape, [2, 1, 1, 2]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.27932149171829224, - accuracy: -0.005586429834365845) - XCTAssertEqual( - result.sum().item(Float.self), -1.117285966873169, - accuracy: -0.02234571933746338) - } - - func testDropout() { - MLXRandom.seed(959) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 8, 16]) - XCTAssertEqual(a.shape, [2, 8, 16]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.5114291906356812, - accuracy: 0.010228583812713623) - XCTAssertEqual( - a.sum().item(Float.self), 130.92587280273438, - accuracy: 2.6185174560546876) - let result = Dropout()(a) - XCTAssertEqual(result.shape, [2, 8, 16]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.47791361808776855, - accuracy: 0.009558272361755372) - XCTAssertEqual( - result.sum().item(Float.self), 122.34588623046875, - accuracy: 2.446917724609375) - } - - func testDropout2d() { - MLXRandom.seed(695) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 8, 16]) - XCTAssertEqual(a.shape, [2, 8, 16]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.4578399062156677, - accuracy: 0.009156798124313355) - XCTAssertEqual( - a.sum().item(Float.self), 117.20701599121094, - accuracy: 2.344140319824219) - let result = Dropout2d()(a) - XCTAssertEqual(result.shape, [2, 8, 16]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.36828434467315674, - accuracy: 0.007365686893463135) - XCTAssertEqual( - result.sum().item(Float.self), 94.28079223632812, - accuracy: 1.8856158447265625) - } - - func testDropout3d() { - MLXRandom.seed(23) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 8, 8, 4]) - XCTAssertEqual(a.shape, [2, 8, 8, 4]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.5006061792373657, - accuracy: 0.010012123584747314) - XCTAssertEqual( - a.sum().item(Float.self), 256.31036376953125, - accuracy: 5.126207275390625) - let result = Dropout3d()(a) - XCTAssertEqual(result.shape, [2, 8, 8, 4]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.23728415369987488, - accuracy: 0.004745683073997498) - XCTAssertEqual( - result.sum().item(Float.self), 121.48948669433594, - accuracy: 2.429789733886719) - } - - func testEmbedding() { - MLXRandom.seed(557) - let a = MLXRandom.randInt(low: 0, high: 10, [2, 8, 8, 4]) - XCTAssertEqual(a.shape, [2, 8, 8, 4]) - XCTAssertEqual(a.dtype, .int32) - XCTAssertEqual( - a.mean().item(Float.self), 4.60546875, - accuracy: 0.09210937500000001) - XCTAssertEqual( - a.sum().item(Float.self), 2358, - accuracy: 47.160000000000004) - let result = Embedding(embeddingCount: 10, dimensions: 8)(a) - XCTAssertEqual(result.shape, [2, 8, 8, 4, 8]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.0011973462533205748, - accuracy: -2.3946925066411497e-05) - XCTAssertEqual( - result.sum().item(Float.self), -4.904330253601074, - accuracy: -0.0980866050720215) - } - - func testInstanceNorm() { - MLXRandom.seed(435) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 8, 16]) - XCTAssertEqual(a.shape, [2, 8, 16]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.5000646114349365, - accuracy: 0.01000129222869873) - XCTAssertEqual( - a.sum().item(Float.self), 128.01654052734375, - accuracy: 2.560330810546875) - let result = InstanceNorm(dimensions: 8)(a)[0, 0] - XCTAssertEqual(result.shape, [16]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.10645411163568497, - accuracy: 0.0021290822327136995) - XCTAssertEqual( - result.sum().item(Float.self), 1.7032657861709595, - accuracy: 0.03406531572341919) - } - - func testLayerNorm() { - MLXRandom.seed(635) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 8, 16]) - XCTAssertEqual(a.shape, [2, 8, 16]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.4926903247833252, - accuracy: 0.009853806495666504) - XCTAssertEqual( - a.sum().item(Float.self), 126.12872314453125, - accuracy: 2.522574462890625) - let result = LayerNorm(dimensions: 16)(a)[.ellipsis, 0] - XCTAssertEqual(result.shape, [2, 8]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.2909903824329376, - accuracy: 0.005819807648658752) - XCTAssertEqual( - result.sum().item(Float.self), 4.655846118927002, - accuracy: 0.09311692237854004) - } - - func testRMSNorm() { - MLXRandom.seed(103) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 8, 16]) - XCTAssertEqual(a.shape, [2, 8, 16]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.5054763555526733, - accuracy: 0.010109527111053467) - XCTAssertEqual( - a.sum().item(Float.self), 129.40194702148438, - accuracy: 2.5880389404296875) - let result = RMSNorm(dimensions: 16)(a) - XCTAssertEqual(result.shape, [2, 8, 16]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.8729387521743774, - accuracy: 0.01745877504348755) - XCTAssertEqual( - result.sum().item(Float.self), 223.47232055664062, - accuracy: 4.469446411132813) - } - - func testGroupNorm() { - MLXRandom.seed(855) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 8, 16]) - XCTAssertEqual(a.shape, [2, 8, 16]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.48666587471961975, - accuracy: 0.009733317494392395) - XCTAssertEqual( - a.sum().item(Float.self), 124.58646392822266, - accuracy: 2.491729278564453) - let result = GroupNorm(groupCount: 4, dimensions: 16)(a)[0, 0] - XCTAssertEqual(result.shape, [16]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), -0.054606519639492035, - accuracy: -0.0010921303927898408) - XCTAssertEqual( - result.sum().item(Float.self), -0.8737043142318726, - accuracy: -0.017474086284637452) - } - - func testBatchNorm() { - MLXRandom.seed(266) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 8, 16]) - XCTAssertEqual(a.shape, [2, 8, 16]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.5058146715164185, - accuracy: 0.010116293430328369) - XCTAssertEqual( - a.sum().item(Float.self), 129.48855590820312, - accuracy: 2.5897711181640624) - let result = BatchNorm(featureCount: 16)(a)[0, 0] - XCTAssertEqual(result.shape, [16]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.4397852420806885, - accuracy: 0.00879570484161377) - XCTAssertEqual( - result.sum().item(Float.self), 7.036563873291016, - accuracy: 0.14073127746582031) - } - - func testRoPE() { - MLXRandom.seed(71) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 8, 16]) - XCTAssertEqual(a.shape, [2, 8, 16]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.5082664489746094, - accuracy: 0.010165328979492188) - XCTAssertEqual( - a.sum().item(Float.self), 130.1162109375, - accuracy: 2.60232421875) - let result = RoPE(dimensions: 8)(a) - XCTAssertEqual(result.shape, [2, 8, 16]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.4562537670135498, - accuracy: 0.009125075340270997) - XCTAssertEqual( - result.sum().item(Float.self), 116.80096435546875, - accuracy: 2.3360192871093752) - } - - func testRoPEArrayOffset() { - MLXRandom.seed(42) - let batch = MLXRandom.uniform(0.0 ..< 1.0, [3, 8, 16]) - XCTAssertEqual(batch.shape, [3, 8, 16]) - - let offsets = MLXArray([50, 20, 0]) - - // Test MLXFast.RoPE with array offset - let result = MLXFast.RoPE( - batch, dimensions: 8, traditional: false, - base: 10000, scale: 1.0, offset: offsets) - XCTAssertEqual(result.shape, [3, 8, 16]) - - // Verify against individual scalar offset calls - for i in 0 ..< 3 { - let single = batch[i].expandedDimensions(axis: 0) - let offsetValue = [50, 20, 0][i] - let expected = MLXFast.RoPE( - single, dimensions: 8, traditional: false, - base: 10000, scale: 1.0, offset: offsetValue) - XCTAssert(allClose(result[i], expected[0]).all().item()) - } - } - - func testRoPEArrayOffsetModule() { - MLXRandom.seed(123) - let batch = MLXRandom.uniform(0.0 ..< 1.0, [3, 8, 16]) - let offsets = MLXArray([10, 5, 0]) - - let rope = RoPE(dimensions: 8) - - // Test MLXNN RoPE module with array offset - let result = rope(batch, offset: offsets) - XCTAssertEqual(result.shape, [3, 8, 16]) - - // Verify shape and dtype preserved - XCTAssertEqual(result.dtype, batch.dtype) - } - - func testRoPEScalarArrayOffset() { - MLXRandom.seed(99) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 8, 16]) - - // Scalar array offset should work the same as int offset - let scalarOffset = MLXArray(5) - let resultArray = MLXFast.RoPE( - a, dimensions: 8, traditional: false, - base: 10000, scale: 1.0, offset: scalarOffset) - - let resultInt = MLXFast.RoPE( - a, dimensions: 8, traditional: false, - base: 10000, scale: 1.0, offset: 5) - - XCTAssert(allClose(resultArray, resultInt).all().item()) - } - - func testSinusoidalPositionalEncoding() { - MLXRandom.seed(226) - let a = MLXRandom.uniform(0.0 ..< 1.0, [2, 8, 16]) - XCTAssertEqual(a.shape, [2, 8, 16]) - XCTAssertEqual(a.dtype, .float32) - XCTAssertEqual( - a.mean().item(Float.self), 0.5026599168777466, - accuracy: 0.010053198337554931) - XCTAssertEqual( - a.sum().item(Float.self), 128.68093872070312, - accuracy: 2.5736187744140624) - let result = SinusoidalPositionalEncoding(dimensions: 8)(a) - XCTAssertEqual(result.shape, [2, 8, 16, 8]) - XCTAssertEqual(result.dtype, .float32) - XCTAssertEqual( - result.mean().item(Float.self), 0.2705308198928833, - accuracy: 0.005410616397857666) - XCTAssertEqual( - result.sum().item(Float.self), 554.047119140625, - accuracy: 11.0809423828125) - } - -} diff --git a/Tests/MLXTests/MLXArray+InitTests.swift b/Tests/MLXTests/MLXArray+InitTests.swift index 8c1017be4..35e344718 100644 --- a/Tests/MLXTests/MLXArray+InitTests.swift +++ b/Tests/MLXTests/MLXArray+InitTests.swift @@ -438,6 +438,46 @@ class MLXArrayInitTests: XCTestCase { + "buffer permanently pins whatever the closure held") } + // MARK: - linspace + + func testLinspaceIntegerBoundsAreFloat() { + // matches python: mx.linspace(0, 1, 3) -> float32 [0, 0.5, 1]. + // an integer result would truncate to [0, 0, 1] + let a = MLXArray.linspace(0, 1, count: 3) + XCTAssertEqual(a.dtype, .float32) + assertEqual(a, MLXArray(converting: [0.0, 0.5, 1.0]), atol: 1e-6) + } + + func testLinspaceIntegerDType() { + // an integer (truncating) result is available, but has to be asked for + let a = MLXArray.linspace(0, 10, count: 6, dtype: .int32) + XCTAssertEqual(a.dtype, .int32) + XCTAssertEqual(a.asArray(Int32.self), [0, 2, 4, 6, 8, 10]) + } + + func testLinspaceDoubleBoundsAreFloat32() { + // Double.dtype is float64, but like MLXArray(1.0) we do not promote to + // float64 -- it is not available on the GPU + let a = MLXArray.linspace(0.0, 1.0, count: 5) + XCTAssertEqual(a.dtype, .float32) + assertEqual(a, MLXArray(converting: [0.0, 0.25, 0.5, 0.75, 1.0]), atol: 1e-6) + } + + func testLinspaceFloatBoundsKeepTheirDType() { + XCTAssertEqual(MLXArray.linspace(Float(0), Float(1), count: 5).dtype, .float32) + #if !arch(x86_64) + XCTAssertEqual(MLXArray.linspace(Float16(0), Float16(1), count: 5).dtype, .float16) + #endif + } + + func testLinspaceEndpoint() { + let inclusive = MLX.linspace(0, 1, count: 5) + assertEqual(inclusive, MLXArray(converting: [0, 0.25, 0.5, 0.75, 1.0]), atol: 1e-6) + + let halfOpen = MLX.linspace(0, 1, count: 5, endpoint: false) + assertEqual(halfOpen, MLXArray(converting: [0, 0.2, 0.4, 0.6, 0.8]), atol: 1e-6) + } + #if canImport(IOSurface) func testIOSurface() { let height = 100 diff --git a/Tests/MLXTests/OpsTests.swift b/Tests/MLXTests/OpsTests.swift index b59970c23..41bde1dcb 100644 --- a/Tests/MLXTests/OpsTests.swift +++ b/Tests/MLXTests/OpsTests.swift @@ -53,6 +53,85 @@ class OpsTests: XCTestCase { assertEqual(c, expected) } + func testTensordotDefaultAxes() { + // axes defaults to 2 (as in numpy and python mlx): sum over the last two + // dimensions of a and the first two of b + let a = MLXArray(0 ..< 24, [2, 3, 4]).asType(.float32) + let b = MLXArray(0 ..< 60, [3, 4, 5]).asType(.float32) + + assertEqual(tensordot(a, b), tensordot(a, b, axes: 2)) + XCTAssertEqual(tensordot(a, b).shape, [2, 5]) + } + + func testNanToNumDefaults() { + // by default infinities become the largest finite value for the dtype + // (matching python mlx) and NaN becomes 0 + let a = MLXArray( + [1.5, Float.nan, Float.infinity, -Float.infinity] as [Float]) + + let result = nanToNum(a) + XCTAssertEqual(result.dtype, .float32) + assertEqual( + result, + MLXArray( + [ + 1.5, 0, Float.greatestFiniteMagnitude, -Float.greatestFiniteMagnitude, + ] as [Float])) + + // and they can be replaced explicitly + assertEqual( + nanToNum(a, nan: -1, posInf: 100, negInf: -100), + MLXArray([1.5, -1, 100, -100] as [Float])) + } + + func testNanToNumDefaultsFloat16() { + // the replacement follows the dtype + let a = MLXArray([Float.infinity, -Float.infinity] as [Float]).asType(.float16) + let limit = Float(DType.float16.finfo!.max) + + let result = nanToNum(a) + XCTAssertEqual(result.dtype, .float16) + XCTAssertEqual(result[0].item(Float.self), limit) + XCTAssertEqual(result[1].item(Float.self), -limit) + } + + func testConvolveModeShapes() { + // matches numpy/python mlx: full is M + K - 1, same is M, valid is M - K + 1 + let a = MLXArray(0 ..< 20).asType(.float32) + + for kernelSize in [3, 4, 5, 6] { + let v = MLXArray(1 ..< (kernelSize + 1)).asType(.float32) + + XCTAssertEqual( + convolve(a, v, mode: .full).shape, [a.size + kernelSize - 1], + "full, kernel \(kernelSize)") + XCTAssertEqual( + convolve(a, v, mode: .same).shape, [a.size], "same, kernel \(kernelSize)") + XCTAssertEqual( + convolve(a, v, mode: .valid).shape, [a.size - kernelSize + 1], + "valid, kernel \(kernelSize)") + } + } + + func testConvolveSameAndValidAreWindowsOfFull() { + // `same` is the centered `a.size` window of the full convolution and + // `valid` is the window with no zero padding at all -- true for both odd + // and even sized weights (even sizes need asymmetric padding) + let a = MLXArray(0 ..< 20).asType(.float32) + + for kernelSize in [3, 4, 5, 6] { + let v = MLXArray(1 ..< (kernelSize + 1)).asType(.float32) + let full = convolve(a, v, mode: .full) + + let sameStart = (kernelSize - 1) / 2 + assertEqual( + convolve(a, v, mode: .same), full[sameStart ..< (sameStart + a.size)]) + + assertEqual( + convolve(a, v, mode: .valid), full[(kernelSize - 1) ..< a.size]) + } + } + func testConvertScalarInt() { let a = MLXArray(0 ..< 10) let b = a .< (a + 1) diff --git a/Tests/MLXTests/OptimizerTests.swift b/Tests/MLXTests/OptimizerTests.swift index 95e04af7d..5d1e3727a 100644 --- a/Tests/MLXTests/OptimizerTests.swift +++ b/Tests/MLXTests/OptimizerTests.swift @@ -13,17 +13,24 @@ class OptimizerTests: XCTestCase { } class ShapeModule: Module { + // ranks 1 through 3: some optimizers (Adafactor, Muon) treat them differently let first = [MLXArray.zeros([10]), MLXArray.zeros([1])] let second = MLXArray.zeros([1]) + let matrix = MLXArray.zeros([3, 5]) + let tensor = MLXArray.zeros([2, 3, 4]) } - func checkShape(optimizer: OptimizerBase) { + func checkShape(optimizer: OptimizerBase, steps: Int = 2) { let model = ShapeModule() let params = model.parameters() let grads = params.mapValues { MLXArray.ones(like: $0) } - let optimizer = SGD(learningRate: 0.1) - let update = optimizer.apply(gradients: grads, modelParameters: model.parameters()) + // note: more than one step, and using the optimizer that was passed in -- + // this used to build its own SGD, so every caller was really testing SGD + var update = params + for _ in 0 ..< steps { + update = optimizer.apply(gradients: grads, modelParameters: update) + } eval(update) let shapesEqual = params.mapValues(update) { (e1, e2) -> Bool in @@ -212,6 +219,167 @@ class OptimizerTests: XCTestCase { checkTrain(optimizer: Adafactor(learningRate: 0.1), compile: true) } + // MARK: - Muon + + func testMuonUpdateIsConditioned() { + // the Newton-Schulz iteration used by Muon is a *crude* polar approximation: + // it pulls the singular values of the update into a band around 1 (it does + // not converge to an exactly orthogonal matrix). With no momentum and + // nesterov off the update direction is the orthogonalized gradient, so the + // singular values of the update should be far better conditioned than the + // gradient's. + let optimizer = Muon( + learningRate: 1.0, momentum: 0.0, weightDecay: 0.0, nesterov: false, nsSteps: 5) + + // full rank and deterministic: `arange` reshaped is rank 2, which the + // iteration cannot orthogonalize + let values = (0 ..< 24).map { Float(sin(Double($0) * 1.7) + 0.1 * Double($0 % 5)) } + let parameter = MLXArray(values, [6, 4]) + let parameters = ModuleParameters.unflattened([("w", parameter)]) + + let updated = optimizer.apply(gradients: parameters, modelParameters: parameters) + + // parameter - learningRate * scale * update, with scale = sqrt(max(1, 6/4)) + let scale = Float((6.0 / 4.0).squareRoot()) + let update = (parameter - updated[unwrapping: "w"]!) / scale + + // singular values via the gram matrix, which avoids picking between the svd + // overloads: eigvalsh returns them ascending + func singularValues(_ array: MLXArray) -> MLXArray { + MLX.sqrt( + MLX.maximum(0, MLX.eigvalsh(matmul(array.T, array), stream: .cpu))) + } + + let inputSingular = singularValues(parameter) + let updateSingular = singularValues(update) + + let inputCondition = + inputSingular.max().item(Float.self) / inputSingular.min().item(Float.self) + let updateCondition = + updateSingular.max().item(Float.self) / updateSingular.min().item(Float.self) + + XCTAssertGreaterThan(inputCondition, 10) + XCTAssertLessThan(updateCondition, 2) + // and the values are pulled toward 1 rather than being rescaled arbitrarily + XCTAssertGreaterThan(updateSingular.min().item(Float.self), 0.5) + XCTAssertLessThan(updateSingular.max().item(Float.self), 1.5) + } + + func testMuonLeavesLowRankParametersAlone() { + // rank 0/1 parameters skip the orthogonalization: the update is the plain + // momentum direction, so a single step moves by learningRate * gradient + let optimizer = Muon( + learningRate: 0.1, momentum: 0.0, weightDecay: 0.0, nesterov: false, nsSteps: 5) + + let parameter = MLXArray(converting: [1.0, 2.0, 3.0]) + let gradients = ModuleParameters.unflattened([("b", MLXArray.ones([3]))]) + let updated = optimizer.apply( + gradients: gradients, + modelParameters: ModuleParameters.unflattened([("b", parameter)])) + + assertEqual( + updated[unwrapping: "b"]!, MLXArray(converting: [0.9, 1.9, 2.9]), atol: 1e-6) + } + + // MARK: - defaults + // + // these have to match python: `tools/audit_defaults.py` compares them + // automatically, and this is the regression test for the one that was wrong + + func testLionDefaults() { + let lion = Lion(learningRate: 0.1) + XCTAssertEqual(lion.betas.0, 0.9) + // python's Lion uses 0.99 here -- Adam is the one with 0.999 + XCTAssertEqual(lion.betas.1, 0.99) + XCTAssertEqual(lion.weightDecay, 0) + } + + func testAdamAndAdamWDefaults() { + XCTAssertEqual(Adam(learningRate: 0.1).betas.1, 0.999) + XCTAssertEqual(AdamW(learningRate: 0.1).betas.1, 0.999) + XCTAssertEqual(AdamW(learningRate: 0.1).weightDecay, 0.01) + } + + func testMuonDefaults() { + let muon = Muon(learningRate: 0.1) + XCTAssertEqual(muon.momentum, 0.95) + XCTAssertEqual(muon.weightDecay, 0.01) + XCTAssertTrue(muon.nesterov) + XCTAssertEqual(muon.nsSteps, 5) + } + + // MARK: - Adafactor + // + // Adafactor keeps *factored* state for parameters of rank 2 and above and + // unfactored state below that, so its state and its update path have to agree + // about the rank -- rank 3+ was broken (an outer product via matmul, which only + // works for rank 2) and the mismatch then crashed on a nil state. + + func testAdafactorRanks() throws { + let optimizer = Adafactor(learningRate: 0.1, relativeStep: false) + + var parameters = ModuleParameters.unflattened([ + ("vector", MLXArray.ones([7])), + ("matrix", MLXArray.ones([3, 5])), + ("tensor", MLXArray.ones([2, 3, 4])), + ("big", MLXArray.ones([2, 3, 4, 5])), + ]) + + for _ in 0 ..< 3 { + let gradients = parameters.mapValues { 0.5 * $0 } + parameters = optimizer.apply(gradients: gradients, modelParameters: parameters) + eval(parameters) + } + + for (key, shape) in [ + ("vector", [7]), ("matrix", [3, 5]), ("tensor", [2, 3, 4]), ("big", [2, 3, 4, 5]), + ] { + let value = parameters[unwrapping: key]! + XCTAssertEqual(value.shape, shape, key) + XCTAssertTrue(isFinite(value).all().item(Bool.self), "\(key) is not finite") + // the update should have moved the parameter off its starting value + XCTAssertFalse( + value.allClose(MLXArray.ones(like: value), rtol: 1e-6).item(Bool.self), + "\(key) did not change") + } + } + + func testAdafactorFactoredUpdateIsBatched() { + // the row/column factors combine per leading index: slice `i` of a rank 3 + // update must equal the rank 2 update of slice `i` of the factors + let optimizer = Adafactor(learningRate: 0.1, relativeStep: false) + + let row = MLXArray(converting: [1.0, 2.0, 3.0, 4.0, 5.0, 6.0], [2, 3]) + let column = MLXArray(converting: [0.5, 1.0, 1.5, 2.0, 2.5, 3.0, 3.5, 4.0], [2, 4]) + + let combined = optimizer.approvateExpMovingAverage( + expAvgSqRow: row, expAvgSqCol: column) + XCTAssertEqual(combined.shape, [2, 3, 4]) + + for i in 0 ..< 2 { + let slice = optimizer.approvateExpMovingAverage( + expAvgSqRow: row[i], expAvgSqCol: column[i]) + assertEqual(combined[i], slice, atol: 1e-6) + } + } + + func testAdafactorRankChangeDoesNotCrash() { + // the state is created for the rank seen on the first step; a later call with + // a different rank for the same key used to hit a nil state + let optimizer = Adafactor(learningRate: 0.1, relativeStep: false) + + var matrix = ModuleParameters.unflattened([("w", MLXArray.ones([3, 4]))]) + matrix = optimizer.apply( + gradients: matrix.mapValues { 0.5 * $0 }, modelParameters: matrix) + + var vector = ModuleParameters.unflattened([("w", MLXArray.ones([5]))]) + vector = optimizer.apply( + gradients: vector.mapValues { 0.5 * $0 }, modelParameters: vector) + eval(vector) + + XCTAssertEqual(vector[unwrapping: "w"]!.shape, [5]) + } + class TwoParameterModel: Module { let weight = MLXArray.zeros([3]) let bias = MLXArray.zeros([3]) diff --git a/Tests/MLXTests/PoolingTests.swift b/Tests/MLXTests/PoolingTests.swift index 90534fca9..2363b21ef 100644 --- a/Tests/MLXTests/PoolingTests.swift +++ b/Tests/MLXTests/PoolingTests.swift @@ -67,4 +67,87 @@ class MLXNNPoolingTests: XCTestCase { let output = pool.callAsFunction(input) assertEqual(output, MLXArray(converting: [2.5, 4.5, 10.5, 12.5], [1, 2, 2, 1])) } + + // MARK: - padding + // + // padding shifts every spatial dimension and must leave the batch and channel + // dimensions alone: the output is + // floor((size + 2 * padding - kernel) / stride) + 1 + + func testMaxPooling1dPadding() { + let input = MLXArray(0 ..< 4, [1, 4, 1]) + let pool = MaxPool1d(kernelSize: 2, stride: 2, padding: 1) + let output = pool(input) + + // padded with -inf: [-inf, 0, 1, 2, 3, -inf] + XCTAssertEqual(output.shape, [1, 3, 1]) + assertEqual(output, MLXArray([0, 2, 3], [1, 3, 1])) + } + + func testAvgPooling1dPadding() { + let input = MLXArray(0 ..< 4, [1, 4, 1]).asType(.float32) + let pool = AvgPool1d(kernelSize: 2, stride: 2, padding: 1) + let output = pool(input) + + // padded with 0, which is included in the average: [0, 0, 1, 2, 3, 0] + XCTAssertEqual(output.shape, [1, 3, 1]) + assertEqual(output, MLXArray(converting: [0, 1.5, 1.5], [1, 3, 1])) + } + + func testMaxPooling2dPadding() { + let input = MLXArray(0 ..< 16, [1, 4, 4, 1]) + let pool = MaxPool2d(kernelSize: 3, stride: 2, padding: 1) + let output = pool(input) + + XCTAssertEqual(output.shape, [1, 2, 2, 1]) + assertEqual(output, MLXArray([5, 7, 13, 15], [1, 2, 2, 1])) + } + + func testAvgPooling2dPadding() { + let input = MLXArray(0 ..< 4, [1, 2, 2, 1]).asType(.float32) + let pool = AvgPool2d(kernelSize: 2, stride: 2, padding: 1) + let output = pool(input) + + XCTAssertEqual(output.shape, [1, 2, 2, 1]) + assertEqual(output, MLXArray(converting: [0, 0.25, 0.5, 0.75], [1, 2, 2, 1])) + } + + func testPoolingPaddingKeepsBatchAndChannels() { + // the padding must not touch the batch or channel dimensions: this shape + // regressed to [2, 3, 4, 6] when the pad widths were built per pair + // instead of per dimension + let input = MLXArray(0 ..< (2 * 8 * 8 * 4), [2, 8, 8, 4]).asType(.float32) + + XCTAssertEqual( + MaxPool2d(kernelSize: 3, stride: 2, padding: 1)(input).shape, [2, 4, 4, 4]) + XCTAssertEqual( + AvgPool2d(kernelSize: 3, stride: 2, padding: 1)(input).shape, [2, 4, 4, 4]) + XCTAssertEqual( + MaxPool2d(kernelSize: 2, stride: 2, padding: 0)(input).shape, [2, 4, 4, 4]) + } + + func testPooling1dAnd3dPaddingShapes() { + let input1d = MLXArray(0 ..< (2 * 16 * 4), [2, 16, 4]).asType(.float32) + XCTAssertEqual( + MaxPool1d(kernelSize: 3, stride: 2, padding: 1)(input1d).shape, [2, 8, 4]) + XCTAssertEqual( + AvgPool1d(kernelSize: 3, stride: 2, padding: 1)(input1d).shape, [2, 8, 4]) + + let input3d = MLXArray(0 ..< (2 * 4 * 8 * 8 * 4), [2, 4, 8, 8, 4]).asType(.float32) + XCTAssertEqual( + MaxPool3d(kernelSize: 3, stride: 2, padding: 1)(input3d).shape, [2, 2, 4, 4, 4]) + XCTAssertEqual( + AvgPool3d(kernelSize: 3, stride: 2, padding: 1)(input3d).shape, [2, 2, 4, 4, 4]) + } + + func testPoolingAsymmetricPaddingPerAxis() { + let input = MLXArray(0 ..< (1 * 8 * 6 * 1), [1, 8, 6, 1]).asType(.float32) + + // padding only the width axis + XCTAssertEqual( + MaxPool2d(kernelSize: 2, stride: 2, padding: [0, 1])(input).shape, [1, 4, 4, 1]) + // and only the height axis + XCTAssertEqual( + MaxPool2d(kernelSize: 2, stride: 2, padding: [1, 0])(input).shape, [1, 5, 3, 1]) + } } diff --git a/Tests/MLXTests/PositionalEncodingTests.swift b/Tests/MLXTests/PositionalEncodingTests.swift new file mode 100644 index 000000000..9d5ac1502 --- /dev/null +++ b/Tests/MLXTests/PositionalEncodingTests.swift @@ -0,0 +1,167 @@ +// Copyright © 2026 Apple Inc. + +import Foundation +import MLX +import XCTest + +@testable import MLXNN + +/// Tests for layers whose formulas are easy to get subtly wrong -- these pin down +/// the parts that a shape check would not catch. +class PositionalEncodingTests: XCTestCase { + + // MARK: - ALiBi + + func testALiBiSlopesPowerOfTwo() { + // matches python `ALiBi.create_alibi_slope`: 2^(-8i/n) for i in 1...n + let slopes = ALiBi.alibiSlope(numHeads: 4).reshaped([-1]) + assertEqual( + slopes, MLXArray(converting: [0.25, 0.0625, 0.015625, 0.00390625]), atol: 1e-7) + } + + func testALiBiSlopesNotPowerOfTwo() { + // python does *not* extend the geometric series: it takes the slopes of the + // power of two below and pads with every other slope of the one above + assertEqual( + ALiBi.alibiSlope(numHeads: 3).reshaped([-1]), + MLXArray(converting: [0.0625, 0.00390625, 0.25]), atol: 1e-7) + assertEqual( + ALiBi.alibiSlope(numHeads: 6).reshaped([-1]), + MLXArray(converting: [0.25, 0.0625, 0.015625, 0.00390625, 0.5, 0.125]), + atol: 1e-7) + } + + func testALiBiMaskIsADistanceMatrix() { + // the mask must be -|i - j| * slope: a (q, k) matrix, not a broadcast of a + // single column (which collapsed to all zeros when q == k) + let scores = MLXArray.zeros([1, 2, 4, 4]) + let mask = ALiBi()(attentionScores: scores) + + XCTAssertEqual(mask.shape, [1, 2, 4, 4]) + + let slopes: [Float] = [0.0625, 0.00390625] // numHeads == 2 + for head in 0 ..< 2 { + for q in 0 ..< 4 { + for k in 0 ..< 4 { + let expected = -Float(abs(q - k)) * slopes[head] + XCTAssertEqual( + mask[0, head, q, k].item(Float.self), expected, accuracy: 1e-6, + "head \(head), q \(q), k \(k)") + } + } + } + } + + func testALiBiOffsetShiftsTheQueryPositions() { + // q positions start at `offset`, so with offset 1 and 3 queries the + // distances are |[1, 2, 3] - [0, 1, 2, 3]| + let scores = MLXArray.zeros([1, 1, 3, 4]) + let mask = ALiBi()(attentionScores: scores, offset: 1) + + XCTAssertEqual(mask.shape, [1, 1, 3, 4]) + + let slope: Float = 0.00390625 // numHeads == 1 + for q in 0 ..< 3 { + for k in 0 ..< 4 { + let expected = -Float(abs((q + 1) - k)) * slope + XCTAssertEqual( + mask[0, 0, q, k].item(Float.self), expected, accuracy: 1e-6, + "q \(q), k \(k)") + } + } + } + + func testALiBiAddsToTheScores() { + let scores = MLXArray.ones([1, 1, 3, 3]) + let mask = ALiBi()(attentionScores: MLXArray.zeros([1, 1, 3, 3])) + assertEqual(ALiBi()(attentionScores: scores), scores + mask) + } +} + +/// Recurrent layers: the formulas have branches that only run for some argument +/// combinations, which is where they drift from python. +class RecurrentTests: XCTestCase { + + private func parameters(_ module: Module) -> ModuleParameters { + // deterministic values so the tests do not depend on the random init + module.mapParameters { parameter in + let size = parameter.size + return MLX.arange(size, dtype: .float32).reshaped(parameter.shape) / Float(size) + - 0.5 + } + } + + func testGRUAppliesHiddenBiasWithoutAnIncomingHiddenState() { + // python applies `r * bhn` even when no hidden state is passed in; when + // this was missing, `bhn` had no effect on the first call at all + let gru = GRU(inputSize: 4, hiddenSize: 3) + gru.update(parameters: parameters(gru)) + + let x = MLX.arange(2 * 5 * 4, dtype: .float32).reshaped([2, 5, 4]) / 40 + + let before = gru(x) + + gru.update(parameters: ModuleParameters.unflattened([("bhn", MLXArray.ones([3]))])) + let after = gru(x) + + assertNotEqual(before, after) + } + + func testGRUMatchesAReferenceImplementation() { + let hiddenSize = 3 + let gru = GRU(inputSize: 4, hiddenSize: hiddenSize) + gru.update(parameters: parameters(gru)) + + let x = MLX.arange(1 * 2 * 4, dtype: .float32).reshaped([1, 2, 4]) / 8 + let result = gru(x) + + // the same formula written out, with no incoming hidden state + let wx = gru.wx + let wh = gru.wh + let b = gru.b! + let bhn = gru.bhn! + + let projected = addMM(b, x, wx.T) + var hidden: MLXArray? = nil + var steps = [MLXArray]() + + for step in 0 ..< x.dim(-2) { + var rz = projected[.ellipsis, step, ..<(2 * hiddenSize)] + var hiddenN: MLXArray? = nil + + if let hidden { + let projectedHidden = matmul(hidden, wh.T) + rz = rz + projectedHidden[.ellipsis, ..<(2 * hiddenSize)] + hiddenN = projectedHidden[.ellipsis, (2 * hiddenSize)...] + bhn + } + + rz = sigmoid(rz) + let r = rz[.ellipsis, ..