Fork for thermals request add-json-schema-dpeq
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}