Compare commits
9 Commits
88aeff7935
...
91d44a924d
| Author | SHA1 | Date | |
|---|---|---|---|
| 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"],
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -1,8 +1,9 @@
|
|||||||
import Foundation
|
import Foundation
|
||||||
import Metal
|
import Metal
|
||||||
|
import MetalPerformanceShaders
|
||||||
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
|
||||||
@@ -41,14 +42,12 @@ public struct EmbeddingGemmaConfig: Codable {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/// EmbeddingGemma - Google's 300M parameter embedding model
|
public final class EmbeddingGemmaModel: @unchecked Sendable {
|
||||||
public final class EmbeddingGemmaModel {
|
|
||||||
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,21 +66,13 @@ 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))")
|
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)!
|
||||||
|
|
||||||
for i in 0..<config.numHiddenLayers {
|
for i in 0..<config.numHiddenLayers {
|
||||||
let p = "layers.\(i)"
|
let p = "layers.\(i)"
|
||||||
layerNorms.append([
|
layerNorms.append([
|
||||||
@@ -90,22 +81,20 @@ 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
|
|
||||||
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)) }
|
||||||
@@ -113,22 +102,37 @@ public final class EmbeddingGemmaModel {
|
|||||||
|
|
||||||
let seqLen = tokens.count, hs = config.hiddenSize
|
let seqLen = tokens.count, hs = config.hiddenSize
|
||||||
|
|
||||||
// Embedding lookup
|
// Test 1: Embedding lookup only
|
||||||
let inputBuf = try lookupEmbeddings(tokens: tokens)
|
print(" TEST: Embedding lookup...")
|
||||||
|
let cmdBuf1 = engine.commandQueue.makeCommandBuffer()!
|
||||||
|
let inputBuf = try lookupEmbeddings(tokens: tokens, cmdBuf: cmdBuf1)
|
||||||
|
cmdBuf1.commit()
|
||||||
|
cmdBuf1.waitUntilCompleted()
|
||||||
|
print(" TEST: Embedding lookup OK")
|
||||||
|
|
||||||
// Forward through layers
|
// Test 2: Single layer forward (layer 0 only)
|
||||||
var hidden = inputBuf
|
print(" TEST: Layer 0 forward...")
|
||||||
|
let cmdBuf2 = engine.commandQueue.makeCommandBuffer()!
|
||||||
|
var hidden = try forwardLayerDebug(hidden: inputBuf, layerIdx: 0, seqLen: seqLen, cmdBuf: cmdBuf2)
|
||||||
|
cmdBuf2.commit()
|
||||||
|
cmdBuf2.waitUntilCompleted()
|
||||||
|
print(" TEST: Layer 0 OK")
|
||||||
|
|
||||||
|
// Full forward pass
|
||||||
|
print(" TEST: Full forward pass...")
|
||||||
|
let cmdBuf = engine.commandQueue.makeCommandBuffer()!
|
||||||
|
hidden = inputBuf
|
||||||
for idx in 0..<config.numHiddenLayers {
|
for idx in 0..<config.numHiddenLayers {
|
||||||
hidden = try forwardLayer(hidden: hidden, layerIdx: idx, seqLen: seqLen)
|
if idx % 6 == 0 { print(" Layer \(idx)/\(config.numHiddenLayers)") }
|
||||||
|
hidden = try forwardLayerDebug(hidden: hidden, layerIdx: idx, seqLen: seqLen, cmdBuf: cmdBuf)
|
||||||
}
|
}
|
||||||
|
print(" TEST: All layers OK")
|
||||||
|
|
||||||
// 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
|
||||||
@@ -136,29 +140,20 @@ public final class EmbeddingGemmaModel {
|
|||||||
}
|
}
|
||||||
let n = Float(seqLen)
|
let n = Float(seqLen)
|
||||||
for i in 0..<hs { embedding[i] /= n }
|
for i in 0..<hs { embedding[i] /= n }
|
||||||
|
|
||||||
var norm: Float = 0
|
var norm: Float = 0
|
||||||
for i in 0..<hs { norm += embedding[i] * embedding[i] }
|
for i in 0..<hs { norm += embedding[i] * embedding[i] }
|
||||||
norm = sqrt(norm)
|
norm = sqrt(norm)
|
||||||
if norm > 0 { for i in 0..<hs { embedding[i] /= norm } }
|
if norm > 0 { for i in 0..<hs { embedding[i] /= norm } }
|
||||||
|
|
||||||
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 +162,45 @@ 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 loadAndTranspose(_ name: String, rows: Int, cols: Int) throws -> MTLBuffer {
|
||||||
|
// Load [rows, cols] and transpose to [cols, rows] for MPS matmul C = A × B
|
||||||
|
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)!
|
||||||
|
}
|
||||||
|
|
||||||
|
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)!
|
print(" lookupEmbeddings: seqLen=\(seqLen), hs=\(hs)")
|
||||||
let cmdBuf = engine.commandQueue.makeCommandBuffer()!
|
// Read embedding table to CPU
|
||||||
let enc = cmdBuf.makeComputeCommandEncoder()!
|
let embedPtr = embedTokens.contents().assumingMemoryBound(to: Float.self)
|
||||||
let pso = try engine.pipeline(named: "lookup_embeddings")
|
let embedCount = embedTokens.length / 4
|
||||||
enc.setComputePipelineState(pso)
|
let embedArray = Array(UnsafeBufferPointer(start: embedPtr, count: embedCount))
|
||||||
enc.setBuffer(embedTokens, offset: 0, index: 0)
|
print(" Read \(embedCount) floats from embedTokens")
|
||||||
enc.setBytes(tokens, length: seqLen * MemoryLayout<Int>.size, index: 1)
|
|
||||||
enc.setBuffer(buf, offset: 0, index: 2)
|
// Lookup embeddings for tokens
|
||||||
var h = UInt32(hs), s = UInt32(seqLen), v = UInt32(config.vocabSize)
|
var embedData = [Float](repeating: 0, count: seqLen * hs)
|
||||||
enc.setBytes(&h, length: 4, index: 3)
|
for (i, token) in tokens.enumerated() {
|
||||||
enc.setBytes(&s, length: 4, index: 4)
|
let start = i * hs
|
||||||
enc.setBytes(&v, length: 4, index: 5)
|
let srcStart = token * hs
|
||||||
enc.dispatchThreads(MTLSize(width: seqLen, height: 1, depth: 1),
|
embedData[start..<start+hs] = embedArray[srcStart..<srcStart+hs]
|
||||||
threadsPerThreadgroup: MTLSize(width: min(256, seqLen), height: 1, depth: 1))
|
}
|
||||||
enc.endEncoding()
|
print(" Looked up \(seqLen) tokens")
|
||||||
cmdBuf.commit(); cmdBuf.waitUntilCompleted()
|
|
||||||
|
let buf = engine.device.makeBuffer(bytes: embedData, length: embedData.count * 4)!
|
||||||
|
print(" Created MTLBuffer")
|
||||||
return buf
|
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)
|
||||||
@@ -200,104 +209,60 @@ public final class EmbeddingGemmaModel {
|
|||||||
var c = UInt32(count), e: Float = config.rmsNormEps
|
var c = UInt32(count), e: Float = config.rmsNormEps
|
||||||
enc.setBytes(&c, length: 4, index: 3)
|
enc.setBytes(&c, length: 4, index: 3)
|
||||||
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()!
|
// Use MPS for optimized matrix multiplication on Apple Silicon
|
||||||
let pso = try engine.pipeline(named: "matmul_f32")
|
// Weight is stored transposed [k, n] for C = A × B
|
||||||
enc.setComputePipelineState(pso)
|
let descA = MPSMatrixDescriptor(rows: m, columns: k, rowBytes: k * 4, dataType: .float32)
|
||||||
enc.setBuffer(input, offset: 0, index: 0)
|
let descB = MPSMatrixDescriptor(rows: k, columns: n, rowBytes: n * 4, dataType: .float32)
|
||||||
enc.setBuffer(weight, offset: 0, index: 1)
|
let descC = MPSMatrixDescriptor(rows: m, columns: n, rowBytes: n * 4, dataType: .float32)
|
||||||
enc.setBuffer(output, offset: 0, index: 2)
|
let matA = MPSMatrix(buffer: input, descriptor: descA)
|
||||||
var mm = UInt32(m), kk = UInt32(k), nn = UInt32(n)
|
let matB = MPSMatrix(buffer: weight, descriptor: descB)
|
||||||
enc.setBytes(&mm, length: 4, index: 3)
|
let matC = MPSMatrix(buffer: output, descriptor: descC)
|
||||||
enc.setBytes(&kk, length: 4, index: 4)
|
|
||||||
enc.setBytes(&nn, length: 4, index: 5)
|
let matMul = MPSMatrixMultiplication(device: engine.device,
|
||||||
enc.dispatchThreads(MTLSize(width: m * n, height: 1, depth: 1),
|
transposeLeft: false,
|
||||||
threadsPerThreadgroup: MTLSize(width: min(256, m * n), height: 1, depth: 1))
|
transposeRight: false,
|
||||||
enc.endEncoding()
|
resultRows: m,
|
||||||
|
resultColumns: n,
|
||||||
|
interiorColumns: k,
|
||||||
|
alpha: 1.0,
|
||||||
|
beta: 0.0)
|
||||||
|
matMul.encode(commandBuffer: cmdBuf, leftMatrix: matA, rightMatrix: matB, resultMatrix: matC)
|
||||||
}
|
}
|
||||||
|
|
||||||
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)
|
||||||
@@ -308,18 +273,12 @@ public final class EmbeddingGemmaModel {
|
|||||||
enc.setBytes(&hd, length: 4, index: 3)
|
enc.setBytes(&hd, length: 4, index: 3)
|
||||||
enc.setBytes(&nh, length: 4, index: 4)
|
enc.setBytes(&nh, length: 4, index: 4)
|
||||||
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 +294,84 @@ 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 forwardLayerDebug(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 {
|
print(" Residual copy...")
|
||||||
let enc = cmdBuf.makeComputeCommandEncoder()!
|
let resid = device.makeBuffer(length: seqLen * hs * 4)!
|
||||||
let pso = try engine.pipeline(named: "gelu_mul_kernel")
|
let blit = cmdBuf.makeBlitCommandEncoder()!
|
||||||
enc.setComputePipelineState(pso)
|
blit.copy(from: hidden, sourceOffset: 0, to: resid, destinationOffset: 0, size: seqLen * hs * 4)
|
||||||
enc.setBuffer(gate, offset: 0, index: 0)
|
blit.endEncoding()
|
||||||
enc.setBuffer(up, offset: 0, index: 1)
|
|
||||||
enc.setBuffer(output, offset: 0, index: 2)
|
print(" Input norm...")
|
||||||
var c = UInt32(count)
|
let h1 = try applyRmsNorm(input: hidden, weight: layerNorms[layerIdx][0], count: seqLen * hs, cmdBuf: cmdBuf)
|
||||||
enc.setBytes(&c, length: 4, index: 3)
|
|
||||||
enc.dispatchThreads(MTLSize(width: count, height: 1, depth: 1),
|
print(" QKV projections...")
|
||||||
threadsPerThreadgroup: MTLSize(width: min(256, count), height: 1, depth: 1))
|
let qBuf = device.makeBuffer(length: seqLen * nH * hDim * 4)!
|
||||||
enc.endEncoding()
|
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)
|
||||||
|
print(" Q done")
|
||||||
|
try matmulSeq(input: h1, weight: kProjs[layerIdx], output: kBuf, m: seqLen, k: hs, n: nKV * hDim, cmdBuf: cmdBuf)
|
||||||
|
print(" K done")
|
||||||
|
try matmulSeq(input: h1, weight: vProjs[layerIdx], output: vBuf, m: seqLen, k: hs, n: nKV * hDim, cmdBuf: cmdBuf)
|
||||||
|
print(" V done")
|
||||||
|
|
||||||
|
print(" RoPE...")
|
||||||
|
try applyRoPE(q: qBuf, k: kBuf, seqLen: seqLen, headDim: hDim, numHeads: nH, cmdBuf: cmdBuf)
|
||||||
|
print(" RoPE done")
|
||||||
|
|
||||||
|
print(" Attention...")
|
||||||
|
let attnOut = device.makeBuffer(length: seqLen * nH * hDim * 4)!
|
||||||
|
try bidirectionalAttention(q: qBuf, k: kBuf, v: vBuf, output: attnOut, seqLen: seqLen, cmdBuf: cmdBuf)
|
||||||
|
print(" Attention done")
|
||||||
|
|
||||||
|
print(" 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)
|
||||||
|
print(" O done")
|
||||||
|
|
||||||
|
print(" Post-attn norm...")
|
||||||
|
let h2n = try applyRmsNorm(input: h2, weight: layerNorms[layerIdx][2], count: seqLen * hs, cmdBuf: cmdBuf)
|
||||||
|
|
||||||
|
print(" Add residual 1...")
|
||||||
|
try eltwiseAdd(a: resid, b: h2n, output: hidden, count: seqLen * hs, cmdBuf: cmdBuf)
|
||||||
|
|
||||||
|
print(" Pre-FF norm...")
|
||||||
|
let h3 = try applyRmsNorm(input: hidden, weight: layerNorms[layerIdx][1], count: seqLen * hs, cmdBuf: cmdBuf)
|
||||||
|
|
||||||
|
print(" 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)
|
||||||
|
print(" Gate done")
|
||||||
|
try matmulSeq(input: h3, weight: upProjs[layerIdx], output: up, m: seqLen, k: hs, n: intermedi, cmdBuf: cmdBuf)
|
||||||
|
print(" Up done")
|
||||||
|
|
||||||
|
print(" GELU mul...")
|
||||||
|
let gated = device.makeBuffer(length: seqLen * intermedi * 4)!
|
||||||
|
try geluMul(gate: gate, up: up, output: gated, count: seqLen * intermedi, cmdBuf: cmdBuf)
|
||||||
|
print(" GELU done")
|
||||||
|
|
||||||
|
print(" 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)
|
||||||
|
print(" Down done")
|
||||||
|
|
||||||
|
print(" Post-FF norm...")
|
||||||
|
let h4n = try applyRmsNorm(input: h4, weight: layerNorms[layerIdx][3], count: seqLen * hs, cmdBuf: cmdBuf)
|
||||||
|
|
||||||
|
print(" Add residual 2...")
|
||||||
|
try eltwiseAdd(a: hidden, b: h4n, output: hidden, count: seqLen * hs, cmdBuf: cmdBuf)
|
||||||
|
print(" Layer \(layerIdx) complete")
|
||||||
|
|
||||||
|
return hidden
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -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)
|
|
||||||
|
// 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)
|
let tokenizer = try TokenizerFactory.load(modelDir: modelPath)
|
||||||
let generator = StreamingGenerator(model: model, tokenizer: tokenizer, engine: engine)
|
var generator: StreamingGenerator? = nil
|
||||||
let embeddingModel = try TextEmbeddingModel(modelDir: modelPath, engine: engine, config: TextEmbeddingConfig())
|
|
||||||
|
if isE4B {
|
||||||
print("✓ E4B loaded (\(model.numHiddenLayers) layers)")
|
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