1use std::{cell::RefCell, collections::HashMap};
2
3use crate::{
4 array,
5 error::Exception,
6 module::{Module, Param},
7 ops::{
8 arange, concatenate, exp,
9 indexing::{NewAxis, TryIndexOp},
10 log,
11 },
12 Array, Dtype,
13};
14use mlx_internal_macros::{generate_builder, Buildable, Builder};
15use mlx_macros::ModuleParameters;
16
17pub type Rope = RotaryPositionalEncoding;
19
20pub type RopeBuilder = RotaryPositionalEncodingBuilder;
22
23generate_builder! {
24 #[derive(Debug, Clone, ModuleParameters, Buildable)]
33 #[module(root = crate)]
34 #[buildable(root = crate)]
35 #[builder(root = crate)]
36 pub struct RotaryPositionalEncoding {
37 pub dimensions: i32,
40
41 #[builder(optional, default = RotaryPositionalEncoding::DEFAULT_TRADITIONAL)]
44 pub traditional: bool,
45
46 #[builder(optional, default = RotaryPositionalEncoding::DEFAULT_BASE)]
49 pub base: f32,
50
51 #[builder(optional, default = RotaryPositionalEncoding::DEFAULT_SCALE)]
53 pub scale: f32,
54 }
55}
56
57impl RotaryPositionalEncoding {
58 pub const DEFAULT_TRADITIONAL: bool = false;
60
61 pub const DEFAULT_BASE: f32 = 10_000.0;
63
64 pub const DEFAULT_SCALE: f32 = 1.0;
66}
67
68generate_builder! {
69 #[derive(Debug, Buildable, Clone)]
71 #[buildable(root = crate)]
72 #[builder(root = crate)]
73 pub struct RopeInput<'a> {
74 pub x: &'a Array,
76
77 #[builder(optional, default = RopeInput::DEFAULT_OFFSET)]
79 pub offset: i32,
80 }
81}
82
83impl RopeInput<'_> {
84 pub const DEFAULT_OFFSET: i32 = 0;
86}
87
88impl<'a> From<&'a Array> for RopeInput<'a> {
89 fn from(x: &'a Array) -> Self {
90 RopeInput {
91 x,
92 offset: Self::DEFAULT_OFFSET,
93 }
94 }
95}
96
97impl<'a> From<(&'a Array,)> for RopeInput<'a> {
98 fn from((x,): (&'a Array,)) -> Self {
99 RopeInput {
100 x,
101 offset: Self::DEFAULT_OFFSET,
102 }
103 }
104}
105
106impl<'a> From<(&'a Array, i32)> for RopeInput<'a> {
107 fn from((x, offset): (&'a Array, i32)) -> Self {
108 RopeInput { x, offset }
109 }
110}
111
112impl<'a, Input> Module<Input> for RotaryPositionalEncoding
113where
114 Input: Into<RopeInput<'a>>,
115{
116 type Error = Exception;
117
118 type Output = Array;
119
120 fn forward(&mut self, input: Input) -> Result<Self::Output, Self::Error> {
121 let RopeInput { x, offset } = input.into();
122 crate::fast::rope(
123 x,
124 self.dimensions,
125 self.traditional,
126 self.base,
127 self.scale,
128 offset,
129 None,
130 )
131 }
132
133 fn training_mode(&mut self, _mode: bool) {}
134}
135
136pub type Sinpe = SinusoidalPositionalEncoding;
138
139pub type SinpeBuilder = SinusoidalPositionalEncodingBuilder;
141
142#[derive(Debug, Clone, ModuleParameters, Buildable)]
147#[module(root = crate)]
148#[buildable(root = crate)]
149pub struct SinusoidalPositionalEncoding {
150 #[param]
151 sigmas: Param<Array>,
152
153 pub scale: f32,
155
156 pub cosine_first: bool,
158}
159
160impl Sinpe {
161 pub const DEFAULT_COSINE_FIRST: bool = false;
163
164 pub const DEFAULT_MIN_FREQUENCY: f32 = 0.0001;
166
167 pub const DEFAULT_MAX_FREQUENCY: f32 = 1.0;
169
170 pub const DEFAULT_FULL_TURNS: bool = false;
172}
173
174#[derive(Debug, Clone, Builder)]
176#[builder(
177 root = crate,
178 build_with = build_sinpe,
179 err = Exception,
180)]
181pub struct SinusoidalPositionalEncodingBuilder {
182 dimensions: i32,
183
184 #[builder(optional, default = Sinpe::DEFAULT_MIN_FREQUENCY)]
185 min_frequency: f32,
186
187 #[builder(optional, default = Sinpe::DEFAULT_MAX_FREQUENCY)]
188 max_frequency: f32,
189
190 #[builder(optional, default = None)]
191 scale: Option<f32>,
192
193 #[builder(optional, default = Sinpe::DEFAULT_COSINE_FIRST)]
194 cosine_first: bool,
195
196 #[builder(optional, default = Sinpe::DEFAULT_FULL_TURNS)]
197 full_turns: bool,
198}
199
200fn build_sinpe(builder: SinpeBuilder) -> Result<SinusoidalPositionalEncoding, Exception> {
201 let SinpeBuilder {
202 dimensions,
203 min_frequency,
204 max_frequency,
205 scale,
206 cosine_first,
207 full_turns,
208 } = builder;
209
210 let half_dim = dimensions / 2;
211 let one_zero = array!(1.0)
212 .subtract(Array::from_iter(0..half_dim, &[half_dim]).divide(array!(half_dim - 1))?)?;
213 let min_frequency = log(array!(min_frequency))?;
214 let max_frequency = log(array!(max_frequency))?;
215
216 let mut sigmas = exp(&one_zero * (&max_frequency - &min_frequency) + &min_frequency)?;
218 if full_turns {
219 sigmas *= array!(2.0 * std::f32::consts::PI);
221 }
222
223 let scale = scale.unwrap_or_else(|| (2.0 / dimensions as f32).sqrt());
224
225 Ok(SinusoidalPositionalEncoding {
226 sigmas: Param::new(sigmas),
227 scale,
228 cosine_first,
229 })
230}
231
232impl Module<&Array> for Sinpe {
233 type Error = Exception;
234 type Output = Array;
235
236 fn forward(&mut self, x: &Array) -> Result<Self::Output, Self::Error> {
237 let mut y = x
238 .expand_dims_axes(&[-1])
239 .and_then(|x| x.multiply(&self.sigmas))?;
240
241 let cosy = y.cos()?;
242 let siny = y.sin()?;
243
244 if self.cosine_first {
245 y = concatenate(&[cosy, siny], -1)?;
246 } else {
247 y = concatenate(&[siny, cosy], -1)?;
248 }
249
250 if self.scale != 1.0 {
251 y *= self.scale;
253 }
254
255 Ok(y)
256 }
257
258 fn training_mode(&mut self, _mode: bool) {}
259}
260
261#[derive(Debug, Clone, Hash, PartialEq, Eq)]
262struct AlibiKey {
263 q_seq_len: i32,
264 k_seq_len: i32,
265 num_heads: i32,
266 offset: i32,
267 dtype: Dtype,
268}
269
270thread_local! {
271 static ALIBI_CACHE: RefCell<HashMap<AlibiKey, Array>> = RefCell::new(HashMap::new());
272}
273
274#[derive(Debug, Clone, ModuleParameters)]
276#[module(root = crate)]
277pub struct Alibi;
278
279impl Alibi {
280 fn slope(num_heads: i32) -> Result<Array, Exception> {
281 let x = 2.0_f32.powi(8).powf(1.0 / num_heads as f32);
282 array!(x)
283 .power(&arange::<_, f32>(1, num_heads + 1, None)?)?
284 .expand_dims_axes(&[-1, -2])
285 }
286
287 fn matrix(key: AlibiKey) -> Result<Array, Exception> {
288 if let Some(value) = ALIBI_CACHE.with(|cache| cache.borrow().get(&key).cloned()) {
289 return Ok(value);
290 }
291
292 let x1 = arange::<_, f32>(key.offset, key.q_seq_len, None)?;
293 let x2 = arange::<_, f32>(0, key.k_seq_len, None)?;
294 let distance_matrix = x1
295 .try_index((.., NewAxis))?
296 .subtract(x2.try_index((NewAxis, ..))?)?
297 .expand_dims_axes(&[0, 1])?
298 .abs()?
299 .negative()?;
300
301 let slope = Self::slope(key.num_heads)?;
302 let mask = distance_matrix.multiply(&slope)?.as_dtype(key.dtype)?;
303
304 ALIBI_CACHE.with(|cache| {
305 cache.borrow_mut().insert(key, mask.clone());
306 });
307
308 Ok(mask)
309 }
310}
311
312generate_builder! {
313 #[derive(Debug, Clone, Buildable)]
315 #[buildable(root = crate)]
316 #[builder(root = crate)]
317 pub struct AlibiInput<'a> {
318 pub attention_scores: &'a Array,
320
321 #[builder(optional, default = AlibiInput::DEFAULT_OFFSET)]
323 pub offset: i32,
324
325 #[builder(optional, default = None)]
327 pub mask: Option<&'a Array>,
328 }
329}
330
331impl AlibiInput<'_> {
332 pub const DEFAULT_OFFSET: i32 = 0;
334}
335
336impl<'a> From<&'a Array> for AlibiInput<'a> {
337 fn from(attention_scores: &'a Array) -> Self {
338 AlibiInput {
339 attention_scores,
340 offset: Self::DEFAULT_OFFSET,
341 mask: None,
342 }
343 }
344}
345
346impl<'a> From<(&'a Array,)> for AlibiInput<'a> {
347 fn from((attention_scores,): (&'a Array,)) -> Self {
348 AlibiInput {
349 attention_scores,
350 offset: Self::DEFAULT_OFFSET,
351 mask: None,
352 }
353 }
354}
355
356impl<'a> From<(&'a Array, i32)> for AlibiInput<'a> {
357 fn from((attention_scores, offset): (&'a Array, i32)) -> Self {
358 AlibiInput {
359 attention_scores,
360 offset,
361 mask: None,
362 }
363 }
364}
365
366impl<'a> From<(&'a Array, i32, &'a Array)> for AlibiInput<'a> {
367 fn from((attention_scores, offset, mask): (&'a Array, i32, &'a Array)) -> Self {
368 AlibiInput {
369 attention_scores,
370 offset,
371 mask: Some(mask),
372 }
373 }
374}
375
376impl<'a> From<(&'a Array, i32, Option<&'a Array>)> for AlibiInput<'a> {
377 fn from((attention_scores, offset, mask): (&'a Array, i32, Option<&'a Array>)) -> Self {
378 AlibiInput {
379 attention_scores,
380 offset,
381 mask,
382 }
383 }
384}
385
386impl<'a, Input> Module<Input> for Alibi
387where
388 Input: Into<AlibiInput<'a>>,
389{
390 type Output = Array;
391 type Error = Exception;
392
393 fn forward(&mut self, input: Input) -> Result<Self::Output, Self::Error> {
394 let AlibiInput {
395 attention_scores,
396 offset,
397 mask,
398 } = input.into();
399
400 let key = AlibiKey {
401 q_seq_len: attention_scores.dim(-2) + offset,
402 k_seq_len: attention_scores.dim(-1),
403 num_heads: attention_scores.dim(1),
404 offset,
405 dtype: attention_scores.dtype(),
406 };
407
408 let mut alibi_mask = Self::matrix(key)?;
409 if let Some(mask) = mask {
410 alibi_mask = alibi_mask.add(mask)?;
411 }
412
413 attention_scores.add(alibi_mask)
414 }
415
416 fn training_mode(&mut self, _mode: bool) {}
417}
418
419#[allow(clippy::excessive_precision)]
420#[cfg(test)]
421mod tests {
422 use crate::{module::Module, nn::AlibiInput, random::uniform, Dtype};
423 use float_eq::assert_float_eq;
424
425 use crate::nn::Rope;
426
427 #[test]
430 fn test_rope() {
431 crate::random::seed(71).unwrap();
432 let a = uniform::<_, f32>(0, 1, &[2, 8, 16], None).unwrap();
433 assert_eq!(a.shape(), &[2, 8, 16]);
434 assert_eq!(a.dtype(), Dtype::Float32);
435 assert_float_eq!(
436 a.mean(None).unwrap().item_exact::<f32>(),
437 0.5082664489746094,
438 abs <= 0.010165328979492188
439 );
440 assert_float_eq!(
441 a.sum(None).unwrap().item_exact::<f32>(),
442 130.1162109375,
443 abs <= 2.60232421875
444 );
445
446 let mut rope = Rope::new(8);
447 let result = rope.forward(&a).unwrap();
448 assert_eq!(result.shape(), &[2, 8, 16]);
449 assert_eq!(result.dtype(), Dtype::Float32);
450 assert_float_eq!(
451 result.mean(None).unwrap().item_exact::<f32>(),
452 0.4562537670135498,
453 abs <= 0.009125075340270997
454 );
455 assert_float_eq!(
456 result.sum(None).unwrap().item_exact::<f32>(),
457 116.80096435546875,
458 abs <= 2.3360192871093752
459 );
460 }
461
462 #[test]
463 fn test_rope_single_position_matches_broadcast_head() {
464 let head = [0.1_f32, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8];
465 let mut heads = Vec::with_capacity(32);
466 for _ in 0..4 {
467 heads.extend_from_slice(&head);
468 }
469 let input = crate::Array::from_slice(&heads, &[1, 4, 1, 8]);
470 let single_head = crate::Array::from_slice(&head, &[1, 1, 1, 8]);
471 let mut rope = Rope::new(8);
472
473 let output = rope.forward((&input, 1)).unwrap();
474 let single_output = rope.forward((&single_head, 1)).unwrap();
475 let expected = crate::ops::broadcast_to(&single_output, &[1, 4, 1, 8]).unwrap();
476 let max_head_difference = output
477 .subtract(&expected)
478 .unwrap()
479 .abs()
480 .unwrap()
481 .max(None)
482 .unwrap()
483 .item_exact::<f32>();
484
485 assert_eq!(output.shape(), &[1, 4, 1, 8]);
486 assert!(max_head_difference < 1e-5);
487
488 let rotation = single_output
489 .subtract(&single_head)
490 .unwrap()
491 .abs()
492 .unwrap()
493 .max(None)
494 .unwrap()
495 .item_exact::<f32>();
496 assert!(rotation > 1e-4);
497 }
498
499 #[test]
502 fn test_sinpe() {
503 crate::random::seed(226).unwrap();
504 let a = uniform::<_, f32>(0, 1, &[2, 8, 16], None).unwrap();
505 assert_eq!(a.shape(), &[2, 8, 16]);
506 assert_eq!(a.dtype(), Dtype::Float32);
507 assert_float_eq!(
508 a.mean(None).unwrap().item_exact::<f32>(),
509 0.5026599168777466,
510 abs <= 0.010053198337554931
511 );
512 assert_float_eq!(
513 a.sum(None).unwrap().item_exact::<f32>(),
514 128.68093872070312,
515 abs <= 2.5736187744140624
516 );
517
518 let mut sinpe = crate::nn::Sinpe::new(8).unwrap();
519 let result = sinpe.forward(&a).unwrap();
520 assert_eq!(result.shape(), &[2, 8, 16, 8]);
521 assert_eq!(result.dtype(), Dtype::Float32);
522 assert_float_eq!(
523 result.mean(None).unwrap().item_exact::<f32>(),
524 0.2705308198928833,
525 abs <= 0.005410616397857666
526 );
527 assert_float_eq!(
528 result.sum(None).unwrap().item_exact::<f32>(),
529 554.047119140625,
530 abs <= 11.0809423828125
531 );
532 }
533
534 #[test]
537 fn test_alibi() {
538 let mut alibi = crate::nn::Alibi;
539 let shape = [1, 8, 20, 20];
540 let x = uniform::<_, f32>(0, 1, &shape, None).unwrap();
541 let input = AlibiInput::from(&x);
542 let y = alibi.forward(input).unwrap();
543 assert_eq!(y.shape(), shape);
544 assert_eq!(y.dtype(), Dtype::Float32);
545
546 let x2 = x.as_dtype(Dtype::Float16).unwrap();
547 let input = AlibiInput::from(&x2);
548 let y = alibi.forward(input).unwrap();
549 assert_eq!(y.dtype(), Dtype::Float16);
550 }
551}