mark / Sources /Mark /Mark.swift
ainouche-abderahmane's picture
Release Mark v0.1.0: 38.7M parameter multilingual diacritics (INT8/FP16/FP32/CoreML)
b5f9085 verified
Raw History Blame Contribute Delete
9.69 kB
//
// Mark.swift
// Mark
//
// Author: Ainouche Abderahmane & Mythologic
// On-device multilingual diacritic and tone restoration across 37 languages.
// Core ML inference on Apple Neural Engine (Swift 6.3).
//
import CoreML
import Foundation
public actor Mark {
public static let supportedLanguages: [String] = [
"ak", "ar", "az", "ca", "cs", "cy", "ee", "es", "ff", "fr",
"ga", "gn", "ha", "he", "hr", "ht", "hu", "ig", "ku", "ln",
"lt", "lv", "mi", "pl", "pt", "qu", "ro", "sk", "sl", "sm",
"sr", "tk", "tr", "uz", "vi", "wo", "yo",
]
private let model: MLModel
private let tags: [String]
private let languageMap: [String: Int]
private let splitRegex: Regex<AnyRegexOutput>
public enum MarkError: LocalizedError {
case modelNotFound
case unsupportedLanguage(String)
case predictionFailed(String)
case resourceMissing(String)
public var errorDescription: String? {
switch self {
case .modelNotFound:
return "Mark Core ML model file not found in bundle or provided path."
case .unsupportedLanguage(let lang):
return "Unsupported language code '\(lang)'. Supported: \(Mark.supportedLanguages.joined(separator: ", "))"
case .predictionFailed(let reason):
return "Prediction failed: \(reason)"
case .resourceMissing(let name):
return "Required resource '\(name)' missing from bundle."
}
}
}
private static var defaultBundle: Bundle {
#if SWIFT_PACKAGE
return Bundle.module
#else
return Bundle.main
#endif
}
/// Initialize Mark with a compiled MLModel or load from a custom URL.
public init(modelURL: URL? = nil, configuration: MLModelConfiguration = MLModelConfiguration()) throws {
let resolvedURL: URL
if let customURL = modelURL {
resolvedURL = customURL
} else if let bundled = Self.defaultBundle.url(forResource: "mark", withExtension: "mlmodelc") {
resolvedURL = bundled
} else if let bundled = Bundle.main.url(forResource: "mark", withExtension: "mlmodelc") {
resolvedURL = bundled
} else {
throw MarkError.modelNotFound
}
self.model = try MLModel(contentsOf: resolvedURL, configuration: configuration)
guard let tagsURL = Self.defaultBundle.url(forResource: "tags", withExtension: "json") ??
Bundle.main.url(forResource: "tags", withExtension: "json") else {
throw MarkError.resourceMissing("tags.json")
}
let data = try Data(contentsOf: tagsURL)
self.tags = try JSONDecoder().decode([String].self, from: data)
var map: [String: Int] = [:]
for (index, code) in Mark.supportedLanguages.enumerated() {
map[code] = index
}
self.languageMap = map
// Precompile regex for whitespace/word tokenization
self.splitRegex = try Regex(#"\S+|\s+"#)
}
/// Restore diacritics and tones on arbitrary input text.
///
/// - Parameters:
/// - text: Raw unmarked input string.
/// - lang: Target ISO 639-1 language code (e.g. "yo", "ar", "es").
/// - margin: Confidence threshold over KEEP (default: 0.0). When set to 0.4+,
/// suppresses hypothesis flicker during streaming keyboard input.
public func restore(_ text: String, lang: String, margin: Float = 0.0) throws -> String {
guard !text.isEmpty else { return "" }
let normalizedLang = lang.trimmingCharacters(in: .whitespacesAndNewlines).lowercased()
guard let langId = languageMap[normalizedLang] else {
throw MarkError.unsupportedLanguage(lang)
}
let chunks = splitIntoChunks(text, maxBytes: 480)
var restoredChunks: [String] = []
for chunk in chunks {
let restored = try restoreChunk(chunk, langId: langId, margin: margin)
restoredChunks.append(restored)
}
return restoredChunks.joined()
}
/// Single chunk restoration (under 512 bytes) using MLShapedArray.
private func restoreChunk(_ chunk: String, langId: Int, margin: Float = 0.0) throws -> String {
let utf8Bytes = Array(chunk.utf8)
guard !utf8Bytes.isEmpty else { return "" }
let seqLen = min(utf8Bytes.count, 512)
// Type-safe MLShapedArray inputs (no manual memory manipulation)
let inputScalars: [Int32] = (0..<512).map { i in
i < seqLen ? Int32(utf8Bytes[i]) : Int32(0)
}
let inputShaped = MLShapedArray<Int32>(scalars: inputScalars, shape: [1, 512])
let langShaped = MLShapedArray<Int32>(scalars: [Int32(langId)], shape: [1])
let featureProvider = try MLDictionaryFeatureProvider(dictionary: [
"input_ids": MLFeatureValue(multiArray: MLMultiArray(inputShaped)),
"lang_ids": MLFeatureValue(multiArray: MLMultiArray(langShaped)),
])
let prediction = try model.prediction(from: featureProvider)
guard let logitsFeature = prediction.featureValue(for: "logits")?.multiArrayValue else {
throw MarkError.predictionFailed("Output logits tensor missing")
}
let shapedLogits = MLShapedArray<Float>(converting: logitsFeature)
let scalars = shapedLogits.scalars
let numTags = self.tags.count
var preds: [Int] = []
preds.reserveCapacity(seqLen)
// Fast contiguous scalar access for calibrated margin or argmax
for pos in 0..<seqLen {
let offset = pos * numTags
if margin > 0.0 && numTags > 1 {
let keepVal = scalars[offset + 0]
var bestVal = -Float.infinity
var bestIdx = 0
for tagIdx in 1..<numTags {
let val = scalars[offset + tagIdx]
if val > bestVal {
bestVal = val
bestIdx = tagIdx
}
}
preds.append((bestVal - keepVal >= margin) ? bestIdx : 0)
} else {
var maxVal = -Float.infinity
var maxIdx = 0
for tagIdx in 0..<numTags {
let val = scalars[offset + tagIdx]
if val > maxVal {
maxVal = val
maxIdx = tagIdx
}
}
preds.append(maxIdx)
}
}
return decodePrediction(chunk: chunk, utf8Bytes: utf8Bytes, predTags: preds)
}
/// Decode raw character bytes and predicted tag sequence into marked text.
private func decodePrediction(chunk: String, utf8Bytes: [UInt8], predTags: [Int]) -> String {
var result = ""
var byteOffset = 0
for char in chunk {
let charByteCount = String(char).utf8.count
let tagStr: String
if byteOffset < predTags.count {
let tagId = predTags[byteOffset]
tagStr = (tagId >= 0 && tagId < tags.count) ? tags[tagId] : "KEEP"
} else {
tagStr = "KEEP"
}
if tagStr.hasPrefix("P:") && tagStr.contains("->") {
let parts = tagStr.dropFirst(2).components(separatedBy: "->")
if parts.count == 2 {
result.append(parts[1])
} else {
result.append(char)
}
} else if tagStr.hasPrefix("M:") {
let hexMarks = tagStr.dropFirst(2).components(separatedBy: "+")
var combiningString = String(char)
for hex in hexMarks {
if let scalarVal = UInt32(hex, radix: 16), let scalar = UnicodeScalar(scalarVal) {
combiningString.append(String(scalar))
}
}
result.append(combiningString.precomposedStringWithCanonicalMapping)
} else {
result.append(char)
}
byteOffset += charByteCount
}
return result
}
/// Swift Regex tokenization preserving whitespace and boundaries.
private func splitIntoChunks(_ text: String, maxBytes: Int) -> [String] {
guard text.utf8.count > maxBytes else { return [text] }
let matches = text.matches(of: splitRegex)
let tokens = matches.map { String(text[$0.range]) }
var chunks: [String] = []
var currentChunk = ""
for token in tokens {
if (currentChunk + token).utf8.count <= maxBytes {
currentChunk += token
} else {
if !currentChunk.isEmpty {
chunks.append(currentChunk)
currentChunk = ""
}
if token.utf8.count <= maxBytes {
currentChunk = token
} else {
// Extremely long token exceeding maxBytes
for char in token {
let s = String(char)
if (currentChunk + s).utf8.count > maxBytes {
chunks.append(currentChunk)
currentChunk = s
} else {
currentChunk += s
}
}
}
}
}
if !currentChunk.isEmpty {
chunks.append(currentChunk)
}
return chunks
}
}