diff --git a/Source/MLX/Documentation.docc/free-functions.md b/Source/MLX/Documentation.docc/free-functions.md index e392e428..947b9564 100644 --- a/Source/MLX/Documentation.docc/free-functions.md +++ b/Source/MLX/Documentation.docc/free-functions.md @@ -119,9 +119,13 @@ operations as methods for convenience. - ``loadArray(url:stream:)`` - ``loadArrays(url:stream:)`` +- ``loadArrays(url:stream:progressHandler:)`` - ``loadArraysAndMetadata(url:stream:)`` +- ``loadArraysAndMetadata(url:stream:progressHandler:)`` - ``save(array:url:stream:)`` - ``save(arrays:metadata:url:stream:)`` +- ``withLoadProgressHandler(_:_:)-(_,()throws->R)`` +- ``withLoadProgressHandler(_:_:)-(_,()async throws->R)`` ### Logical diff --git a/Source/MLX/IO.swift b/Source/MLX/IO.swift index aee73a56..43954cc5 100644 --- a/Source/MLX/IO.swift +++ b/Source/MLX/IO.swift @@ -3,6 +3,94 @@ import Cmlx import Foundation +/// Byte-level progress for loading arrays from disk. +/// +/// `completedUnitCount` and `totalUnitCount` are bytes and describe a single file. +/// Progress callbacks for a given file are delivered in monotonically increasing +/// order, but callbacks for _different_ files may interleave -- use ``url`` to +/// aggregate progress when loading a model made of several `safetensors` shards. +public struct LoadProgress: Sendable, Equatable { + /// The file being read. + public let url: URL + + public let completedUnitCount: Int64 + public let totalUnitCount: Int64 + + public var fractionCompleted: Double { + guard totalUnitCount > 0 else { return 0 } + return min(1, max(0, Double(completedUnitCount) / Double(totalUnitCount))) + } + + public init(url: URL, completedUnitCount: Int64, totalUnitCount: Int64) { + self.url = url + self.completedUnitCount = completedUnitCount + self.totalUnitCount = totalUnitCount + } +} + +/// Holder for the scoped ``withLoadProgressHandler(_:_:)-(_,()throws->R)`` handler. +enum LoadProgressHandler { + + /// The stack of installed handlers -- the innermost scope wins. + @TaskLocal + static var handlers: [@Sendable (LoadProgress) -> Void] = [] + + static var current: (@Sendable (LoadProgress) -> Void)? { + handlers.last + } +} + +/// Evaluate the block with a scoped byte-progress handler for file loads. +/// +/// Any ``loadArrays(url:stream:)`` or ``loadArraysAndMetadata(url:stream:)`` performed +/// inside `body` reports byte progress to `handler`, without the call site having to pass +/// a progress handler explicitly. This makes it possible to drive a precise loading +/// progress bar for code -- such as a model loading library -- that you do not control: +/// +/// ```swift +/// let tracker = LoadProgressTracker(totalBytes: totalBytesOfSafetensors(in: directory)) +/// let model = try withLoadProgressHandler({ tracker.update($0) }) { +/// try loadModel(from: directory) +/// } +/// ``` +/// +/// Loading is lazy: progress is reported as the returned arrays are evaluated, so `body` +/// should include the evaluation of the loaded arrays. Arrays that are never evaluated +/// are never read, so the reported progress may legitimately stop short of the file size. +/// +/// - Note: `handler` is called from MLX worker threads, potentially concurrently for +/// different files, and is on the critical path of the read. It should be cheap and +/// must not call back into MLX loading. +/// +/// - Parameters: +/// - handler: the scoped progress handler +/// - body: the code where the handler is to be active +/// +/// ### See Also +/// - ``loadArrays(url:stream:progressHandler:)`` +public func withLoadProgressHandler( + _ handler: @escaping @Sendable (LoadProgress) -> Void, _ body: () throws -> R +) rethrows -> R { + try LoadProgressHandler.$handlers.withValue(LoadProgressHandler.handlers + [handler]) { + try body() + } +} + +/// Evaluate the block with a scoped byte-progress handler for file loads (async). +/// +/// See ``withLoadProgressHandler(_:_:)-(_,()throws->R)`` for details. +/// +/// - Parameters: +/// - handler: the scoped progress handler +/// - body: the code where the handler is to be active +public func withLoadProgressHandler( + _ handler: @escaping @Sendable (LoadProgress) -> Void, _ body: () async throws -> R +) async rethrows -> R { + try await LoadProgressHandler.$handlers.withValue(LoadProgressHandler.handlers + [handler]) { + try await body() + } +} + public enum LoadSaveError: Error { case unableToOpen(URL, String) case unknownExtension(String) @@ -117,6 +205,9 @@ public func loadArray(url: URL, stream: StreamOrDevice = .cpu) throws -> MLXArra /// - url: URL of file to load /// - stream: stream or device to evaluate on /// +/// - Note: when a scoped progress handler is installed with +/// ``withLoadProgressHandler(_:_:)-(_,()throws->R)`` this reports byte progress to it. +/// /// ### See Also /// - ``loadArray(url:stream:)`` /// - ``loadArraysAndMetadata(url:stream:)`` @@ -128,6 +219,10 @@ public func loadArrays(url: URL, stream: StreamOrDevice = .cpu) throws -> [Strin switch url.pathExtension { case "safetensors": + if let progressHandler = LoadProgressHandler.current { + return try loadArrays(url: url, stream: stream, progressHandler: progressHandler) + } + var r0 = mlx_map_string_to_array_new() var r1 = mlx_map_string_to_string_new() defer { mlx_map_string_to_array_free(r0) } @@ -143,12 +238,36 @@ public func loadArrays(url: URL, stream: StreamOrDevice = .cpu) throws -> [Strin } } +/// Load dictionary of ``MLXArray`` from a `safetensors` file, reporting byte progress as +/// lazy arrays are evaluated. +/// +/// - Parameters: +/// - url: URL of file to load +/// - stream: stream or device to evaluate on +/// - progressHandler: progress callback. This may be called from MLX worker threads. +/// Progress is reported in byte chunks while the returned lazy arrays are evaluated. +/// +/// ### See Also +/// - ``loadArrays(url:stream:)`` +/// - ``loadArraysAndMetadata(url:stream:progressHandler:)`` +public func loadArrays( + url: URL, stream: StreamOrDevice = .cpu, + progressHandler: @Sendable @escaping (LoadProgress) -> Void +) throws -> [String: MLXArray] { + let (arrays, _) = try loadArraysAndMetadata( + url: url, stream: stream, progressHandler: progressHandler) + return arrays +} + /// Load dictionary of ``MLXArray`` and metadata `[String:String]` from a `safetensors` file. /// /// - Parameters: /// - url: URL of file to load /// - stream: stream or device to evaluate on /// +/// - Note: when a scoped progress handler is installed with +/// ``withLoadProgressHandler(_:_:)-(_,()throws->R)`` this reports byte progress to it. +/// /// ### See Also /// - ``loadArrays(url:stream:)`` /// - ``loadArray(url:stream:)`` @@ -160,6 +279,11 @@ public func loadArraysAndMetadata(url: URL, stream: StreamOrDevice = .cpu) throw switch url.pathExtension { case "safetensors": + if let progressHandler = LoadProgressHandler.current { + return try loadArraysAndMetadata( + url: url, stream: stream, progressHandler: progressHandler) + } + var r0 = mlx_map_string_to_array_new() var r1 = mlx_map_string_to_string_new() defer { mlx_map_string_to_array_free(r0) } @@ -175,6 +299,44 @@ public func loadArraysAndMetadata(url: URL, stream: StreamOrDevice = .cpu) throw } } +/// Load dictionary of ``MLXArray`` and metadata from a `safetensors` file, reporting byte +/// progress as lazy arrays are evaluated. +/// +/// - Parameters: +/// - url: URL of file to load +/// - stream: stream or device to evaluate on +/// - progressHandler: progress callback. This may be called from MLX worker threads. +/// Progress is reported in byte chunks while the returned lazy arrays are evaluated. +/// +/// ### See Also +/// - ``loadArraysAndMetadata(url:stream:)`` +/// - ``loadArrays(url:stream:progressHandler:)`` +public func loadArraysAndMetadata( + url: URL, stream: StreamOrDevice = .cpu, + progressHandler: @Sendable @escaping (LoadProgress) -> Void +) throws -> ([String: MLXArray], [String: String]) { + precondition(url.isFileURL) + + switch url.pathExtension { + case "safetensors": + var r0 = mlx_map_string_to_array_new() + var r1 = mlx_map_string_to_string_new() + defer { mlx_map_string_to_array_free(r0) } + defer { mlx_map_string_to_string_free(r1) } + + let reader = try new_mlx_io_reader_fileIO(url, progressHandler: progressHandler) + defer { mlx_io_reader_free(reader) } + + _ = try withError { + mlx_load_safetensors_reader(&r0, &r1, reader, stream.ctx) + } + + return (mlx_map_array_values(r0), mlx_map_string_values(r1)) + default: + throw LoadSaveError.unknownExtension(url.pathExtension) + } +} + // MARK: - Memory I/O private class IOState { @@ -187,6 +349,160 @@ private class IOState { } } +private final class FileIOState { + private static let maximumReadChunkSize = 4 * 1024 * 1024 + + private let descriptor: CInt + private let lock = NSLock() + private let progressLock = NSLock() + private var offset: Int64 = 0 + private var completedUnitCount: Int64 = 0 + private var readError: String? + private let progressHandler: @Sendable (LoadProgress) -> Void + private let labelPointer: UnsafeMutablePointer + private let url: URL + + let totalUnitCount: Int64 + + init(url: URL, progressHandler: @Sendable @escaping (LoadProgress) -> Void) throws { + let path = url.path(percentEncoded: false) + let descriptor = path.withCString { open($0, O_RDONLY) } + guard descriptor >= 0 else { + throw LoadSaveError.unableToOpen(url, String(cString: strerror(errno))) + } + + var statBuffer = stat() + guard fstat(descriptor, &statBuffer) == 0 else { + let message = String(cString: strerror(errno)) + close(descriptor) + throw LoadSaveError.unableToOpen(url, message) + } + + guard let labelPointer = strdup("file \(path)") else { + close(descriptor) + throw LoadSaveError.unableToOpen(url, String(cString: strerror(errno))) + } + + self.descriptor = descriptor + self.totalUnitCount = max(0, Int64(statBuffer.st_size)) + self.progressHandler = progressHandler + self.labelPointer = labelPointer + self.url = url + + progressHandler( + .init(url: url, completedUnitCount: 0, totalUnitCount: totalUnitCount)) + } + + deinit { + close(descriptor) + free(labelPointer) + } + + var isOpen: Bool { + descriptor >= 0 + } + + var good: Bool { + lock.withLock { + readError == nil + } + } + + var label: UnsafePointer { + UnsafePointer(labelPointer) + } + + func tell() -> Int { + lock.withLock { + Int(offset) + } + } + + func seek(offset newOffset: Int64, whence: Int32) { + lock.withLock { + switch whence { + case SEEK_SET: + offset = newOffset + case SEEK_CUR: + offset += newOffset + case SEEK_END: + offset = totalUnitCount + newOffset + default: + break + } + } + } + + func read(to data: UnsafeMutablePointer?, count: Int) { + guard let data else { return } + + let readOffset = lock.withLock { + offset + } + let bytesRead = read(to: data, count: count, offset: readOffset) + + lock.withLock { + offset += Int64(bytesRead) + } + } + + func read(to data: UnsafeMutablePointer?, count: Int, offset readOffset: Int64) { + guard let data else { return } + + _ = read(to: data, count: count, offset: readOffset) + } + + @discardableResult + private func read(to data: UnsafeMutablePointer, count: Int, offset readOffset: Int64) + -> Int + { + var totalRead = 0 + while totalRead < count { + let chunkSize = min(count - totalRead, Self.maximumReadChunkSize) + let bytesRead = pread( + descriptor, + UnsafeMutableRawPointer(data.advanced(by: totalRead)), + chunkSize, + off_t(readOffset + Int64(totalRead))) + guard bytesRead > 0 else { + recordReadError(bytesRead: bytesRead, requestedCount: count - totalRead) + break + } + totalRead += bytesRead + reportProgress(bytesRead: bytesRead) + } + + return totalRead + } + + private func recordReadError(bytesRead: Int, requestedCount: Int) { + let message: String + if bytesRead < 0 { + message = String(cString: strerror(errno)) + } else { + message = "unexpected end of file while reading \(requestedCount) bytes" + } + lock.withLock { + if readError == nil { + readError = message + } + } + } + + private func reportProgress(bytesRead: Int) { + guard bytesRead > 0 else { return } + + progressLock.withLock { + completedUnitCount = min(totalUnitCount, completedUnitCount + Int64(bytesRead)) + let progress = LoadProgress( + url: url, + completedUnitCount: completedUnitCount, + totalUnitCount: totalUnitCount) + progressHandler(progress) + } + } +} + private let label: StaticString = "\0" private func getData(_ writer: mlx_io_writer) -> Data { @@ -214,7 +530,7 @@ private func new_mlx_io_vtable_dataIO() -> mlx_io_vtable { case SEEK_CUR: state.offset += Int(offset) case SEEK_END: - state.offset = state.offset - Int(offset) + state.offset = state.data.count + Int(offset) default: break } @@ -260,6 +576,49 @@ private func new_mlx_io_reader_dataIO(_ data: Data) -> mlx_io_reader { return mlx_io_reader_new(ptr, new_mlx_io_vtable_dataIO()) } +private func new_mlx_io_vtable_fileIO() -> mlx_io_vtable { + mlx_io_vtable { ptr in + guard let ptr else { return false } + return Unmanaged.fromOpaque(ptr).takeUnretainedValue().isOpen + } good: { ptr in + guard let ptr else { return false } + let state = Unmanaged.fromOpaque(ptr).takeUnretainedValue() + return state.isOpen && state.good + } tell: { ptr in + let state = Unmanaged.fromOpaque(ptr!).takeUnretainedValue() + return state.tell() + + } seek: { ptr, offset, whence in + let state = Unmanaged.fromOpaque(ptr!).takeUnretainedValue() + state.seek(offset: Int64(offset), whence: whence) + + } read: { ptr, data, n in + let state = Unmanaged.fromOpaque(ptr!).takeUnretainedValue() + state.read(to: data, count: n) + + } read_at_offset: { ptr, data, n, offset in + let state = Unmanaged.fromOpaque(ptr!).takeUnretainedValue() + state.read(to: data, count: n, offset: Int64(offset)) + + } write: { _, _, _ in + + } label: { ptr in + let state = Unmanaged.fromOpaque(ptr!).takeUnretainedValue() + return state.label + + } free: { ptr in + Unmanaged.fromOpaque(ptr!).release() + } +} + +private func new_mlx_io_reader_fileIO( + _ url: URL, progressHandler: @Sendable @escaping (LoadProgress) -> Void +) throws -> mlx_io_reader { + let ptr = Unmanaged.passRetained(try FileIOState(url: url, progressHandler: progressHandler)) + .toOpaque() + return mlx_io_reader_new(ptr, new_mlx_io_vtable_fileIO()) +} + private func new_mlx_io_writer_dataIO() -> mlx_io_writer { let ptr = Unmanaged.passRetained(IOState()).toOpaque() return mlx_io_writer_new(ptr, new_mlx_io_vtable_dataIO()) diff --git a/Tests/MLXTests/SaveTests.swift b/Tests/MLXTests/SaveTests.swift index 6f700df4..02f9fd22 100644 --- a/Tests/MLXTests/SaveTests.swift +++ b/Tests/MLXTests/SaveTests.swift @@ -7,6 +7,45 @@ import MLX import XCTest +import os + +private final class ProgressRecorder: Sendable { + private let progress = OSAllocatedUnfairLock(initialState: [LoadProgress]()) + + func record(_ progress: LoadProgress) { + self.progress.withLock { values in + values.append(progress) + } + } + + var reported: [LoadProgress] { + progress.withLock { values in + values + } + } + + var values: [Double] { + reported.map { $0.fractionCompleted } + } + + /// Fractions reported for a single file, in order. + func values(for url: URL) -> [Double] { + reported.filter { $0.url == url }.map { $0.fractionCompleted } + } + + /// Aggregate fraction across every file seen, by bytes. + var aggregateFraction: Double { + var completed = [URL: Int64]() + var total = [URL: Int64]() + for progress in reported { + completed[progress.url] = progress.completedUnitCount + total[progress.url] = progress.totalUnitCount + } + let totalBytes = total.values.reduce(0, +) + guard totalBytes > 0 else { return 0 } + return Double(completed.values.reduce(0, +)) / Double(totalBytes) + } +} final class SaveTests: XCTestCase { @@ -16,7 +55,6 @@ final class SaveTests: XCTestCase { ) override func setUpWithError() throws { - setDefaultDevice() try FileManager.default.createDirectory( at: temporaryPath, withIntermediateDirectories: false @@ -28,72 +66,234 @@ final class SaveTests: XCTestCase { } public func testSaveArrays() throws { - let safetensorsPath = temporaryPath.appending( - path: "arrays.safetensors", - directoryHint: .notDirectory - ) + try MLX.Device.withDefaultDevice(.cpu) { + let safetensorsPath = temporaryPath.appending( + path: "arrays.safetensors", + directoryHint: .notDirectory + ) + + let arrays: [String: MLXArray] = [ + "foo": MLX.ones([1, 2]), + "bar": MLX.zeros([2, 1]), + ] + + try MLX.save(arrays: arrays, url: safetensorsPath) + + let loadedArrays = try MLX.loadArrays(url: safetensorsPath) + XCTAssertEqual(loadedArrays.keys.sorted(), arrays.keys.sorted()) + + assertEqual(try XCTUnwrap(loadedArrays["foo"]), try XCTUnwrap(arrays["foo"])) + assertEqual(try XCTUnwrap(loadedArrays["bar"]), try XCTUnwrap(arrays["bar"])) + } + } + + public func testLoadArraysProgressReportsThroughEvaluation() throws { + try MLX.Device.withDefaultDevice(.cpu) { + let safetensorsPath = temporaryPath.appending( + path: "arrays.safetensors", + directoryHint: .notDirectory + ) + + let arrays: [String: MLXArray] = [ + "foo": MLX.ones([128, 128]), + "bar": MLX.zeros([64, 256]), + ] + try MLX.save(arrays: arrays, url: safetensorsPath) + + let recorder = ProgressRecorder() + let loadedArrays = try MLX.loadArrays( + url: safetensorsPath + ) { @Sendable progress in + recorder.record(progress) + } + + assertEqual(try XCTUnwrap(loadedArrays["foo"]), try XCTUnwrap(arrays["foo"])) + assertEqual(try XCTUnwrap(loadedArrays["bar"]), try XCTUnwrap(arrays["bar"])) - let arrays: [String: MLXArray] = [ - "foo": MLX.ones([1, 2]), - "bar": MLX.zeros([2, 1]), - ] + let fractions = recorder.values + XCTAssertGreaterThan(fractions.count, 1) + XCTAssertEqual(fractions.first, 0) + XCTAssertEqual(fractions.last, 1) + XCTAssertEqual(fractions, fractions.sorted()) + } + } + + public func testLoadArraysProgressFailsOnTruncatedTensorData() throws { + try MLX.Device.withDefaultDevice(.cpu) { + let safetensorsPath = temporaryPath.appending( + path: "truncated.safetensors", + directoryHint: .notDirectory + ) + + let arrays: [String: MLXArray] = [ + "foo": MLX.ones([128, 128]), + "bar": MLX.zeros([64, 256]), + ] + try MLX.save(arrays: arrays, url: safetensorsPath) + + var data = try Data(contentsOf: safetensorsPath) + data.removeLast(32) + try data.write(to: safetensorsPath) + + // A truncated file has to be reported either eagerly, while the header is + // parsed (mlx >= 0.32.1 validates the tensor data offsets against the size of + // the file), or lazily, when the arrays are evaluated and the read fails + // (ml-explore/mlx#3742 + ml-explore/mlx-c#126). + var thrownError: Error? + do { + let loadedArrays = try MLX.loadArrays(url: safetensorsPath) { _ in } + try checkedEval(Array(loadedArrays.values) as [Any]) + } catch { + thrownError = error + } + + if thrownError == nil { + throw XCTSkip( + """ + the vendored mlx/mlx-c silently ignores a failed read from a custom \ + io reader -- requires mlx >= 0.32.1 (ml-explore/mlx#3742) and \ + ml-explore/mlx-c#126 + """) + } + } + } + + public func testScopedLoadProgressHandler() throws { + try MLX.Device.withDefaultDevice(.cpu) { + let safetensorsPath = temporaryPath.appending( + path: "scoped.safetensors", + directoryHint: .notDirectory + ) + + let arrays: [String: MLXArray] = [ + "foo": MLX.ones([128, 128]), + "bar": MLX.zeros([64, 256]), + ] + try MLX.save(arrays: arrays, url: safetensorsPath) + + let recorder = ProgressRecorder() + + // note: the plain loadArrays(url:) -- no progress handler passed at the call site + let loadedArrays = try withLoadProgressHandler({ @Sendable in recorder.record($0) }) { + let loadedArrays = try MLX.loadArrays(url: safetensorsPath) + MLX.eval(Array(loadedArrays.values)) + return loadedArrays + } - try MLX.save(arrays: arrays, url: safetensorsPath) + assertEqual(try XCTUnwrap(loadedArrays["foo"]), try XCTUnwrap(arrays["foo"])) - let loadedArrays = try MLX.loadArrays(url: safetensorsPath) - XCTAssertEqual(loadedArrays.keys.sorted(), arrays.keys.sorted()) + let fractions = recorder.values + XCTAssertGreaterThan(fractions.count, 1) + XCTAssertEqual(fractions.first, 0) + XCTAssertEqual(fractions.last, 1) + XCTAssertEqual(fractions, fractions.sorted()) + XCTAssertEqual(Set(recorder.reported.map(\.url)), [safetensorsPath]) + } + } + + public func testScopedLoadProgressHandlerIsScoped() throws { + try MLX.Device.withDefaultDevice(.cpu) { + let safetensorsPath = temporaryPath.appending( + path: "unscoped.safetensors", + directoryHint: .notDirectory + ) + try MLX.save(arrays: ["foo": MLX.ones([128, 128])], url: safetensorsPath) + + let recorder = ProgressRecorder() + withLoadProgressHandler({ @Sendable in recorder.record($0) }) { + } - assertEqual(try XCTUnwrap(loadedArrays["foo"]), try XCTUnwrap(arrays["foo"])) - assertEqual(try XCTUnwrap(loadedArrays["bar"]), try XCTUnwrap(arrays["bar"])) + // outside the scope nothing is reported + let loadedArrays = try MLX.loadArrays(url: safetensorsPath) + MLX.eval(Array(loadedArrays.values)) + + XCTAssertTrue(recorder.reported.isEmpty) + } + } + + public func testScopedLoadProgressAggregatesAcrossFiles() throws { + try MLX.Device.withDefaultDevice(.cpu) { + let shards = try (0 ..< 3).map { index -> URL in + let url = temporaryPath.appending( + path: "shard-\(index).safetensors", + directoryHint: .notDirectory + ) + try MLX.save(arrays: ["w\(index)": MLX.ones([64, 128])], url: url) + return url + } + + let recorder = ProgressRecorder() + try withLoadProgressHandler({ @Sendable in recorder.record($0) }) { + // this mimics a model loader: several shards loaded lazily, then evaluated + var weights = [String: MLXArray]() + for url in shards { + let (w, _) = try MLX.loadArraysAndMetadata(url: url) + weights.merge(w) { _, new in new } + } + MLX.eval(Array(weights.values)) + } + + XCTAssertEqual(Set(recorder.reported.map(\.url)), Set(shards)) + for url in shards { + XCTAssertEqual(recorder.values(for: url).last, 1) + } + XCTAssertEqual(recorder.aggregateFraction, 1, accuracy: 1e-9) + } } public func testSaveArray() throws { - // single array npy file - let path = temporaryPath.appending( - path: "array.npy", - directoryHint: .notDirectory - ) + try MLX.Device.withDefaultDevice(.cpu) { + // single array npy file + let path = temporaryPath.appending( + path: "array.npy", + directoryHint: .notDirectory + ) - let array = MLX.ones([2, 4]) + let array = MLX.ones([2, 4]) - try MLX.save(array: array, url: path) + try MLX.save(array: array, url: path) - let loaded = try MLX.loadArray(url: path) + let loaded = try MLX.loadArray(url: path) - assertEqual(array, loaded) + assertEqual(array, loaded) + } } public func testSaveArraysData() throws { - let arrays: [String: MLXArray] = [ - "foo": MLX.ones([1, 2]), - "bar": MLX.zeros([2, 1]), - ] + try MLX.Device.withDefaultDevice(.cpu) { + let arrays: [String: MLXArray] = [ + "foo": MLX.ones([1, 2]), + "bar": MLX.zeros([2, 1]), + ] - let data = try saveToData(arrays: arrays) - let loadedArrays = try loadArrays(data: data) - XCTAssertEqual(loadedArrays.keys.sorted(), arrays.keys.sorted()) + let data = try saveToData(arrays: arrays) + let loadedArrays = try loadArrays(data: data) + XCTAssertEqual(loadedArrays.keys.sorted(), arrays.keys.sorted()) - assertEqual(try XCTUnwrap(loadedArrays["foo"]), try XCTUnwrap(arrays["foo"])) - assertEqual(try XCTUnwrap(loadedArrays["bar"]), try XCTUnwrap(arrays["bar"])) + assertEqual(try XCTUnwrap(loadedArrays["foo"]), try XCTUnwrap(arrays["foo"])) + assertEqual(try XCTUnwrap(loadedArrays["bar"]), try XCTUnwrap(arrays["bar"])) + } } public func testSaveArraysMetadataData() throws { - let arrays: [String: MLXArray] = [ - "foo": MLX.ones([1, 2]), - "bar": MLX.zeros([2, 1]), - ] - let metadata = [ - "key": "value", - "key2": "value2", - ] - - let data = try saveToData(arrays: arrays, metadata: metadata) - let (loadedArrays, loadedMetadata) = try loadArraysAndMetadata(data: data) - XCTAssertEqual(loadedArrays.keys.sorted(), arrays.keys.sorted()) - - assertEqual(try XCTUnwrap(loadedArrays["foo"]), try XCTUnwrap(arrays["foo"])) - assertEqual(try XCTUnwrap(loadedArrays["bar"]), try XCTUnwrap(arrays["bar"])) - XCTAssertEqual(loadedMetadata, metadata) + try MLX.Device.withDefaultDevice(.cpu) { + let arrays: [String: MLXArray] = [ + "foo": MLX.ones([1, 2]), + "bar": MLX.zeros([2, 1]), + ] + let metadata = [ + "key": "value", + "key2": "value2", + ] + + let data = try saveToData(arrays: arrays, metadata: metadata) + let (loadedArrays, loadedMetadata) = try loadArraysAndMetadata(data: data) + XCTAssertEqual(loadedArrays.keys.sorted(), arrays.keys.sorted()) + + assertEqual(try XCTUnwrap(loadedArrays["foo"]), try XCTUnwrap(arrays["foo"])) + assertEqual(try XCTUnwrap(loadedArrays["bar"]), try XCTUnwrap(arrays["bar"])) + XCTAssertEqual(loadedMetadata, metadata) + } } }