Section 9/171 menit

9. Updatable Models dengan MLUpdateTask

9. Updatable Models dengan MLUpdateTask

Core ML mendukung fine-tuning model on-device menggunakan data pengguna — tanpa mengirim data ke server. Ini ideal untuk personalisasi dan federated learning.

Model yang Bisa Di-update

Model harus di-buat dengan isUpdatable = True saat konversi. Layer terakhir (biasanya fully-connected) yang di-update. Di Xcode, ini ditandai dengan ikon khusus di model inspector.

Tipe model yang mendukung updating:

  • Neural network classifier dengan updatable layers
  • KNN classifier (k-Nearest Neighbor)
  • Image classifier dengan transfer learning layer

MLUpdateTask untuk Training On-Device

swift
// MARK: - On-Device Model Personalization

@MainActor
@Observable
final class PersonalizedClassifier {
    var trainingState: TrainingState = .idle
    private var currentModelURL: URL
    private var updateTask: MLUpdateTask?

    enum TrainingState {
        case idle
        case training(progress: Double)
        case completed
        case error(String)
    }

    struct TrainingExample {
        let image: CVPixelBuffer
        let label: String
    }

    init() {
        // Mulai dari model base yang di-bundle di app
        currentModelURL = Bundle.main.url(forResource: "BaseClassifier", withExtension: "mlmodelc")!
    }

    func trainWithExamples(_ examples: [TrainingExample]) async throws {
        guard !examples.isEmpty else { return }
        trainingState = .training(progress: 0)

        // Buat training data sebagai MLArrayBatchProvider
        let providers: [MLFeatureProvider] = try examples.map { example in
            let input = UpdatableModelInput(
                image: example.image,
                label: example.label
            )
            return input
        }
        let trainingData = MLArrayBatchProvider(array: providers)

        // Konfigurasi update task
        let updateConfig = MLUpdateTask.progressHandlers(
            forStoreAt: currentModelURL,
            trainingData: trainingData,
            configuration: MLModelConfiguration()
        ) { [weak self] context in
            // Progress callback (dipanggil di background thread)
            let progress = context.metrics[.epochIndex] as? Double ?? 0
            let maxEpoch = context.metrics[.maxEpochIndex] as? Double ?? 1
            let progressFraction = progress / max(maxEpoch, 1)
            Task { @MainActor [weak self] in
                self?.trainingState = .training(progress: progressFraction)
            }
        } completionHandler: { [weak self] context in
            Task { @MainActor [weak self] in
                guard let self else { return }
                if context.task.error != nil {
                    self.trainingState = .error(context.task.error!.localizedDescription)
                    return
                }

                // Simpan model yang di-update ke Documents
                do {
                    let updatedURL = try self.saveUpdatedModel(from: context)
                    self.currentModelURL = updatedURL
                    self.trainingState = .completed
                } catch {
                    self.trainingState = .error(error.localizedDescription)
                }
            }
        }

        let task = try MLUpdateTask(
            forModelAt: currentModelURL,
            trainingData: trainingData,
            configuration: MLModelConfiguration(),
            progressHandlers: updateConfig
        )

        self.updateTask = task
        task.resume()

        // Tunggu selesai
        try await withCheckedThrowingContinuation { (continuation: CheckedContinuation<Void, Error>) in
            task.resume()
        }
    }

    private func saveUpdatedModel(from context: MLUpdateContext) throws -> URL {
        let documentsURL = FileManager.default.urls(
            for: .documentDirectory, in: .userDomainMask
        ).first!
        let updatedModelURL = documentsURL.appending(component: "PersonalizedModel.mlmodelc")

        // Hapus model lama jika ada
        try? FileManager.default.removeItem(at: updatedModelURL)

        // Simpan model yang baru di-update
        try context.model.write(to: updatedModelURL)
        return updatedModelURL
    }

    func cancelTraining() {
        updateTask?.cancel()
        trainingState = .idle
    }

    func resetToBaseModel() {
        let baseURL = Bundle.main.url(forResource: "BaseClassifier", withExtension: "mlmodelc")!
        currentModelURL = baseURL
        trainingState = .idle
    }
}

// Placeholder input type
class UpdatableModelInput: MLFeatureProvider {
    let image: CVPixelBuffer
    let label: String
    var featureNames: Set<String> { ["image", "classLabel"] }
    init(image: CVPixelBuffer, label: String) { self.image = image; self.label = label }
    func featureValue(for featureName: String) -> MLFeatureValue? {
        switch featureName {
        case "image": return try? MLFeatureValue(pixelBuffer: image, pixelFormatType: kCVPixelFormatType_32BGRA)
        case "classLabel": return MLFeatureValue(string: label)
        default: return nil
        }
    }
}