File size: 9,685 Bytes
98682c3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b5f9085
 
 
 
 
 
 
98682c3
 
 
 
 
 
 
 
 
 
 
b5f9085
98682c3
 
 
 
 
 
 
b5f9085
98682c3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b5f9085
98682c3
 
b5f9085
 
 
 
 
 
 
 
 
 
98682c3
b5f9085
 
 
 
 
 
 
 
 
 
 
 
98682c3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
//
//  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
    }
}