Compare commits
13 Commits
88aeff7935
..
v2
| Author | SHA1 | Date | |
|---|---|---|---|
| b33947feaf | |||
| c0347d5213 | |||
| e41e244b25 | |||
| 319f29bf69 | |||
| 91d44a924d | |||
| 6af56f58ea | |||
| 31d5e8adaf | |||
| c48983a413 | |||
| 492b779634 | |||
| f122d854dc | |||
| 5e060c7aea | |||
| dbec6b20ea | |||
| e7a94b3203 |
+13
-1
@@ -7,6 +7,7 @@ let package = Package(
|
|||||||
products: [
|
products: [
|
||||||
.library(name: "MarkBase", targets: ["MarkBase"]),
|
.library(name: "MarkBase", targets: ["MarkBase"]),
|
||||||
.executable(name: "MarkBaseServer", targets: ["MarkBaseServer"]),
|
.executable(name: "MarkBaseServer", targets: ["MarkBaseServer"]),
|
||||||
|
.executable(name: "EmbeddingServer", targets: ["EmbeddingServer"]),
|
||||||
.executable(name: "CLITest", targets: ["CLITest"]),
|
.executable(name: "CLITest", targets: ["CLITest"]),
|
||||||
],
|
],
|
||||||
dependencies: [
|
dependencies: [
|
||||||
@@ -16,7 +17,7 @@ let package = Package(
|
|||||||
targets: [
|
targets: [
|
||||||
.target(
|
.target(
|
||||||
name: "MarkBase",
|
name: "MarkBase",
|
||||||
exclude: ["Metal/MetalKernels.metal", "Metal/OptimizedKernels.metal", "Metal/FusionKernels.metal", "Metal/MetalKernels.metallib", "Metal/metallib"],
|
exclude: ["Metal/MetalKernels.metal", "Metal/OptimizedKernels.metal", "Metal/FusionKernels.metal", "Metal/EmbeddingKernels.metal", "Metal/MetalKernels.metallib", "Metal/metallib"],
|
||||||
linkerSettings: [
|
linkerSettings: [
|
||||||
.linkedFramework("Metal"),
|
.linkedFramework("Metal"),
|
||||||
.linkedFramework("Foundation"),
|
.linkedFramework("Foundation"),
|
||||||
@@ -34,6 +35,17 @@ let package = Package(
|
|||||||
.linkedFramework("Foundation"),
|
.linkedFramework("Foundation"),
|
||||||
]
|
]
|
||||||
),
|
),
|
||||||
|
.executableTarget(
|
||||||
|
name: "EmbeddingServer",
|
||||||
|
dependencies: [
|
||||||
|
"MarkBase",
|
||||||
|
.product(name: "Hummingbird", package: "hummingbird"),
|
||||||
|
],
|
||||||
|
linkerSettings: [
|
||||||
|
.linkedFramework("Metal"),
|
||||||
|
.linkedFramework("Foundation"),
|
||||||
|
]
|
||||||
|
),
|
||||||
.executableTarget(
|
.executableTarget(
|
||||||
name: "CLITest",
|
name: "CLITest",
|
||||||
dependencies: ["MarkBase"],
|
dependencies: ["MarkBase"],
|
||||||
|
|||||||
@@ -6,14 +6,14 @@
|
|||||||
<string>com.markbase.embedding</string>
|
<string>com.markbase.embedding</string>
|
||||||
<key>ProgramArguments</key>
|
<key>ProgramArguments</key>
|
||||||
<array>
|
<array>
|
||||||
<string>/Users/accusys/MarkBaseEngine/.build/arm64-apple-macosx/release/MarkBaseServer</string>
|
<string>/Users/accusys/MarkBaseEngine/.build/arm64-apple-macosx/release/EmbeddingServer</string>
|
||||||
<string>E4B-MarkBase</string>
|
<string>embeddinggemma-300m</string>
|
||||||
<string>8084</string>
|
<string>8084</string>
|
||||||
</array>
|
</array>
|
||||||
<key>RunAtLoad</key>
|
<key>RunAtLoad</key>
|
||||||
<true/>
|
<true/>
|
||||||
<key>KeepAlive</key>
|
<key>KeepAlive</key>
|
||||||
<true/>
|
<false/>
|
||||||
<key>StandardOutPath</key>
|
<key>StandardOutPath</key>
|
||||||
<string>/Users/accusys/MarkBaseEngine/logs/embedding.log</string>
|
<string>/Users/accusys/MarkBaseEngine/logs/embedding.log</string>
|
||||||
<key>StandardErrorPath</key>
|
<key>StandardErrorPath</key>
|
||||||
|
|||||||
@@ -0,0 +1,98 @@
|
|||||||
|
import Foundation
|
||||||
|
import MarkBase
|
||||||
|
import Hummingbird
|
||||||
|
|
||||||
|
@main
|
||||||
|
struct EmbeddingServerMain {
|
||||||
|
static func main() async throws {
|
||||||
|
let args = CommandLine.arguments
|
||||||
|
let modelName = args.count > 1 ? args[1] : "embeddinggemma-300m"
|
||||||
|
let port = args.count > 2 ? Int(args[2]) ?? 8084 : 8084
|
||||||
|
|
||||||
|
let modelPath = NSString(string: "~/MarkBaseEngine/models/\(modelName)").expandingTildeInPath
|
||||||
|
|
||||||
|
print("MarkBaseEngine Embedding Server")
|
||||||
|
print(" Model: \(modelName)")
|
||||||
|
print(" Port: \(port)")
|
||||||
|
print(" Path: \(modelPath)")
|
||||||
|
|
||||||
|
let engine = try MarkBaseEngine(autoCompile: true)
|
||||||
|
let embedModel = try EmbeddingGemmaModel(modelDir: modelPath, engine: engine)
|
||||||
|
let layers = embedModel.config.numHiddenLayers
|
||||||
|
let hiddenSize = embedModel.config.hiddenSize
|
||||||
|
|
||||||
|
print("EmbeddingGemma loaded (\(layers) layers, hidden=\(hiddenSize))")
|
||||||
|
|
||||||
|
// Use actor to serialize access to embedModel
|
||||||
|
let embedder = Embedder(model: embedModel)
|
||||||
|
|
||||||
|
let router = Router()
|
||||||
|
|
||||||
|
router.get("/") { _, _ in
|
||||||
|
return "{\"server\":\"MarkBaseEngine Embedding\",\"model\":\"embeddinggemma-300m\",\"layers\":\(layers),\"hidden_size\":\(hiddenSize)}"
|
||||||
|
}
|
||||||
|
|
||||||
|
router.get("/health") { _, _ in
|
||||||
|
return "{\"status\":\"healthy\",\"model\":\"embeddinggemma-300m\",\"layers\":\(layers)}"
|
||||||
|
}
|
||||||
|
|
||||||
|
router.post("/v1/embeddings") { request, _ in
|
||||||
|
let buffer = try await request.body.collect(upTo: .max)
|
||||||
|
let data = Data(buffer: buffer)
|
||||||
|
|
||||||
|
guard let json = try JSONSerialization.jsonObject(with: data) as? [String: Any],
|
||||||
|
let input = json["input"] else {
|
||||||
|
return "{\"error\":\"missing 'input' field\"}"
|
||||||
|
}
|
||||||
|
|
||||||
|
let inputs: [String]
|
||||||
|
if let str = input as? String { inputs = [str] }
|
||||||
|
else if let arr = input as? [String] { inputs = arr }
|
||||||
|
else { return "{\"error\":\"'input' must be string or array\"}" }
|
||||||
|
|
||||||
|
var embeddings: [[String: Any]] = []
|
||||||
|
for (i, text) in inputs.enumerated() {
|
||||||
|
let result = try await embedder.embed(text: text)
|
||||||
|
embeddings.append(["object": "embedding", "index": i, "embedding": result.embedding, "usage_ms": result.ms])
|
||||||
|
}
|
||||||
|
|
||||||
|
let id = UUID().uuidString
|
||||||
|
let ts = Int(Date().timeIntervalSince1970)
|
||||||
|
let response: [String: Any] = [
|
||||||
|
"id": id, "object": "list", "created": ts, "model": "embeddinggemma-300m",
|
||||||
|
"data": embeddings
|
||||||
|
]
|
||||||
|
|
||||||
|
let jsonData = try JSONSerialization.data(withJSONObject: response)
|
||||||
|
return String(data: jsonData, encoding: .utf8) ?? "{}"
|
||||||
|
}
|
||||||
|
|
||||||
|
let app = Application(
|
||||||
|
router: router,
|
||||||
|
configuration: .init(address: .hostname("0.0.0.0", port: port))
|
||||||
|
)
|
||||||
|
|
||||||
|
print("Server starting on port \(port)...")
|
||||||
|
try await app.run()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
struct EmbeddingResult {
|
||||||
|
let embedding: [Float]
|
||||||
|
let ms: Int
|
||||||
|
}
|
||||||
|
|
||||||
|
actor Embedder {
|
||||||
|
private let model: EmbeddingGemmaModel
|
||||||
|
|
||||||
|
init(model: EmbeddingGemmaModel) {
|
||||||
|
self.model = model
|
||||||
|
}
|
||||||
|
|
||||||
|
func embed(text: String) throws -> EmbeddingResult {
|
||||||
|
let t0 = Date()
|
||||||
|
let embedding = try model.embed(text: text)
|
||||||
|
let ms = Int(Date().timeIntervalSince(t0) * 1000)
|
||||||
|
return EmbeddingResult(embedding: embedding, ms: ms)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -2,7 +2,7 @@ import Foundation
|
|||||||
import Metal
|
import Metal
|
||||||
import Accelerate
|
import Accelerate
|
||||||
|
|
||||||
/// EmbeddingGemmaConfig - Configuration for EmbeddingGemma model
|
/// EmbeddingGemma configuration
|
||||||
public struct EmbeddingGemmaConfig: Codable {
|
public struct EmbeddingGemmaConfig: Codable {
|
||||||
public let hiddenSize: Int
|
public let hiddenSize: Int
|
||||||
public let numHiddenLayers: Int
|
public let numHiddenLayers: Int
|
||||||
@@ -42,13 +42,12 @@ public struct EmbeddingGemmaConfig: Codable {
|
|||||||
}
|
}
|
||||||
|
|
||||||
/// EmbeddingGemma - Google's 300M parameter embedding model
|
/// EmbeddingGemma - Google's 300M parameter embedding model
|
||||||
public final class EmbeddingGemmaModel {
|
public final class EmbeddingGemmaModel: @unchecked Sendable {
|
||||||
public let config: EmbeddingGemmaConfig
|
public let config: EmbeddingGemmaConfig
|
||||||
public let engine: MarkBaseEngine
|
public let engine: MarkBaseEngine
|
||||||
public let tokenizer: Tokenizer
|
public let tokenizer: Tokenizer
|
||||||
public let reader: SafeTensorsReader
|
public let reader: SafeTensorsReader
|
||||||
|
|
||||||
// GPU Buffers
|
|
||||||
public var embedTokens: MTLBuffer!
|
public var embedTokens: MTLBuffer!
|
||||||
public var finalNorm: MTLBuffer!
|
public var finalNorm: MTLBuffer!
|
||||||
public var layerNorms: [[MTLBuffer]] = []
|
public var layerNorms: [[MTLBuffer]] = []
|
||||||
@@ -67,18 +66,10 @@ public final class EmbeddingGemmaModel {
|
|||||||
self.config = try EmbeddingGemmaConfig.load(from: modelDir)
|
self.config = try EmbeddingGemmaConfig.load(from: modelDir)
|
||||||
self.tokenizer = try TokenizerFactory.load(modelDir: modelDir)
|
self.tokenizer = try TokenizerFactory.load(modelDir: modelDir)
|
||||||
self.reader = try SafeTensorsReader(path: modelDir + "/model.safetensors")
|
self.reader = try SafeTensorsReader(path: modelDir + "/model.safetensors")
|
||||||
|
|
||||||
try loadWeights()
|
try loadWeights()
|
||||||
print("✓ EmbeddingGemma loaded (\(config.numHiddenLayers) layers, hidden=\(config.hiddenSize))")
|
|
||||||
}
|
}
|
||||||
|
|
||||||
private func loadWeights() throws {
|
private func loadWeights() throws {
|
||||||
let hs = config.hiddenSize
|
|
||||||
let intermedi = config.intermediateSize
|
|
||||||
let nKV = config.numKeyValueHeads
|
|
||||||
let hDim = config.headDim
|
|
||||||
|
|
||||||
// Embedding table [vocab, hidden]
|
|
||||||
let embedData = try readTensor("embed_tokens.weight")
|
let embedData = try readTensor("embed_tokens.weight")
|
||||||
embedTokens = engine.device.makeBuffer(bytes: embedData, length: embedData.count * 4)!
|
embedTokens = engine.device.makeBuffer(bytes: embedData, length: embedData.count * 4)!
|
||||||
|
|
||||||
@@ -90,45 +81,53 @@ public final class EmbeddingGemmaModel {
|
|||||||
try loadBuffer("\(p).post_attention_layernorm.weight"),
|
try loadBuffer("\(p).post_attention_layernorm.weight"),
|
||||||
try loadBuffer("\(p).post_feedforward_layernorm.weight"),
|
try loadBuffer("\(p).post_feedforward_layernorm.weight"),
|
||||||
])
|
])
|
||||||
qProjs.append(try loadBuffer("\(p).self_attn.q_proj.weight")) // [hs, hs]
|
qProjs.append(try loadAndTranspose("\(p).self_attn.q_proj.weight", rows: config.hiddenSize, cols: config.hiddenSize))
|
||||||
kProjs.append(try loadBuffer("\(p).self_attn.k_proj.weight")) // [nKV*hDim, hs]
|
kProjs.append(try loadAndTranspose("\(p).self_attn.k_proj.weight", rows: config.numKeyValueHeads * config.headDim, cols: config.hiddenSize))
|
||||||
vProjs.append(try loadBuffer("\(p).self_attn.v_proj.weight")) // [nKV*hDim, hs]
|
vProjs.append(try loadAndTranspose("\(p).self_attn.v_proj.weight", rows: config.numKeyValueHeads * config.headDim, cols: config.hiddenSize))
|
||||||
oProjs.append(try loadBuffer("\(p).self_attn.o_proj.weight")) // [hs, nH*hDim]
|
oProjs.append(try loadAndTranspose("\(p).self_attn.o_proj.weight", rows: config.numAttentionHeads * config.headDim, cols: config.hiddenSize))
|
||||||
qNorms.append(try loadBuffer("\(p).self_attn.q_norm.weight")) // [hDim]
|
qNorms.append(try loadBuffer("\(p).self_attn.q_norm.weight"))
|
||||||
kNorms.append(try loadBuffer("\(p).self_attn.k_norm.weight")) // [hDim]
|
kNorms.append(try loadBuffer("\(p).self_attn.k_norm.weight"))
|
||||||
gateProjs.append(try loadBuffer("\(p).mlp.gate_proj.weight")) // [intermedi, hs]
|
gateProjs.append(try loadAndTranspose("\(p).mlp.gate_proj.weight", rows: config.intermediateSize, cols: config.hiddenSize))
|
||||||
upProjs.append(try loadBuffer("\(p).mlp.up_proj.weight")) // [intermedi, hs]
|
upProjs.append(try loadAndTranspose("\(p).mlp.up_proj.weight", rows: config.intermediateSize, cols: config.hiddenSize))
|
||||||
downProjs.append(try loadBuffer("\(p).mlp.down_proj.weight")) // [hs, intermedi]
|
downProjs.append(try loadAndTranspose("\(p).mlp.down_proj.weight", rows: config.hiddenSize, cols: config.intermediateSize))
|
||||||
}
|
}
|
||||||
|
|
||||||
let fnData = try readTensor("norm.weight")
|
let fnData = try readTensor("norm.weight")
|
||||||
finalNorm = engine.device.makeBuffer(bytes: fnData, length: fnData.count * 4)!
|
finalNorm = engine.device.makeBuffer(bytes: fnData, length: fnData.count * 4)!
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Generate embedding for text
|
private func loadAndTranspose(_ name: String, rows: Int, cols: Int) throws -> MTLBuffer {
|
||||||
|
let data = try readTensor(name)
|
||||||
|
var transposed = [Float](repeating: 0, count: data.count)
|
||||||
|
for r in 0..<rows {
|
||||||
|
for c in 0..<cols {
|
||||||
|
transposed[c * rows + r] = data[r * cols + c]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return engine.device.makeBuffer(bytes: transposed, length: transposed.count * 4)!
|
||||||
|
}
|
||||||
|
|
||||||
public func embed(text: String, maxLen: Int = 2048) throws -> [Float] {
|
public func embed(text: String, maxLen: Int = 2048) throws -> [Float] {
|
||||||
var tokens = tokenizer.encode(text: text)
|
var tokens = tokenizer.encode(text: text)
|
||||||
if tokens.count > maxLen { tokens = Array(tokens.prefix(maxLen)) }
|
if tokens.count > maxLen { tokens = Array(tokens.prefix(maxLen)) }
|
||||||
guard !tokens.isEmpty else { return [] }
|
guard !tokens.isEmpty else { return [] }
|
||||||
|
|
||||||
let seqLen = tokens.count, hs = config.hiddenSize
|
let seqLen = tokens.count, hs = config.hiddenSize
|
||||||
|
let cmdBuf = engine.commandQueue.makeCommandBuffer()!
|
||||||
|
|
||||||
// Embedding lookup
|
let inputBuf = try lookupEmbeddings(tokens: tokens, cmdBuf: cmdBuf)
|
||||||
let inputBuf = try lookupEmbeddings(tokens: tokens)
|
|
||||||
|
|
||||||
// Forward through layers
|
|
||||||
var hidden = inputBuf
|
var hidden = inputBuf
|
||||||
for idx in 0..<config.numHiddenLayers {
|
for idx in 0..<config.numHiddenLayers {
|
||||||
hidden = try forwardLayer(hidden: hidden, layerIdx: idx, seqLen: seqLen)
|
hidden = try forwardLayer(hidden: hidden, layerIdx: idx, seqLen: seqLen, cmdBuf: cmdBuf)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Final norm
|
let output = try applyRmsNorm(input: hidden, weight: finalNorm, count: seqLen * hs, cmdBuf: cmdBuf)
|
||||||
let output = try applyRmsNorm(input: hidden, weight: finalNorm, count: seqLen * hs)
|
|
||||||
|
cmdBuf.commit()
|
||||||
|
cmdBuf.waitUntilCompleted()
|
||||||
|
|
||||||
// Readback
|
|
||||||
let data = engine.readFloats(from: output, count: seqLen * hs)
|
let data = engine.readFloats(from: output, count: seqLen * hs)
|
||||||
|
|
||||||
// Mean pool + L2 normalize
|
|
||||||
var embedding = [Float](repeating: 0, count: hs)
|
var embedding = [Float](repeating: 0, count: hs)
|
||||||
for i in 0..<seqLen {
|
for i in 0..<seqLen {
|
||||||
let start = i * hs
|
let start = i * hs
|
||||||
@@ -145,20 +144,13 @@ public final class EmbeddingGemmaModel {
|
|||||||
return embedding
|
return embedding
|
||||||
}
|
}
|
||||||
|
|
||||||
// MARK: - Helpers
|
|
||||||
|
|
||||||
private func readTensor(_ name: String) throws -> [Float] {
|
private func readTensor(_ name: String) throws -> [Float] {
|
||||||
guard let desc = reader.tensor(named: name) else {
|
guard let desc = reader.tensor(named: name) else { throw WeightError.tensorNotFound(name) }
|
||||||
throw WeightError.tensorNotFound(name)
|
|
||||||
}
|
|
||||||
let data = try reader.read(tensor: desc)
|
let data = try reader.read(tensor: desc)
|
||||||
switch desc.dtype {
|
switch desc.dtype {
|
||||||
case .f32:
|
case .f32: return data.withUnsafeBytes { Array(UnsafeBufferPointer(start: $0.baseAddress?.assumingMemoryBound(to: Float.self), count: data.count/4)) }
|
||||||
return data.withUnsafeBytes { Array(UnsafeBufferPointer(start: $0.baseAddress?.assumingMemoryBound(to: Float.self), count: data.count/4)) }
|
case .bf16: return try SafeTensorsReader.bf16ToFloat32(data)
|
||||||
case .bf16:
|
default: throw WeightError.unsupportedDtype(desc.dtype.rawValue)
|
||||||
return try SafeTensorsReader.bf16ToFloat32(data)
|
|
||||||
default:
|
|
||||||
throw WeightError.unsupportedDtype(desc.dtype.rawValue)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -167,31 +159,25 @@ public final class EmbeddingGemmaModel {
|
|||||||
return engine.device.makeBuffer(bytes: data, length: data.count * 4)!
|
return engine.device.makeBuffer(bytes: data, length: data.count * 4)!
|
||||||
}
|
}
|
||||||
|
|
||||||
private func lookupEmbeddings(tokens: [Int]) throws -> MTLBuffer {
|
private func lookupEmbeddings(tokens: [Int], cmdBuf: MTLCommandBuffer) throws -> MTLBuffer {
|
||||||
let seqLen = tokens.count, hs = config.hiddenSize
|
let seqLen = tokens.count, hs = config.hiddenSize
|
||||||
let buf = engine.device.makeBuffer(length: seqLen * hs * 4)!
|
// CPU-based embedding lookup
|
||||||
let cmdBuf = engine.commandQueue.makeCommandBuffer()!
|
let embedPtr = embedTokens.contents().assumingMemoryBound(to: Float.self)
|
||||||
let enc = cmdBuf.makeComputeCommandEncoder()!
|
var embedData = [Float](repeating: 0, count: seqLen * hs)
|
||||||
let pso = try engine.pipeline(named: "lookup_embeddings")
|
for (i, token) in tokens.enumerated() {
|
||||||
enc.setComputePipelineState(pso)
|
let dstStart = i * hs
|
||||||
enc.setBuffer(embedTokens, offset: 0, index: 0)
|
let srcStart = token * hs
|
||||||
enc.setBytes(tokens, length: seqLen * MemoryLayout<Int>.size, index: 1)
|
for j in 0..<hs {
|
||||||
enc.setBuffer(buf, offset: 0, index: 2)
|
embedData[dstStart + j] = embedPtr[srcStart + j]
|
||||||
var h = UInt32(hs), s = UInt32(seqLen), v = UInt32(config.vocabSize)
|
}
|
||||||
enc.setBytes(&h, length: 4, index: 3)
|
}
|
||||||
enc.setBytes(&s, length: 4, index: 4)
|
return engine.device.makeBuffer(bytes: embedData, length: embedData.count * 4)!
|
||||||
enc.setBytes(&v, length: 4, index: 5)
|
|
||||||
enc.dispatchThreads(MTLSize(width: seqLen, height: 1, depth: 1),
|
|
||||||
threadsPerThreadgroup: MTLSize(width: min(256, seqLen), height: 1, depth: 1))
|
|
||||||
enc.endEncoding()
|
|
||||||
cmdBuf.commit(); cmdBuf.waitUntilCompleted()
|
|
||||||
return buf
|
|
||||||
}
|
}
|
||||||
|
|
||||||
private func applyRmsNorm(input: MTLBuffer, weight: MTLBuffer, count: Int) throws -> MTLBuffer {
|
private func applyRmsNorm(input: MTLBuffer, weight: MTLBuffer, count: Int, cmdBuf: MTLCommandBuffer) throws -> MTLBuffer {
|
||||||
let output = engine.device.makeBuffer(length: count * 4)!
|
let output = engine.device.makeBuffer(length: count * 4)!
|
||||||
let cmdBuf = engine.commandQueue.makeCommandBuffer()!
|
|
||||||
let enc = cmdBuf.makeComputeCommandEncoder()!
|
let enc = cmdBuf.makeComputeCommandEncoder()!
|
||||||
|
defer { enc.endEncoding() }
|
||||||
let pso = try engine.pipeline(named: "rms_norm")
|
let pso = try engine.pipeline(named: "rms_norm")
|
||||||
enc.setComputePipelineState(pso)
|
enc.setComputePipelineState(pso)
|
||||||
enc.setBuffer(input, offset: 0, index: 0)
|
enc.setBuffer(input, offset: 0, index: 0)
|
||||||
@@ -202,84 +188,9 @@ public final class EmbeddingGemmaModel {
|
|||||||
enc.setBytes(&e, length: 4, index: 4)
|
enc.setBytes(&e, length: 4, index: 4)
|
||||||
enc.dispatchThreads(MTLSize(width: count, height: 1, depth: 1),
|
enc.dispatchThreads(MTLSize(width: count, height: 1, depth: 1),
|
||||||
threadsPerThreadgroup: MTLSize(width: min(256, count), height: 1, depth: 1))
|
threadsPerThreadgroup: MTLSize(width: min(256, count), height: 1, depth: 1))
|
||||||
enc.endEncoding()
|
|
||||||
cmdBuf.commit(); cmdBuf.waitUntilCompleted()
|
|
||||||
return output
|
return output
|
||||||
}
|
}
|
||||||
|
|
||||||
private func forwardLayer(hidden: MTLBuffer, layerIdx: Int, seqLen: Int) throws -> MTLBuffer {
|
|
||||||
let hs = config.hiddenSize, device = engine.device
|
|
||||||
let hDim = config.headDim, nH = config.numAttentionHeads, nKV = config.numKeyValueHeads
|
|
||||||
let intermedi = config.intermediateSize
|
|
||||||
|
|
||||||
let cmdBuf = engine.commandQueue.makeCommandBuffer()!
|
|
||||||
|
|
||||||
// Residual
|
|
||||||
let resid = device.makeBuffer(length: seqLen * hs * 4)!
|
|
||||||
let blit = cmdBuf.makeBlitCommandEncoder()!
|
|
||||||
blit.copy(from: hidden, sourceOffset: 0, to: resid, destinationOffset: 0, size: seqLen * hs * 4)
|
|
||||||
blit.endEncoding()
|
|
||||||
|
|
||||||
// Input norm
|
|
||||||
let h1 = try applyRmsNorm(input: hidden, weight: layerNorms[layerIdx][0], count: seqLen * hs)
|
|
||||||
|
|
||||||
// Q, K, V projections (using optimized matmul)
|
|
||||||
let qBuf = device.makeBuffer(length: seqLen * nH * hDim * 4)!
|
|
||||||
let kBuf = device.makeBuffer(length: seqLen * nKV * hDim * 4)!
|
|
||||||
let vBuf = device.makeBuffer(length: seqLen * nKV * hDim * 4)!
|
|
||||||
try matmulSeq(input: h1, weight: qProjs[layerIdx], output: qBuf, m: seqLen, k: hs, n: nH * hDim, cmdBuf: cmdBuf)
|
|
||||||
try matmulSeq(input: h1, weight: kProjs[layerIdx], output: kBuf, m: seqLen, k: hs, n: nKV * hDim, cmdBuf: cmdBuf)
|
|
||||||
try matmulSeq(input: h1, weight: vProjs[layerIdx], output: vBuf, m: seqLen, k: hs, n: nKV * hDim, cmdBuf: cmdBuf)
|
|
||||||
|
|
||||||
// RoPE
|
|
||||||
try applyRoPE(q: qBuf, k: kBuf, seqLen: seqLen, headDim: hDim, numHeads: nH, numKVHeads: nKV, cmdBuf: cmdBuf)
|
|
||||||
|
|
||||||
// Q/K Norm
|
|
||||||
try applyQKNorm(q: qBuf, k: kBuf, qNorm: qNorms[layerIdx], kNorm: kNorms[layerIdx], seqLen: seqLen, headDim: hDim, numHeads: nH, numKVHeads: nKV, cmdBuf: cmdBuf)
|
|
||||||
|
|
||||||
// Bidirectional sliding window attention
|
|
||||||
let attnOut = device.makeBuffer(length: seqLen * nH * hDim * 4)!
|
|
||||||
try bidirectionalAttention(q: qBuf, k: kBuf, v: vBuf, output: attnOut, seqLen: seqLen, cmdBuf: cmdBuf)
|
|
||||||
|
|
||||||
// O projection
|
|
||||||
let h2 = device.makeBuffer(length: seqLen * hs * 4)!
|
|
||||||
try matmulSeq(input: attnOut, weight: oProjs[layerIdx], output: h2, m: seqLen, k: nH * hDim, n: hs, cmdBuf: cmdBuf)
|
|
||||||
|
|
||||||
// Post-attn norm
|
|
||||||
let h2n = try applyRmsNorm(input: h2, weight: layerNorms[layerIdx][2], count: seqLen * hs)
|
|
||||||
|
|
||||||
// Add residual: hidden = resid + h2n
|
|
||||||
try eltwiseAdd(a: resid, b: h2n, output: hidden, count: seqLen * hs, cmdBuf: cmdBuf)
|
|
||||||
|
|
||||||
// Pre-FF norm
|
|
||||||
let h3 = try applyRmsNorm(input: hidden, weight: layerNorms[layerIdx][1], count: seqLen * hs)
|
|
||||||
|
|
||||||
// MLP: gate, up
|
|
||||||
let gate = device.makeBuffer(length: seqLen * intermedi * 4)!
|
|
||||||
let up = device.makeBuffer(length: seqLen * intermedi * 4)!
|
|
||||||
try matmulSeq(input: h3, weight: gateProjs[layerIdx], output: gate, m: seqLen, k: hs, n: intermedi, cmdBuf: cmdBuf)
|
|
||||||
try matmulSeq(input: h3, weight: upProjs[layerIdx], output: up, m: seqLen, k: hs, n: intermedi, cmdBuf: cmdBuf)
|
|
||||||
|
|
||||||
// GELU(gate) * up
|
|
||||||
let gated = device.makeBuffer(length: seqLen * intermedi * 4)!
|
|
||||||
try geluMul(gate: gate, up: up, output: gated, count: seqLen * intermedi, cmdBuf: cmdBuf)
|
|
||||||
|
|
||||||
// Down projection
|
|
||||||
let h4 = device.makeBuffer(length: seqLen * hs * 4)!
|
|
||||||
try matmulSeq(input: gated, weight: downProjs[layerIdx], output: h4, m: seqLen, k: intermedi, n: hs, cmdBuf: cmdBuf)
|
|
||||||
|
|
||||||
// Post-FF norm
|
|
||||||
let h4n = try applyRmsNorm(input: h4, weight: layerNorms[layerIdx][3], count: seqLen * hs)
|
|
||||||
|
|
||||||
// Add residual: hidden = hidden + h4n
|
|
||||||
try eltwiseAdd(a: hidden, b: h4n, output: hidden, count: seqLen * hs, cmdBuf: cmdBuf)
|
|
||||||
|
|
||||||
cmdBuf.commit(); cmdBuf.waitUntilCompleted()
|
|
||||||
return hidden
|
|
||||||
}
|
|
||||||
|
|
||||||
// MARK: - Metal Kernels
|
|
||||||
|
|
||||||
private func matmulSeq(input: MTLBuffer, weight: MTLBuffer, output: MTLBuffer, m: Int, k: Int, n: Int, cmdBuf: MTLCommandBuffer) throws {
|
private func matmulSeq(input: MTLBuffer, weight: MTLBuffer, output: MTLBuffer, m: Int, k: Int, n: Int, cmdBuf: MTLCommandBuffer) throws {
|
||||||
let enc = cmdBuf.makeComputeCommandEncoder()!
|
let enc = cmdBuf.makeComputeCommandEncoder()!
|
||||||
let pso = try engine.pipeline(named: "matmul_f32")
|
let pso = try engine.pipeline(named: "matmul_f32")
|
||||||
@@ -291,13 +202,43 @@ public final class EmbeddingGemmaModel {
|
|||||||
enc.setBytes(&mm, length: 4, index: 3)
|
enc.setBytes(&mm, length: 4, index: 3)
|
||||||
enc.setBytes(&kk, length: 4, index: 4)
|
enc.setBytes(&kk, length: 4, index: 4)
|
||||||
enc.setBytes(&nn, length: 4, index: 5)
|
enc.setBytes(&nn, length: 4, index: 5)
|
||||||
enc.dispatchThreads(MTLSize(width: m * n, height: 1, depth: 1),
|
let total = m * n
|
||||||
threadsPerThreadgroup: MTLSize(width: min(256, m * n), height: 1, depth: 1))
|
enc.dispatchThreads(MTLSize(width: total, height: 1, depth: 1),
|
||||||
|
threadsPerThreadgroup: MTLSize(width: min(256, total), height: 1, depth: 1))
|
||||||
enc.endEncoding()
|
enc.endEncoding()
|
||||||
}
|
}
|
||||||
|
|
||||||
private func applyRoPE(q: MTLBuffer, k: MTLBuffer, seqLen: Int, headDim: Int, numHeads: Int, numKVHeads: Int, cmdBuf: MTLCommandBuffer) throws {
|
private func eltwiseAdd(a: MTLBuffer, b: MTLBuffer, output: MTLBuffer, count: Int, cmdBuf: MTLCommandBuffer) throws {
|
||||||
let enc = cmdBuf.makeComputeCommandEncoder()!
|
let enc = cmdBuf.makeComputeCommandEncoder()!
|
||||||
|
defer { enc.endEncoding() }
|
||||||
|
let pso = try engine.pipeline(named: "eltwise_add")
|
||||||
|
enc.setComputePipelineState(pso)
|
||||||
|
enc.setBuffer(a, offset: 0, index: 0)
|
||||||
|
enc.setBuffer(b, offset: 0, index: 1)
|
||||||
|
enc.setBuffer(output, offset: 0, index: 2)
|
||||||
|
var c = UInt32(count)
|
||||||
|
enc.setBytes(&c, length: 4, index: 3)
|
||||||
|
enc.dispatchThreads(MTLSize(width: count, height: 1, depth: 1),
|
||||||
|
threadsPerThreadgroup: MTLSize(width: min(256, count), height: 1, depth: 1))
|
||||||
|
}
|
||||||
|
|
||||||
|
private func geluMul(gate: MTLBuffer, up: MTLBuffer, output: MTLBuffer, count: Int, cmdBuf: MTLCommandBuffer) throws {
|
||||||
|
let enc = cmdBuf.makeComputeCommandEncoder()!
|
||||||
|
defer { enc.endEncoding() }
|
||||||
|
let pso = try engine.pipeline(named: "gelu_mul_kernel")
|
||||||
|
enc.setComputePipelineState(pso)
|
||||||
|
enc.setBuffer(gate, offset: 0, index: 0)
|
||||||
|
enc.setBuffer(up, offset: 0, index: 1)
|
||||||
|
enc.setBuffer(output, offset: 0, index: 2)
|
||||||
|
var c = UInt32(count)
|
||||||
|
enc.setBytes(&c, length: 4, index: 3)
|
||||||
|
enc.dispatchThreads(MTLSize(width: count, height: 1, depth: 1),
|
||||||
|
threadsPerThreadgroup: MTLSize(width: min(256, count), height: 1, depth: 1))
|
||||||
|
}
|
||||||
|
|
||||||
|
private func applyRoPE(q: MTLBuffer, k: MTLBuffer, seqLen: Int, headDim: Int, numHeads: Int, cmdBuf: MTLCommandBuffer) throws {
|
||||||
|
let enc = cmdBuf.makeComputeCommandEncoder()!
|
||||||
|
defer { enc.endEncoding() }
|
||||||
let pso = try engine.pipeline(named: "apply_rope")
|
let pso = try engine.pipeline(named: "apply_rope")
|
||||||
enc.setComputePipelineState(pso)
|
enc.setComputePipelineState(pso)
|
||||||
enc.setBuffer(q, offset: 0, index: 0)
|
enc.setBuffer(q, offset: 0, index: 0)
|
||||||
@@ -310,16 +251,11 @@ public final class EmbeddingGemmaModel {
|
|||||||
enc.setBytes(&rt, length: 4, index: 5)
|
enc.setBytes(&rt, length: 4, index: 5)
|
||||||
enc.dispatchThreads(MTLSize(width: numHeads * headDim / 2, height: 1, depth: 1),
|
enc.dispatchThreads(MTLSize(width: numHeads * headDim / 2, height: 1, depth: 1),
|
||||||
threadsPerThreadgroup: MTLSize(width: min(256, numHeads * headDim / 2), height: 1, depth: 1))
|
threadsPerThreadgroup: MTLSize(width: min(256, numHeads * headDim / 2), height: 1, depth: 1))
|
||||||
enc.endEncoding()
|
|
||||||
}
|
|
||||||
|
|
||||||
private func applyQKNorm(q: MTLBuffer, k: MTLBuffer, qNorm: MTLBuffer, kNorm: MTLBuffer, seqLen: Int, headDim: Int, numHeads: Int, numKVHeads: Int, cmdBuf: MTLCommandBuffer) throws {
|
|
||||||
// Apply RMSNorm per head
|
|
||||||
// TODO: Implement per-head normalization
|
|
||||||
}
|
}
|
||||||
|
|
||||||
private func bidirectionalAttention(q: MTLBuffer, k: MTLBuffer, v: MTLBuffer, output: MTLBuffer, seqLen: Int, cmdBuf: MTLCommandBuffer) throws {
|
private func bidirectionalAttention(q: MTLBuffer, k: MTLBuffer, v: MTLBuffer, output: MTLBuffer, seqLen: Int, cmdBuf: MTLCommandBuffer) throws {
|
||||||
let enc = cmdBuf.makeComputeCommandEncoder()!
|
let enc = cmdBuf.makeComputeCommandEncoder()!
|
||||||
|
defer { enc.endEncoding() }
|
||||||
let pso = try engine.pipeline(named: "bidirectional_sliding_attn")
|
let pso = try engine.pipeline(named: "bidirectional_sliding_attn")
|
||||||
enc.setComputePipelineState(pso)
|
enc.setComputePipelineState(pso)
|
||||||
enc.setBuffer(q, offset: 0, index: 0)
|
enc.setBuffer(q, offset: 0, index: 0)
|
||||||
@@ -335,38 +271,58 @@ public final class EmbeddingGemmaModel {
|
|||||||
enc.setBytes(&nkv, length: 4, index: 7)
|
enc.setBytes(&nkv, length: 4, index: 7)
|
||||||
enc.setBytes(&sw, length: 4, index: 8)
|
enc.setBytes(&sw, length: 4, index: 8)
|
||||||
enc.setBytes(&scale, length: 4, index: 9)
|
enc.setBytes(&scale, length: 4, index: 9)
|
||||||
let tgMem = config.slidingWindow * 4 // shared memory for scores
|
let tgMem = config.slidingWindow * 4
|
||||||
enc.setThreadgroupMemoryLength(tgMem, index: 0)
|
enc.setThreadgroupMemoryLength(tgMem, index: 0)
|
||||||
enc.dispatchThreads(MTLSize(width: seqLen * config.numAttentionHeads, height: 1, depth: 1),
|
enc.dispatchThreads(MTLSize(width: seqLen * config.numAttentionHeads, height: 1, depth: 1),
|
||||||
threadsPerThreadgroup: MTLSize(width: min(256, seqLen * config.numAttentionHeads), height: 1, depth: 1))
|
threadsPerThreadgroup: MTLSize(width: min(256, seqLen * config.numAttentionHeads), height: 1, depth: 1))
|
||||||
enc.endEncoding()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
private func eltwiseAdd(a: MTLBuffer, b: MTLBuffer, output: MTLBuffer, count: Int, cmdBuf: MTLCommandBuffer) throws {
|
private func forwardLayer(hidden: MTLBuffer, layerIdx: Int, seqLen: Int, cmdBuf: MTLCommandBuffer) throws -> MTLBuffer {
|
||||||
let enc = cmdBuf.makeComputeCommandEncoder()!
|
let hs = config.hiddenSize, device = engine.device
|
||||||
let pso = try engine.pipeline(named: "eltwise_add")
|
let hDim = config.headDim, nH = config.numAttentionHeads, nKV = config.numKeyValueHeads
|
||||||
enc.setComputePipelineState(pso)
|
let intermedi = config.intermediateSize
|
||||||
enc.setBuffer(a, offset: 0, index: 0)
|
|
||||||
enc.setBuffer(b, offset: 0, index: 1)
|
|
||||||
enc.setBuffer(output, offset: 0, index: 2)
|
|
||||||
var c = UInt32(count)
|
|
||||||
enc.setBytes(&c, length: 4, index: 3)
|
|
||||||
enc.dispatchThreads(MTLSize(width: count, height: 1, depth: 1),
|
|
||||||
threadsPerThreadgroup: MTLSize(width: min(256, count), height: 1, depth: 1))
|
|
||||||
enc.endEncoding()
|
|
||||||
}
|
|
||||||
|
|
||||||
private func geluMul(gate: MTLBuffer, up: MTLBuffer, output: MTLBuffer, count: Int, cmdBuf: MTLCommandBuffer) throws {
|
let resid = device.makeBuffer(length: seqLen * hs * 4)!
|
||||||
let enc = cmdBuf.makeComputeCommandEncoder()!
|
let blit = cmdBuf.makeBlitCommandEncoder()!
|
||||||
let pso = try engine.pipeline(named: "gelu_mul_kernel")
|
blit.copy(from: hidden, sourceOffset: 0, to: resid, destinationOffset: 0, size: seqLen * hs * 4)
|
||||||
enc.setComputePipelineState(pso)
|
blit.endEncoding()
|
||||||
enc.setBuffer(gate, offset: 0, index: 0)
|
|
||||||
enc.setBuffer(up, offset: 0, index: 1)
|
let h1 = try applyRmsNorm(input: hidden, weight: layerNorms[layerIdx][0], count: seqLen * hs, cmdBuf: cmdBuf)
|
||||||
enc.setBuffer(output, offset: 0, index: 2)
|
|
||||||
var c = UInt32(count)
|
let qBuf = device.makeBuffer(length: seqLen * nH * hDim * 4)!
|
||||||
enc.setBytes(&c, length: 4, index: 3)
|
let kBuf = device.makeBuffer(length: seqLen * nKV * hDim * 4)!
|
||||||
enc.dispatchThreads(MTLSize(width: count, height: 1, depth: 1),
|
let vBuf = device.makeBuffer(length: seqLen * nKV * hDim * 4)!
|
||||||
threadsPerThreadgroup: MTLSize(width: min(256, count), height: 1, depth: 1))
|
try matmulSeq(input: h1, weight: qProjs[layerIdx], output: qBuf, m: seqLen, k: hs, n: nH * hDim, cmdBuf: cmdBuf)
|
||||||
enc.endEncoding()
|
try matmulSeq(input: h1, weight: kProjs[layerIdx], output: kBuf, m: seqLen, k: hs, n: nKV * hDim, cmdBuf: cmdBuf)
|
||||||
|
try matmulSeq(input: h1, weight: vProjs[layerIdx], output: vBuf, m: seqLen, k: hs, n: nKV * hDim, cmdBuf: cmdBuf)
|
||||||
|
|
||||||
|
try applyRoPE(q: qBuf, k: kBuf, seqLen: seqLen, headDim: hDim, numHeads: nH, cmdBuf: cmdBuf)
|
||||||
|
|
||||||
|
let attnOut = device.makeBuffer(length: seqLen * nH * hDim * 4)!
|
||||||
|
try bidirectionalAttention(q: qBuf, k: kBuf, v: vBuf, output: attnOut, seqLen: seqLen, cmdBuf: cmdBuf)
|
||||||
|
|
||||||
|
let h2 = device.makeBuffer(length: seqLen * hs * 4)!
|
||||||
|
try matmulSeq(input: attnOut, weight: oProjs[layerIdx], output: h2, m: seqLen, k: nH * hDim, n: hs, cmdBuf: cmdBuf)
|
||||||
|
|
||||||
|
let h2n = try applyRmsNorm(input: h2, weight: layerNorms[layerIdx][2], count: seqLen * hs, cmdBuf: cmdBuf)
|
||||||
|
try eltwiseAdd(a: resid, b: h2n, output: hidden, count: seqLen * hs, cmdBuf: cmdBuf)
|
||||||
|
|
||||||
|
let h3 = try applyRmsNorm(input: hidden, weight: layerNorms[layerIdx][1], count: seqLen * hs, cmdBuf: cmdBuf)
|
||||||
|
|
||||||
|
let gate = device.makeBuffer(length: seqLen * intermedi * 4)!
|
||||||
|
let up = device.makeBuffer(length: seqLen * intermedi * 4)!
|
||||||
|
try matmulSeq(input: h3, weight: gateProjs[layerIdx], output: gate, m: seqLen, k: hs, n: intermedi, cmdBuf: cmdBuf)
|
||||||
|
try matmulSeq(input: h3, weight: upProjs[layerIdx], output: up, m: seqLen, k: hs, n: intermedi, cmdBuf: cmdBuf)
|
||||||
|
|
||||||
|
let gated = device.makeBuffer(length: seqLen * intermedi * 4)!
|
||||||
|
try geluMul(gate: gate, up: up, output: gated, count: seqLen * intermedi, cmdBuf: cmdBuf)
|
||||||
|
|
||||||
|
let h4 = device.makeBuffer(length: seqLen * hs * 4)!
|
||||||
|
try matmulSeq(input: gated, weight: downProjs[layerIdx], output: h4, m: seqLen, k: intermedi, n: hs, cmdBuf: cmdBuf)
|
||||||
|
|
||||||
|
let h4n = try applyRmsNorm(input: h4, weight: layerNorms[layerIdx][3], count: seqLen * hs, cmdBuf: cmdBuf)
|
||||||
|
try eltwiseAdd(a: hidden, b: h4n, output: hidden, count: seqLen * hs, cmdBuf: cmdBuf)
|
||||||
|
|
||||||
|
return hidden
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1637,10 +1637,11 @@ kernel void matmul_f32(
|
|||||||
uint id [[thread_position_in_grid]]
|
uint id [[thread_position_in_grid]]
|
||||||
) {
|
) {
|
||||||
// Each thread computes one output element
|
// Each thread computes one output element
|
||||||
uint row = 0; // For single token, M=1
|
uint total = M * N;
|
||||||
uint col = id;
|
if (id >= total) return;
|
||||||
|
|
||||||
if (col >= N) return;
|
uint row = id / N;
|
||||||
|
uint col = id % N;
|
||||||
|
|
||||||
float sum = 0.0;
|
float sum = 0.0;
|
||||||
for (uint k = 0; k < K; k++) {
|
for (uint k = 0; k < K; k++) {
|
||||||
|
|||||||
@@ -109,6 +109,21 @@ public enum MetalKernels {
|
|||||||
.replacingOccurrences(of: "using namespace metal;\n", with: "")
|
.replacingOccurrences(of: "using namespace metal;\n", with: "")
|
||||||
result += "\n" + fusedStripped
|
result += "\n" + fusedStripped
|
||||||
|
|
||||||
|
// Strip preamble from embedding kernels source
|
||||||
|
let embStripped = embeddingKernelsSource
|
||||||
|
.replacingOccurrences(of: "#include <metal_stdlib>\n", with: "")
|
||||||
|
.replacingOccurrences(of: "using namespace metal;\n", with: "")
|
||||||
|
result += "\n" + embStripped
|
||||||
|
|
||||||
return result
|
return result
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Embedding kernel source for EmbeddingGemma.
|
||||||
|
/// Includes RoPE, bidirectional sliding window attention, Q/K norm, and GELU.
|
||||||
|
public static var embeddingKernelsSource: String {
|
||||||
|
let url = URL(fileURLWithPath: #filePath)
|
||||||
|
.deletingLastPathComponent()
|
||||||
|
.appendingPathComponent("Metal/EmbeddingKernels.metal")
|
||||||
|
return try! String(contentsOf: url, encoding: .utf8)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
@@ -90,15 +90,35 @@ public enum TokenizerError: Error, LocalizedError {
|
|||||||
public final class TokenizerFactory: @unchecked Sendable {
|
public final class TokenizerFactory: @unchecked Sendable {
|
||||||
/// Load tokenizer from model directory
|
/// Load tokenizer from model directory
|
||||||
public static func load(modelDir: String) throws -> Tokenizer {
|
public static func load(modelDir: String) throws -> Tokenizer {
|
||||||
|
// Check tokenizer_config.json for tokenizer_class
|
||||||
|
let configPath = modelDir + "/tokenizer_config.json"
|
||||||
|
var tokenizerClass: String? = nil
|
||||||
|
if FileManager.default.fileExists(atPath: configPath),
|
||||||
|
let configData = try? Data(contentsOf: URL(fileURLWithPath: configPath)),
|
||||||
|
let config = try? JSONSerialization.jsonObject(with: configData) as? [String: Any] {
|
||||||
|
tokenizerClass = config["tokenizer_class"] as? String
|
||||||
|
}
|
||||||
|
|
||||||
|
// For GemmaTokenizer or SentencePiece-based tokenizers, prefer .model file
|
||||||
|
if tokenizerClass == "GemmaTokenizer" || tokenizerClass?.contains("SentencePiece") == true {
|
||||||
|
let modelPath = modelDir + "/tokenizer.model"
|
||||||
|
if FileManager.default.fileExists(atPath: modelPath) {
|
||||||
|
print(" Using SentencePieceTokenizer (tokenizer_class=\(tokenizerClass ?? "unknown"))")
|
||||||
|
return try SentencePieceTokenizer(modelPath: modelPath)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// Try tokenizer.json first (HuggingFace format)
|
// Try tokenizer.json first (HuggingFace format)
|
||||||
let tokenizerJsonPath = modelDir + "/tokenizer.json"
|
let tokenizerJsonPath = modelDir + "/tokenizer.json"
|
||||||
if FileManager.default.fileExists(atPath: tokenizerJsonPath) {
|
if FileManager.default.fileExists(atPath: tokenizerJsonPath) {
|
||||||
|
print(" Using BPETokenizer (tokenizer.json)")
|
||||||
return try BPETokenizer(jsonPath: tokenizerJsonPath)
|
return try BPETokenizer(jsonPath: tokenizerJsonPath)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Try .model file (SentencePiece format)
|
// Try .model file (SentencePiece format)
|
||||||
let modelPath = modelDir + "/tokenizer.model"
|
let modelPath = modelDir + "/tokenizer.model"
|
||||||
if FileManager.default.fileExists(atPath: modelPath) {
|
if FileManager.default.fileExists(atPath: modelPath) {
|
||||||
|
print(" Using SentencePieceTokenizer (tokenizer.model)")
|
||||||
return try SentencePieceTokenizer(modelPath: modelPath)
|
return try SentencePieceTokenizer(modelPath: modelPath)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,113 @@
|
|||||||
|
import Foundation
|
||||||
|
import MarkBase
|
||||||
|
import Hummingbird
|
||||||
|
|
||||||
|
struct EmbeddingServerApp {
|
||||||
|
static func main() async throws {
|
||||||
|
let args = CommandLine.arguments
|
||||||
|
let modelName = args.count > 1 ? args[1] : "embeddinggemma-300m"
|
||||||
|
let port = args.count > 2 ? Int(args[2]) ?? 8084 : 8084
|
||||||
|
|
||||||
|
let modelPath = NSString(string: "~/MarkBaseEngine/models/\(modelName)").expandingTildeInPath
|
||||||
|
|
||||||
|
print("═══════════════════════════════════════════════════════════════════")
|
||||||
|
print(" MarkBaseEngine Embedding Server")
|
||||||
|
print("═══════════════════════════════════════════════════════════════════")
|
||||||
|
print(" Model: \(modelName)")
|
||||||
|
print(" Port: \(port)")
|
||||||
|
print(" Path: \(modelPath)")
|
||||||
|
print("")
|
||||||
|
|
||||||
|
let engine = try MarkBaseEngine(autoCompile: true)
|
||||||
|
let embedModel = try EmbeddingGemmaModel(modelDir: modelPath, engine: engine)
|
||||||
|
let layers = embedModel.config.numHiddenLayers
|
||||||
|
let hiddenSize = embedModel.config.hiddenSize
|
||||||
|
|
||||||
|
print("✓ EmbeddingGemma loaded (\(layers) layers, hidden=\(hiddenSize))")
|
||||||
|
|
||||||
|
let router = Router()
|
||||||
|
|
||||||
|
router.get("/") { _, _ in
|
||||||
|
return """
|
||||||
|
{
|
||||||
|
"server": {
|
||||||
|
"name": "MarkBaseEngine Embedding",
|
||||||
|
"model": "embeddinggemma-300m",
|
||||||
|
"layers": \(layers),
|
||||||
|
"hidden_size": \(hiddenSize),
|
||||||
|
"output_dim": \(hiddenSize),
|
||||||
|
"framework": "Hummingbird 2.x + Metal GPU"
|
||||||
|
},
|
||||||
|
"endpoints": [
|
||||||
|
{"method": "POST", "path": "/v1/embeddings", "summary": "Generate text embeddings"}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
"""
|
||||||
|
}
|
||||||
|
|
||||||
|
router.get("/health") { _, _ in
|
||||||
|
return "{\"status\":\"healthy\",\"model\":\"embeddinggemma-300m\",\"layers\":\(layers),\"hidden_size\":\(hiddenSize)}"
|
||||||
|
}
|
||||||
|
|
||||||
|
router.post("/v1/embeddings") { request, _ in
|
||||||
|
let buffer = try await request.body.collect(upTo: .max)
|
||||||
|
let data = Data(buffer: buffer)
|
||||||
|
|
||||||
|
guard let json = try JSONSerialization.jsonObject(with: data) as? [String: Any],
|
||||||
|
let input = json["input"] else {
|
||||||
|
return "{\"error\":\"invalid request\",\"type\":\"invalid_request_error\",\"code\":400,\"message\":\"missing 'input' field\"}"
|
||||||
|
}
|
||||||
|
|
||||||
|
let modelId = (json["model"] as? String) ?? "embeddinggemma-300m"
|
||||||
|
let encodingFormat = (json["encoding_format"] as? String) ?? "float"
|
||||||
|
|
||||||
|
let inputs: [String]
|
||||||
|
if let str = input as? String { inputs = [str] }
|
||||||
|
else if let arr = input as? [String] { inputs = arr }
|
||||||
|
else {
|
||||||
|
return "{\"error\":\"invalid request\",\"type\":\"invalid_request_error\",\"code\":400,\"message\":\"'input' must be string or array of strings\"}"
|
||||||
|
}
|
||||||
|
|
||||||
|
var embeddings: [[String: Any]] = []
|
||||||
|
for (i, text) in inputs.enumerated() {
|
||||||
|
let t0 = Date()
|
||||||
|
let embedding = try embedModel.embed(text: text)
|
||||||
|
let duration = Date().timeIntervalSince(t0)
|
||||||
|
|
||||||
|
let embeddingData: [String: Any]
|
||||||
|
if encodingFormat == "base64" {
|
||||||
|
let base64 = embedding.withUnsafeBytes { Data($0).base64EncodedString() }
|
||||||
|
embeddingData = ["object": "embedding", "index": i, "embedding": base64, "usage_ms": Int(duration * 1000)]
|
||||||
|
} else {
|
||||||
|
embeddingData = ["object": "embedding", "index": i, "embedding": embedding, "usage_ms": Int(duration * 1000)]
|
||||||
|
}
|
||||||
|
embeddings.append(embeddingData)
|
||||||
|
}
|
||||||
|
|
||||||
|
let id = UUID().uuidString
|
||||||
|
let ts = Int(Date().timeIntervalSince1970)
|
||||||
|
let response: [String: Any] = [
|
||||||
|
"id": id, "object": "list", "created": ts, "model": modelId,
|
||||||
|
"data": embeddings,
|
||||||
|
"usage": ["prompt_tokens": inputs.reduce(0) { $0 + $1.components(separatedBy: .whitespaces).count }, "total_tokens": inputs.reduce(0) { $0 + $1.components(separatedBy: .whitespaces).count }]
|
||||||
|
]
|
||||||
|
|
||||||
|
let jsonData = try JSONSerialization.data(withJSONObject: response)
|
||||||
|
return String(data: jsonData, encoding: .utf8) ?? "{}"
|
||||||
|
}
|
||||||
|
|
||||||
|
let app = Application(
|
||||||
|
router: router,
|
||||||
|
configuration: .init(address: .hostname("0.0.0.0", port: port))
|
||||||
|
)
|
||||||
|
|
||||||
|
print("Server starting on port \(port)...")
|
||||||
|
print("Endpoints:")
|
||||||
|
print(" GET / - Info")
|
||||||
|
print(" GET /health - Health check")
|
||||||
|
print(" POST /v1/embeddings - Text embeddings")
|
||||||
|
print("")
|
||||||
|
|
||||||
|
try await app.run()
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -20,16 +20,36 @@ struct SimpleServerApp {
|
|||||||
print("")
|
print("")
|
||||||
|
|
||||||
let engine = try MarkBaseEngine(autoCompile: true)
|
let engine = try MarkBaseEngine(autoCompile: true)
|
||||||
let model = try E4BModel(modelDir: modelPath, engine: engine, maxContextLength: 512)
|
|
||||||
let tokenizer = try TokenizerFactory.load(modelDir: modelPath)
|
|
||||||
let generator = StreamingGenerator(model: model, tokenizer: tokenizer, engine: engine)
|
|
||||||
let embeddingModel = try TextEmbeddingModel(modelDir: modelPath, engine: engine, config: TextEmbeddingConfig())
|
|
||||||
|
|
||||||
print("✓ E4B loaded (\(model.numHiddenLayers) layers)")
|
// Detect model type
|
||||||
|
let isEmbeddingGemma = modelName.contains("embeddinggemma") || modelName.contains("embedding")
|
||||||
|
let isE4B = !isEmbeddingGemma
|
||||||
|
|
||||||
|
var e4bModel: E4BModel? = nil
|
||||||
|
var embedModel: EmbeddingGemmaModel? = nil
|
||||||
|
let tokenizer = try TokenizerFactory.load(modelDir: modelPath)
|
||||||
|
var generator: StreamingGenerator? = nil
|
||||||
|
|
||||||
|
if isE4B {
|
||||||
|
e4bModel = try E4BModel(modelDir: modelPath, engine: engine, maxContextLength: 512)
|
||||||
|
generator = StreamingGenerator(model: e4bModel!, tokenizer: tokenizer, engine: engine)
|
||||||
|
print("✓ E4B loaded (\(e4bModel!.numHiddenLayers) layers)")
|
||||||
|
} else {
|
||||||
|
embedModel = try EmbeddingGemmaModel(modelDir: modelPath, engine: engine)
|
||||||
|
print("✓ EmbeddingGemma loaded (\(embedModel!.config.numHiddenLayers) layers, hidden=\(embedModel!.config.hiddenSize))")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Embedding model (use EmbeddingGemma if available, fallback to E4B)
|
||||||
|
let textEmbedModel: TextEmbeddingModel?
|
||||||
|
if isE4B {
|
||||||
|
textEmbedModel = try TextEmbeddingModel(modelDir: modelPath, engine: engine, config: TextEmbeddingConfig())
|
||||||
|
} else {
|
||||||
|
textEmbedModel = nil
|
||||||
|
}
|
||||||
|
|
||||||
let router = Router()
|
let router = Router()
|
||||||
|
|
||||||
let layers = model.numHiddenLayers
|
let layers = isE4B ? e4bModel!.numHiddenLayers : embedModel!.config.numHiddenLayers
|
||||||
|
|
||||||
@Sendable func helpJSON() -> String {
|
@Sendable func helpJSON() -> String {
|
||||||
return """
|
return """
|
||||||
@@ -184,7 +204,8 @@ struct SimpleServerApp {
|
|||||||
}
|
}
|
||||||
|
|
||||||
@Sendable func healthResponse() -> String {
|
@Sendable func healthResponse() -> String {
|
||||||
return "{\"status\":\"healthy\",\"model\":\"e4b\",\"layers\":\(layers)}"
|
let m = isEmbeddingGemma ? "embeddinggemma-300m" : "e4b"
|
||||||
|
return "{\"status\":\"healthy\",\"model\":\"\(m)\",\"layers\":\(layers)}"
|
||||||
}
|
}
|
||||||
|
|
||||||
router.get("/") { _, _ in
|
router.get("/") { _, _ in
|
||||||
@@ -229,7 +250,7 @@ struct SimpleServerApp {
|
|||||||
topP: topP
|
topP: topP
|
||||||
)
|
)
|
||||||
|
|
||||||
let generatedTokens = try generator.generateTokens(promptTokens: promptTokens, config: config)
|
let generatedTokens = try generator!.generateTokens(promptTokens: promptTokens, config: config)
|
||||||
|
|
||||||
let id = UUID().uuidString
|
let id = UUID().uuidString
|
||||||
let ts = Int(Date().timeIntervalSince1970)
|
let ts = Int(Date().timeIntervalSince1970)
|
||||||
@@ -313,7 +334,14 @@ struct SimpleServerApp {
|
|||||||
var embeddings: [[String: Any]] = []
|
var embeddings: [[String: Any]] = []
|
||||||
for (i, text) in inputs.enumerated() {
|
for (i, text) in inputs.enumerated() {
|
||||||
let t0 = Date()
|
let t0 = Date()
|
||||||
let embedding = try embeddingModel.embed(text: text)
|
let embedding: [Float]
|
||||||
|
if let eg = embedModelCopy {
|
||||||
|
embedding = try eg.embed(text: text)
|
||||||
|
} else if let tem = textEmbedModelCopy {
|
||||||
|
embedding = try tem.embed(text: text)
|
||||||
|
} else {
|
||||||
|
return "{\"error\":\"no embedding model available\",\"type\":\"server_error\",\"code\":500}"
|
||||||
|
}
|
||||||
let duration = Date().timeIntervalSince(t0)
|
let duration = Date().timeIntervalSince(t0)
|
||||||
|
|
||||||
let embeddingData: [String: Any]
|
let embeddingData: [String: Any]
|
||||||
@@ -338,7 +366,7 @@ struct SimpleServerApp {
|
|||||||
|
|
||||||
let id = UUID().uuidString
|
let id = UUID().uuidString
|
||||||
let ts = Int(Date().timeIntervalSince1970)
|
let ts = Int(Date().timeIntervalSince1970)
|
||||||
let totalTokens = inputs.reduce(0) { $0 + tokenizer.encode(text: $1).count }
|
let totalTokens = inputs.reduce(0) { $0 + $1.components(separatedBy: .whitespaces).count }
|
||||||
|
|
||||||
let response: [String: Any] = [
|
let response: [String: Any] = [
|
||||||
"id": id,
|
"id": id,
|
||||||
|
|||||||
@@ -1,9 +1,17 @@
|
|||||||
import Foundation
|
import Foundation
|
||||||
|
|
||||||
// Entry point — avoids @main conflict with top-level code
|
// Entry point — routes to appropriate server based on model name
|
||||||
Task {
|
Task {
|
||||||
do {
|
do {
|
||||||
try await SimpleServerApp.main()
|
let args = CommandLine.arguments
|
||||||
|
let modelName = args.count > 1 ? args[1] : "E4B-MarkBase"
|
||||||
|
let isEmbedding = modelName.contains("embedding") || modelName.contains("embeddinggemma")
|
||||||
|
|
||||||
|
if isEmbedding {
|
||||||
|
try await EmbeddingServerApp.main()
|
||||||
|
} else {
|
||||||
|
try await SimpleServerApp.main()
|
||||||
|
}
|
||||||
} catch {
|
} catch {
|
||||||
print("Server error: \(error)")
|
print("Server error: \(error)")
|
||||||
exit(1)
|
exit(1)
|
||||||
|
|||||||
@@ -0,0 +1,29 @@
|
|||||||
|
import XCTest
|
||||||
|
@testable import MarkBase
|
||||||
|
|
||||||
|
final class EmbeddingGemmaTest: XCTestCase {
|
||||||
|
func testLoadAndEmbed() throws {
|
||||||
|
let modelDir = "/Users/accusys/MarkBaseEngine/models/embeddinggemma-300m"
|
||||||
|
guard FileManager.default.fileExists(atPath: modelDir + "/model.safetensors") else {
|
||||||
|
XCTFail("Model not found"); return
|
||||||
|
}
|
||||||
|
let engine = try MarkBaseEngine(autoCompile: true)
|
||||||
|
let model = try EmbeddingGemmaModel(modelDir: modelDir, engine: engine)
|
||||||
|
XCTAssertEqual(model.config.numHiddenLayers, 24)
|
||||||
|
XCTAssertEqual(model.config.hiddenSize, 768)
|
||||||
|
|
||||||
|
let emb = try model.embed(text: "Hello world")
|
||||||
|
XCTAssertEqual(emb.count, 768)
|
||||||
|
|
||||||
|
// Check L2 norm is 1.0
|
||||||
|
var norm: Float = 0
|
||||||
|
for v in emb { norm += v * v }
|
||||||
|
norm = sqrt(norm)
|
||||||
|
XCTAssertEqual(norm, 1.0, accuracy: 0.001, "L2 norm should be 1.0")
|
||||||
|
|
||||||
|
// Check no NaN
|
||||||
|
XCTAssertFalse(emb.contains { $0.isNaN }, "No NaN values")
|
||||||
|
|
||||||
|
print("✓ Embedding dim=\(emb.count), norm=\(String(format: "%.4f", norm))")
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user