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
}
}
}