Fork for thermals request add-json-schema-dpeq
0

Configure Feed

Select the types of activity you want to include in your feed.

vmlx-swift / Libraries / MLXLLM / Models / Qwen3Next.swift
27 kB 721 lines
1// 2// Qwen3Next.swift 3// mlx-swift-lm 4// 5// Port of https://github.com/ml-explore/mlx-lm/blob/main/mlx_lm/models/qwen3_next.py 6// 7 8import Foundation 9import MLX 10import MLXLMCommon 11import MLXNN 12 13// MARK: - Helpers 14 15/// Compiled sigmoid gate: fuses sigmoid + multiply into one Metal dispatch. 16/// Used per-layer in Qwen3.5 attention (GatedDeltaNet) 40 layers per forward. 17private let _compiledSigmoidMultiply: @Sendable (MLXArray, MLXArray) -> MLXArray = { 18 let body: @Sendable (MLXArray, MLXArray) -> MLXArray = { (x: MLXArray, gate: MLXArray) -> MLXArray in 19 x * sigmoid(gate) 20 } 21 return HardwareInfo.isCompiledDecodeSupported ? compile(shapeless: true, body) : body 22}() 23 24func sigmoidMultiply(_ x: MLXArray, _ gate: MLXArray) -> MLXArray { 25 _compiledSigmoidMultiply(x, gate) 26} 27 28/// Compiled precise SwiGLU matches Python `@mx.compile` `_precise_swiglu`. 29/// Fuses silu(gate) in float32 + cast-to-fp32 + multiply + cast-to-hdtype into 30/// ONE Metal dispatch. Without this, every GatedDeltaNet layer issues 5 separate 31/// Metal kernels (silu + asType + asType + multiply + asType) and at ~0.1ms per 32/// dispatch times ~30 layers, this is the dominant per-step overhead vs Python. 33private let _compiledPreciseSwiGLU: @Sendable (MLXArray, MLXArray, MLXArray) -> MLXArray = { 34 let body: @Sendable (MLXArray, MLXArray, MLXArray) -> MLXArray = { (h: MLXArray, gate: MLXArray, x: MLXArray) -> MLXArray in 35 let gateF32 = silu(gate.asType(.float32)) 36 let xF32 = x.asType(.float32) 37 return (gateF32 * xF32).asType(h.dtype) 38 } 39 return HardwareInfo.isCompiledDecodeSupported ? compile(shapeless: true, body) : body 40}() 41 42// MARK: - Model Components 43 44// Compiled swiglu for MLP: fuses silu(gate) * up into 1 Metal dispatch. 45private let _compiledSwiGLU: @Sendable (MLXArray, MLXArray) -> MLXArray = { 46 let body: @Sendable (MLXArray, MLXArray) -> MLXArray = { (gate: MLXArray, x: MLXArray) -> MLXArray in 47 silu(gate) * x 48 } 49 return HardwareInfo.isCompiledDecodeSupported ? compile(shapeless: true, body) : body 50}() 51 52final class Qwen3NextRMSNormGated: Module { 53 @ParameterInfo(key: "weight") var weight: MLXArray 54 let eps: Float 55 56 init(dimensions: Int, eps: Float) { 57 self.eps = eps 58 self._weight.wrappedValue = MLXArray.ones([dimensions]) 59 super.init() 60 } 61 62 func callAsFunction(_ hiddenStates: MLXArray, gate: MLXArray? = nil) -> MLXArray { 63 let x = MLXFast.rmsNorm(hiddenStates, weight: weight, eps: eps) 64 if let gate { 65 // Match Python's _precise_swiglu: float32 cast + silu + multiply, then cast back. 66 return _compiledPreciseSwiGLU(hiddenStates, gate, x) 67 } 68 return x 69 } 70} 71 72public final class Qwen3NextAttention: Module { 73 let args: Qwen3NextConfiguration 74 let scale: Float 75 76 @ModuleInfo(key: "q_proj") var qProj: Linear 77 @ModuleInfo(key: "k_proj") var kProj: Linear 78 @ModuleInfo(key: "v_proj") var vProj: Linear 79 @ModuleInfo(key: "o_proj") var oProj: Linear 80 81 @ModuleInfo(key: "q_norm") var qNorm: RMSNorm 82 @ModuleInfo(key: "k_norm") var kNorm: RMSNorm 83 84 let rope: RoPELayer 85 86 init(_ args: Qwen3NextConfiguration) { 87 self.args = args 88 89 let headDim = args.headDim ?? (args.hiddenSize / args.attentionHeads) 90 self.scale = pow(Float(headDim), -0.5) 91 92 _qProj.wrappedValue = Linear( 93 args.hiddenSize, args.attentionHeads * headDim * 2, bias: args.attentionBias) 94 _kProj.wrappedValue = Linear( 95 args.hiddenSize, args.kvHeads * headDim, bias: args.attentionBias) 96 _vProj.wrappedValue = Linear( 97 args.hiddenSize, args.kvHeads * headDim, bias: args.attentionBias) 98 _oProj.wrappedValue = Linear( 99 args.attentionHeads * headDim, args.hiddenSize, bias: args.attentionBias) 100 101 _qNorm.wrappedValue = RMSNorm(dimensions: headDim, eps: args.rmsNormEps) 102 _kNorm.wrappedValue = RMSNorm(dimensions: headDim, eps: args.rmsNormEps) 103 104 let ropeDims = Int(Float(headDim) * args.partialRotaryFactor) 105 self.rope = initializeRope( 106 dims: max(1, ropeDims), 107 base: args.ropeTheta, 108 traditional: false, 109 scalingConfig: args.ropeScaling, 110 maxPositionEmbeddings: args.maxPositionEmbeddings 111 ) 112 113 super.init() 114 } 115 116 public func callAsFunction( 117 _ x: MLXArray, mask: MLXFast.ScaledDotProductAttentionMaskMode, cache: KVCache? 118 ) -> MLXArray { 119 let B = x.dim(0) 120 let L = x.dim(1) 121 122 let qProjOutput = qProj(x) 123 let qSplit = qProjOutput.reshaped(B, L, args.attentionHeads, -1).split(parts: 2, axis: -1) 124 // Bail on a failed split rather than trapping on the subscripts below. 125 guard qSplit.count == 2 else { return x } 126 var queries = qSplit[0] 127 let gate = qSplit[1].reshaped(B, L, -1) 128 129 var keys = kProj(x) 130 var values = vProj(x) 131 132 queries = qNorm(queries).transposed(0, 2, 1, 3) 133 keys = kNorm(keys.reshaped(B, L, args.kvHeads, -1)).transposed(0, 2, 1, 3) 134 values = values.reshaped(B, L, args.kvHeads, -1).transposed(0, 2, 1, 3) 135 136 queries = applyRotaryPosition(rope, to: queries, cache: cache) 137 keys = applyRotaryPosition(rope, to: keys, cache: cache) 138 139 let output = attentionWithCacheUpdate( 140 queries: queries, 141 keys: keys, 142 values: values, 143 cache: cache, 144 scale: scale, 145 mask: mask 146 ) 147 .transposed(0, 2, 1, 3) 148 .reshaped(B, L, -1) 149 150 return oProj(sigmoidMultiply(output, gate)) 151 } 152} 153 154final class Qwen3NextMLP: Module, UnaryLayer { 155 @ModuleInfo(key: "gate_proj") var gateProj: Linear 156 @ModuleInfo(key: "down_proj") var downProj: Linear 157 @ModuleInfo(key: "up_proj") var upProj: Linear 158 159 init(dimensions: Int, hiddenDimensions: Int) { 160 _gateProj.wrappedValue = Linear(dimensions, hiddenDimensions, bias: false) 161 _downProj.wrappedValue = Linear(hiddenDimensions, dimensions, bias: false) 162 _upProj.wrappedValue = Linear(dimensions, hiddenDimensions, bias: false) 163 } 164 165 func callAsFunction(_ x: MLXArray) -> MLXArray { 166 // Fuses silu(gate) * up into 1 Metal dispatch instead of 2. 167 let activated = _compiledSwiGLU(gateProj(x), upProj(x)) 168 return downProj(activated) 169 } 170} 171 172public final class Qwen3NextGatedDeltaNet: Module { 173 let hiddenSize: Int 174 let numVHeads: Int 175 let numKHeads: Int 176 let headKDim: Int 177 let headVDim: Int 178 let keyDim: Int 179 let valueDim: Int 180 let convKernelSize: Int 181 let convDim: Int 182 183 @ModuleInfo(key: "conv1d") var conv1d: Conv1d 184 @ModuleInfo(key: "in_proj_qkvz") var inProjQKVZ: Linear 185 @ModuleInfo(key: "in_proj_ba") var inProjBA: Linear 186 187 @ParameterInfo(key: "dt_bias") var dtBias: MLXArray 188 @ParameterInfo(key: "A_log") var aLog: MLXArray 189 190 @ModuleInfo(key: "norm") var norm: Qwen3NextRMSNormGated 191 @ModuleInfo(key: "out_proj") var outProj: Linear 192 193 init(_ args: Qwen3NextConfiguration) { 194 self.hiddenSize = args.hiddenSize 195 self.numVHeads = args.linearNumValueHeads 196 self.numKHeads = args.linearNumKeyHeads 197 self.headKDim = args.linearKeyHeadDim 198 self.headVDim = args.linearValueHeadDim 199 self.keyDim = headKDim * numKHeads 200 self.valueDim = headVDim * numVHeads 201 self.convKernelSize = args.linearConvKernelDim 202 self.convDim = keyDim * 2 + valueDim 203 204 precondition(numVHeads % numKHeads == 0, "num_v_heads must be divisible by num_k_heads") 205 206 _conv1d.wrappedValue = Conv1d( 207 inputChannels: convDim, 208 outputChannels: convDim, 209 kernelSize: convKernelSize, 210 stride: 1, 211 padding: 0, 212 dilation: 1, 213 groups: convDim, 214 bias: false 215 ) 216 217 _inProjQKVZ.wrappedValue = Linear( 218 hiddenSize, keyDim * 2 + valueDim * 2, bias: false) 219 _inProjBA.wrappedValue = Linear(hiddenSize, numVHeads * 2, bias: false) 220 221 _dtBias.wrappedValue = MLXArray.ones([numVHeads]) 222 let a = MLXRandom.uniform(low: 0, high: 16, [numVHeads]) 223 _aLog.wrappedValue = log(a) 224 225 _norm.wrappedValue = Qwen3NextRMSNormGated(dimensions: headVDim, eps: args.rmsNormEps) 226 _outProj.wrappedValue = Linear(valueDim, hiddenSize, bias: false) 227 228 super.init() 229 } 230 231 private func fixQueryKeyValueOrdering( 232 mixedQKVZ: MLXArray, 233 mixedBA: MLXArray 234 ) -> (MLXArray, MLXArray, MLXArray, MLXArray, MLXArray, MLXArray) { 235 let B = mixedQKVZ.dim(0) 236 let S = mixedQKVZ.dim(1) 237 let nk = numKHeads 238 let dn = headKDim 239 let nv = numVHeads 240 let dv = headVDim 241 let vHeadsPerK = nv / nk 242 243 let qkvz = mixedQKVZ.reshaped(B, S, nk, -1) 244 let ba = mixedBA.reshaped(B, S, nk, -1) 245 246 let qkvzSplit = MLX.split( 247 qkvz, 248 indices: [dn, 2 * dn, 2 * dn + vHeadsPerK * dv], 249 axis: -1 250 ) 251 // Bail on a failed split rather than trapping on the subscripts below; 252 // the recorded MLX error surfaces at the next eval. 253 guard qkvzSplit.count == 4 else { return (qkvz, qkvz, qkvz, qkvz, ba, ba) } 254 let q = qkvzSplit[0] 255 let k = qkvzSplit[1] 256 let v = qkvzSplit[2].reshaped(B, S, -1, dv) 257 let z = qkvzSplit[3].reshaped(B, S, -1, dv) 258 259 let baSplit = MLX.split(ba, indices: [vHeadsPerK], axis: -1) 260 guard baSplit.count == 2 else { return (q, k, v, z, ba, ba) } 261 let b = baSplit[0].reshaped(B, S, nv) 262 let a = baSplit[1].reshaped(B, S, nv) 263 264 return (q, k, v, z, b, a) 265 } 266 267 public func callAsFunction( 268 _ inputs: MLXArray, 269 mask: MLXArray? = nil, 270 cache: MambaCache? = nil 271 ) -> MLXArray { 272 let B = inputs.dim(0) 273 let S = inputs.dim(1) 274 275 let (q, k, v, z, b, a) = fixQueryKeyValueOrdering( 276 mixedQKVZ: inProjQKVZ(inputs), 277 mixedBA: inProjBA(inputs) 278 ) 279 280 let dtype = inputs.dtype 281 let convState: MLXArray 282 if let cacheState = cache?[0] { 283 convState = cacheState 284 } else { 285 convState = MLXArray.zeros([B, convKernelSize - 1, convDim], dtype: dtype) 286 } 287 288 var mixedQKV = concatenated( 289 [q.reshaped(B, S, -1), k.reshaped(B, S, -1), v.reshaped(B, S, -1)], 290 axis: -1 291 ) 292 293 if let mask { 294 mixedQKV = MLX.where( 295 expandedDimensions(mask, axis: -1), mixedQKV, MLXArray.zeros(like: mixedQKV)) 296 } 297 298 let convInput = concatenated([convState, mixedQKV], axis: 1) 299 if let cache { 300 cache[0] = convInput[0..., (1 - convKernelSize)..., 0...] 301 } 302 303 let convOut = silu(conv1d(convInput)) 304 let convSplit = MLX.split(convOut, indices: [keyDim, 2 * keyDim], axis: -1) 305 306 // Bail on a failed split rather than trapping on the subscripts below. 307 guard convSplit.count == 3 else { return inputs } 308 var qOut = convSplit[0].reshaped(B, S, numKHeads, headKDim) 309 var kOut = convSplit[1].reshaped(B, S, numKHeads, headKDim) 310 let vOut = convSplit[2].reshaped(B, S, numVHeads, headVDim) 311 312 let invScale = pow(Float(headKDim), -0.5) 313 qOut = 314 MLXArray(invScale * invScale, dtype: qOut.dtype) 315 * MLXFast.rmsNorm(qOut, weight: MLXArray.mlxNone, eps: 1e-6) 316 kOut = 317 MLXArray(invScale, dtype: kOut.dtype) 318 * MLXFast.rmsNorm(kOut, weight: MLXArray.mlxNone, eps: 1e-6) 319 320 let (out, newState) = gatedDeltaUpdate( 321 q: qOut, 322 k: kOut, 323 v: vOut, 324 a: a, 325 b: b, 326 aLog: aLog, 327 dtBias: dtBias, 328 state: cache?[1], 329 mask: mask 330 ) 331 332 if let cache { 333 cache[1] = newState 334 } 335 336 let normalized = norm(out, gate: z) 337 return outProj(normalized.reshaped(B, S, -1)) 338 } 339} 340 341final class Qwen3NextSparseMoeBlock: Module { 342 let layerIdx: Int 343 let normTopkProb: Bool 344 let numExperts: Int 345 let topK: Int 346 347 @ModuleInfo(key: "gate") var gate: Linear 348 @ModuleInfo(key: "switch_mlp") var switchMLP: SwitchGLU 349 350 @ModuleInfo(key: "shared_expert") var sharedExpert: Qwen3NextMLP 351 @ModuleInfo(key: "shared_expert_gate") var sharedExpertGate: Linear 352 353 init(_ args: Qwen3NextConfiguration, layerIdx: Int) { 354 self.layerIdx = layerIdx 355 self.normTopkProb = args.normTopkProb 356 self.numExperts = args.numExperts 357 self.topK = args.numExpertsPerTok 358 359 _gate.wrappedValue = Linear(args.hiddenSize, args.numExperts, bias: false) 360 _switchMLP.wrappedValue = SwitchGLU( 361 inputDims: args.hiddenSize, 362 hiddenDims: args.moeIntermediateSize, 363 numExperts: args.numExperts 364 ) 365 366 _sharedExpert.wrappedValue = Qwen3NextMLP( 367 dimensions: args.hiddenSize, 368 hiddenDimensions: args.sharedExpertIntermediateSize 369 ) 370 _sharedExpertGate.wrappedValue = Linear(args.hiddenSize, 1, bias: false) 371 } 372 373 func callAsFunction(_ x: MLXArray) -> MLXArray { 374 var gates = gate(x) 375 gates = MLX.softmax(gates, axis: -1, precise: true) 376 377 let k = topK 378 let kth = gates.dim(-1) - k 379 let inds = MLX.argPartition(gates, kth: kth, axis: -1)[.ellipsis, (kth)...] 380 JangPressCanonicalExpertAdvisor.shared.observe(layer: layerIdx, indices: inds) 381 var scores = MLX.takeAlong(gates, inds, axis: -1) 382 if normTopkProb { 383 scores = scores / scores.sum(axis: -1, keepDims: true) 384 } 385 386 let y = switchMLP(x, inds) 387 let combined = (y * scores[.ellipsis, .newAxis]).sum(axis: -2) 388 389 var sharedY = sharedExpert(x) 390 sharedY = sigmoid(sharedExpertGate(x)) * sharedY 391 392 return combined + sharedY 393 } 394} 395 396final class Qwen3NextDecoderLayer: Module { 397 let isLinear: Bool 398 399 @ModuleInfo(key: "self_attn") var selfAttn: Qwen3NextAttention? 400 @ModuleInfo(key: "linear_attn") var linearAttn: Qwen3NextGatedDeltaNet? 401 402 @ModuleInfo(key: "input_layernorm") var inputLayerNorm: RMSNorm 403 @ModuleInfo(key: "post_attention_layernorm") var postAttentionLayerNorm: RMSNorm 404 405 @ModuleInfo(key: "mlp") var mlp: Module 406 407 init(_ args: Qwen3NextConfiguration, layerIdx: Int) { 408 self.isLinear = (layerIdx + 1) % args.fullAttentionInterval != 0 409 410 if isLinear { 411 _linearAttn.wrappedValue = Qwen3NextGatedDeltaNet(args) 412 } else { 413 _selfAttn.wrappedValue = Qwen3NextAttention(args) 414 } 415 416 _inputLayerNorm.wrappedValue = RMSNorm(dimensions: args.hiddenSize, eps: args.rmsNormEps) 417 _postAttentionLayerNorm.wrappedValue = RMSNorm( 418 dimensions: args.hiddenSize, eps: args.rmsNormEps) 419 420 let useMoE = 421 !args.mlpOnlyLayers.contains(layerIdx) 422 && args.numExperts > 0 423 && (layerIdx + 1) % args.decoderSparseStep == 0 424 425 if useMoE { 426 _mlp.wrappedValue = Qwen3NextSparseMoeBlock(args, layerIdx: layerIdx) 427 } else { 428 _mlp.wrappedValue = Qwen3NextMLP( 429 dimensions: args.hiddenSize, 430 hiddenDimensions: args.intermediateSize 431 ) 432 } 433 434 super.init() 435 } 436 437 func callAsFunction( 438 _ x: MLXArray, 439 attentionMask: MLXFast.ScaledDotProductAttentionMaskMode, 440 ssmMask: MLXArray?, 441 cache: KVCache? 442 ) -> MLXArray { 443 let h: MLXArray 444 if isLinear { 445 h = linearAttn!(inputLayerNorm(x), mask: ssmMask, cache: cache as? MambaCache) 446 } else { 447 h = selfAttn!(inputLayerNorm(x), mask: attentionMask, cache: cache) 448 } 449 450 let r = x + h 451 let normed = postAttentionLayerNorm(r) 452 if let moe = mlp as? Qwen3NextSparseMoeBlock { 453 return r + moe(normed) 454 } 455 return r + (mlp as! Qwen3NextMLP)(normed) 456 } 457} 458 459public class Qwen3NextModelInner: Module { 460 @ModuleInfo(key: "embed_tokens") var embedTokens: Embedding 461 462 fileprivate let layers: [Qwen3NextDecoderLayer] 463 let norm: RMSNorm 464 465 let ssmIdx: Int 466 let faIdx: Int 467 468 init(_ args: Qwen3NextConfiguration) { 469 precondition(args.vocabularySize > 0) 470 471 _embedTokens.wrappedValue = Embedding( 472 embeddingCount: args.vocabularySize, 473 dimensions: args.hiddenSize 474 ) 475 476 self.layers = (0 ..< args.hiddenLayers).map { layerIdx in 477 Qwen3NextDecoderLayer(args, layerIdx: layerIdx) 478 } 479 480 self.norm = RMSNorm(dimensions: args.hiddenSize, eps: args.rmsNormEps) 481 482 self.ssmIdx = 0 483 self.faIdx = args.fullAttentionInterval - 1 484 485 super.init() 486 } 487 488 func callAsFunction(_ inputs: MLXArray, cache: [KVCache?]? = nil) -> MLXArray { 489 var hiddenStates = embedTokens(inputs) 490 491 var cacheArray = cache 492 if cacheArray == nil { 493 cacheArray = Array(repeating: nil as KVCache?, count: layers.count) 494 } 495 496 let faMask = createAttentionMask(h: hiddenStates, cache: cacheArray?[faIdx]) 497 let ssmMask = createSSMMask(h: hiddenStates, cache: cacheArray?[ssmIdx] as? MambaCache) 498 499 for (i, layer) in layers.enumerated() { 500 let mask = layer.isLinear ? ssmMask : nil 501 let attnMask = layer.isLinear ? MLXFast.ScaledDotProductAttentionMaskMode.none : faMask 502 hiddenStates = layer( 503 hiddenStates, attentionMask: attnMask, ssmMask: mask, cache: cacheArray?[i]) 504 } 505 506 return norm(hiddenStates) 507 } 508} 509 510public class Qwen3NextModel: Module, LLMModel, KVCacheDimensionProvider { 511 public let vocabularySize: Int 512 public let kvHeads: [Int] 513 514 public let model: Qwen3NextModelInner 515 let configuration: Qwen3NextConfiguration 516 517 @ModuleInfo(key: "lm_head") var lmHead: Linear? 518 519 public init(_ args: Qwen3NextConfiguration) { 520 self.configuration = args 521 self.vocabularySize = args.vocabularySize 522 self.kvHeads = (0 ..< args.hiddenLayers).map { _ in args.kvHeads } 523 self.model = Qwen3NextModelInner(args) 524 525 if !args.tieWordEmbeddings { 526 _lmHead.wrappedValue = Linear(args.hiddenSize, args.vocabularySize, bias: false) 527 } 528 } 529 530 public func callAsFunction(_ inputs: MLXArray, cache: [KVCache]?) -> MLXArray { 531 var out = model(inputs, cache: cache) 532 if let lmHead { 533 out = lmHead(out) 534 } else { 535 out = model.embedTokens.asLinear(out) 536 } 537 return out 538 } 539 540 public func newCache(parameters: GenerateParameters?) -> [KVCache] { 541 // 2026-05-01: honor `parameters.maxKVSize` for attention slots 542 // (parity with Qwen35 / NemotronH; Mamba layers ignore the bound). 543 return model.layers.map { layer in 544 if layer.isLinear { 545 return MambaCache() 546 } 547 if let maxKVSize = parameters?.maxKVSize { 548 return RotatingKVCache(maxSize: maxKVSize, keep: 4) 549 } 550 return KVCacheSimple() 551 } 552 } 553 554 public func makeCache() -> [KVCache] { 555 return newCache(parameters: nil) 556 } 557 558 public func sanitize(weights: [String: MLXArray]) -> [String: MLXArray] { 559 var sanitizedWeights = weights 560 561 if configuration.tieWordEmbeddings { 562 sanitizedWeights["lm_head.weight"] = nil 563 } 564 565 let mtpKeys = sanitizedWeights.keys.filter { $0.contains("mtp.") } 566 for key in mtpKeys { 567 sanitizedWeights[key] = nil 568 } 569 570 if sanitizedWeights["model.layers.0.mlp.experts.0.up_proj.weight"] == nil { 571 return sanitizedWeights 572 } 573 574 for l in 0 ..< configuration.hiddenLayers { 575 let prefix = "model.layers.\(l).mlp" 576 for n in ["up_proj", "down_proj", "gate_proj"] { 577 let key = "\(prefix).experts.0.\(n).weight" 578 if sanitizedWeights[key] != nil { 579 let toJoin = (0 ..< configuration.numExperts).map { e in 580 sanitizedWeights.removeValue( 581 forKey: "\(prefix).experts.\(e).\(n).weight")! 582 } 583 sanitizedWeights["\(prefix).switch_mlp.\(n).weight"] = MLX.stacked(toJoin) 584 } 585 } 586 } 587 588 let normSuffixes = [ 589 ".input_layernorm.weight", 590 ".post_attention_layernorm.weight", 591 "model.norm.weight", 592 ".q_norm.weight", 593 ".k_norm.weight", 594 ] 595 596 for key in Array(sanitizedWeights.keys) { 597 guard let value = sanitizedWeights[key] else { continue } 598 if key.contains("conv1d.weight") && value.dim(-1) != 1 { 599 sanitizedWeights[key] = value.movedAxis(source: 2, destination: 1) 600 continue 601 } 602 if normSuffixes.contains(where: { key.hasSuffix($0) }) && value.ndim == 1 { 603 sanitizedWeights[key] = value + MLXArray(1, dtype: value.dtype) 604 } 605 } 606 607 return sanitizedWeights 608 } 609} 610 611public struct Qwen3NextConfiguration: Codable, Sendable { 612 var modelType: String = "qwen3_next" 613 var hiddenSize: Int 614 var hiddenLayers: Int 615 var intermediateSize: Int 616 var attentionHeads: Int 617 var linearNumValueHeads: Int 618 var linearNumKeyHeads: Int 619 var linearKeyHeadDim: Int 620 var linearValueHeadDim: Int 621 var linearConvKernelDim: Int 622 var numExperts: Int 623 var numExpertsPerTok: Int 624 var decoderSparseStep: Int 625 var sharedExpertIntermediateSize: Int 626 var mlpOnlyLayers: [Int] 627 var moeIntermediateSize: Int 628 var rmsNormEps: Float 629 var vocabularySize: Int 630 var kvHeads: Int 631 var ropeTheta: Float 632 var partialRotaryFactor: Float 633 var maxPositionEmbeddings: Int 634 var normTopkProb: Bool 635 var tieWordEmbeddings: Bool 636 var attentionBias: Bool 637 var headDim: Int? 638 var ropeScaling: [String: StringOrNumber]? 639 var fullAttentionInterval: Int 640 641 enum CodingKeys: String, CodingKey { 642 case modelType = "model_type" 643 case hiddenSize = "hidden_size" 644 case hiddenLayers = "num_hidden_layers" 645 case intermediateSize = "intermediate_size" 646 case attentionHeads = "num_attention_heads" 647 case linearNumValueHeads = "linear_num_value_heads" 648 case linearNumKeyHeads = "linear_num_key_heads" 649 case linearKeyHeadDim = "linear_key_head_dim" 650 case linearValueHeadDim = "linear_value_head_dim" 651 case linearConvKernelDim = "linear_conv_kernel_dim" 652 case numExperts = "num_experts" 653 case numExpertsPerTok = "num_experts_per_tok" 654 case decoderSparseStep = "decoder_sparse_step" 655 case sharedExpertIntermediateSize = "shared_expert_intermediate_size" 656 case mlpOnlyLayers = "mlp_only_layers" 657 case moeIntermediateSize = "moe_intermediate_size" 658 case rmsNormEps = "rms_norm_eps" 659 case vocabularySize = "vocab_size" 660 case kvHeads = "num_key_value_heads" 661 case ropeTheta = "rope_theta" 662 case partialRotaryFactor = "partial_rotary_factor" 663 case maxPositionEmbeddings = "max_position_embeddings" 664 case normTopkProb = "norm_topk_prob" 665 case tieWordEmbeddings = "tie_word_embeddings" 666 case attentionBias = "attention_bias" 667 case headDim = "head_dim" 668 case ropeScaling = "rope_scaling" 669 case fullAttentionInterval = "full_attention_interval" 670 } 671 672 public init(from decoder: Decoder) throws { 673 let container: KeyedDecodingContainer<Qwen3NextConfiguration.CodingKeys> = 674 try decoder.container(keyedBy: Qwen3NextConfiguration.CodingKeys.self) 675 676 self.modelType = 677 try container.decodeIfPresent(String.self, forKey: .modelType) ?? "qwen3_next" 678 self.hiddenSize = try container.decode(Int.self, forKey: .hiddenSize) 679 self.hiddenLayers = try container.decode(Int.self, forKey: .hiddenLayers) 680 self.intermediateSize = try container.decode(Int.self, forKey: .intermediateSize) 681 self.attentionHeads = try container.decode(Int.self, forKey: .attentionHeads) 682 self.linearNumValueHeads = try container.decode(Int.self, forKey: .linearNumValueHeads) 683 self.linearNumKeyHeads = try container.decode(Int.self, forKey: .linearNumKeyHeads) 684 self.linearKeyHeadDim = try container.decode(Int.self, forKey: .linearKeyHeadDim) 685 self.linearValueHeadDim = try container.decode(Int.self, forKey: .linearValueHeadDim) 686 self.linearConvKernelDim = try container.decode(Int.self, forKey: .linearConvKernelDim) 687 self.numExperts = try container.decode(Int.self, forKey: .numExperts) 688 self.numExpertsPerTok = try container.decode(Int.self, forKey: .numExpertsPerTok) 689 self.decoderSparseStep = try container.decode(Int.self, forKey: .decoderSparseStep) 690 self.sharedExpertIntermediateSize = try container.decode( 691 Int.self, forKey: .sharedExpertIntermediateSize) 692 self.mlpOnlyLayers = try container.decodeIfPresent([Int].self, forKey: .mlpOnlyLayers) ?? [] 693 self.moeIntermediateSize = try container.decode(Int.self, forKey: .moeIntermediateSize) 694 self.rmsNormEps = try container.decode(Float.self, forKey: .rmsNormEps) 695 self.vocabularySize = try container.decode(Int.self, forKey: .vocabularySize) 696 self.kvHeads = try container.decode(Int.self, forKey: .kvHeads) 697 self.ropeTheta = try container.decodeIfPresent(Float.self, forKey: .ropeTheta) ?? 1_000_000 698 self.partialRotaryFactor = 699 try container.decodeIfPresent(Float.self, forKey: .partialRotaryFactor) ?? 1.0 700 self.maxPositionEmbeddings = 701 try container.decodeIfPresent(Int.self, forKey: .maxPositionEmbeddings) ?? 32768 702 self.normTopkProb = try container.decodeIfPresent(Bool.self, forKey: .normTopkProb) ?? false 703 self.tieWordEmbeddings = 704 try container.decodeIfPresent(Bool.self, forKey: .tieWordEmbeddings) ?? false 705 self.attentionBias = 706 try container.decodeIfPresent(Bool.self, forKey: .attentionBias) ?? false 707 self.headDim = try container.decodeIfPresent(Int.self, forKey: .headDim) 708 self.ropeScaling = try container.decodeIfPresent( 709 [String: StringOrNumber].self, forKey: .ropeScaling) 710 self.fullAttentionInterval = 711 try container.decodeIfPresent(Int.self, forKey: .fullAttentionInterval) ?? 4 712 } 713} 714 715// MARK: - LoRA 716 717extension Qwen3NextModel: LoRAModel { 718 public var loraLayers: [Module] { 719 model.layers 720 } 721}