Skip to main content

laddu_physics/math/
wigner.rs

1//! Efficient helper methods to calculate Clebsch-Gordon coefficients and Wigner 3j-symbols.
2//!
3//! Ported from <https://github.com/0382/WignerSymbol> to Rust by N. D. Hoffman, 2026.
4
5use laddu_expr::{Expr, cis};
6use serde::{Deserialize, Serialize};
7
8use self::utils::{binomial, const_imax, const_umin};
9use crate::{
10    LadduPhysicsError, LadduPhysicsResult,
11    quantum::{J, M},
12};
13
14mod utils {
15    const MAX_BINOMIAL: u64 = 67;
16    const SIZE: usize = table_size(MAX_BINOMIAL);
17
18    const fn table_size(n: u64) -> usize {
19        let x = n / 2 + 1;
20        (x * (x + (n & 1))) as usize
21    }
22    const fn index(n: u64, k: u64) -> usize {
23        let x = n / 2 + 1;
24        (x * (x - (1 - (n & 1))) + k) as usize
25    }
26    pub(crate) const fn const_umin(a: u64, b: u64) -> u64 {
27        if a < b { a } else { b }
28    }
29    pub(crate) const fn const_imax(a: i64, b: i64) -> i64 {
30        if a > b { a } else { b }
31    }
32    const fn build_binomial_table() -> [u64; SIZE] {
33        let mut data = [0u64; SIZE];
34        data[0] = 1;
35        let mut n = 1;
36        while n <= MAX_BINOMIAL {
37            let mut k = 0;
38            while k <= n / 2 {
39                let value = if k == 0 {
40                    1
41                } else {
42                    let nm1 = n - 1;
43                    let a_k = const_umin(k, nm1 - k);
44                    let km1 = k - 1;
45                    let b_k = const_umin(km1, nm1 - km1);
46                    data[index(nm1, a_k)] + data[index(nm1, b_k)]
47                };
48                data[index(n, k)] = value;
49                k += 1;
50            }
51            n += 1;
52        }
53        data
54    }
55    static BINOMIAL_TABLE: [u64; SIZE] = build_binomial_table();
56
57    /// Compute the binomial coefficient C(n, k).
58    ///
59    /// # Note
60    /// For n > 67, this value would exceed the u64 max, so it will instead return 0.
61    #[inline]
62    pub(crate) const fn binomial(n: u64, k: u64) -> u64 {
63        if n > MAX_BINOMIAL || k > n {
64            return 0;
65        }
66        let k = const_umin(k, n - k);
67        BINOMIAL_TABLE[index(n, k)]
68    }
69
70    #[cfg(test)]
71    mod tests {
72        use super::binomial;
73        #[test]
74        fn test_binomial() {
75            assert_eq!(binomial(0, 0), 1);
76            assert_eq!(binomial(5, 0), 1);
77            assert_eq!(binomial(5, 1), 5);
78            assert_eq!(binomial(5, 2), 10);
79            assert_eq!(binomial(5, 3), 10);
80            assert_eq!(binomial(5, 4), 5);
81            assert_eq!(binomial(5, 5), 1);
82            assert_eq!(binomial(67, 33), 14_226_520_737_620_288_370);
83            assert_eq!(binomial(67, 34), 14_226_520_737_620_288_370);
84            assert_eq!(binomial(68, 1), 0);
85            assert_eq!(binomial(10, 11), 0);
86        }
87    }
88}
89
90/// (-1)^x but efficient
91#[inline]
92const fn phase(x: u64) -> i64 {
93    1 - (2 * (x & 1) as i64)
94}
95
96/// true if both j and m are either both integers or both half-integers, false for mixed cases
97#[inline]
98const fn check_parity(dj: i64, dm: i64) -> bool {
99    (dj ^ dm) & 1 == 0
100}
101#[inline]
102const fn check_jm(dj: i64, dm: i64) -> bool {
103    check_parity(dj, dm) && (dm.abs() <= dj)
104}
105#[inline]
106const fn check_coupling(dj1: i64, dj2: i64, dj3: i64) -> bool {
107    (dj1 >= 0)
108        && (dj2 >= 0)
109        && (dj3 >= (dj1 - dj2).abs())
110        && check_parity(dj1 + dj2, dj3)
111        && (dj3 <= (dj1 + dj2))
112}
113
114/// Computes the Clebsch–Gordan coefficient $`\langle j_1 m_1; j_2 m_2 \mid j m\rangle`$.
115///
116/// # Parameters
117///
118/// All angular momenta and projections are passed as strongly typed quantum numbers.
119///
120/// # Returns
121///
122/// Returns the Clebsch–Gordan coefficient as `f64`.
123///
124/// The function returns `0.0` when any standard selection rule is violated:
125/// - `(j, m)` is not a valid angular-momentum projection pair
126/// - $`j_1`$, $`j_2`$, and $`j`$ do not satisfy the triangle relation
127/// - $`m_1 + m_2 \ne m`$
128///
129/// # Examples
130///
131/// ```rust
132/// # use laddu_physics::{j, m};
133/// # use laddu_physics::math::clebsch_gordan;
134/// let cg = clebsch_gordan(j!(1), m!(1), j!(1), m!(-1), j!(1), m!(0));
135/// assert_eq!(cg, f64::sqrt(1.0 / 2.0));
136/// ```
137pub fn clebsch_gordan(j1: J, m1: M, j2: J, m2: M, j: J, m: M) -> f64 {
138    clebsch_gordan_doubled(
139        j1.doubled() as u64,
140        j2.doubled() as u64,
141        j.doubled() as u64,
142        m1.doubled() as i64,
143        m2.doubled() as i64,
144        m.doubled() as i64,
145    )
146}
147
148fn clebsch_gordan_doubled(dj1: u64, dj2: u64, dj3: u64, dm1: i64, dm2: i64, dm3: i64) -> f64 {
149    if !(check_jm(dj1 as i64, dm1) && check_jm(dj2 as i64, dm2) && check_jm(dj3 as i64, dm3)) {
150        return 0.0;
151    }
152    if !check_coupling(dj1 as i64, dj2 as i64, dj3 as i64) {
153        return 0.0;
154    }
155    if dm1 + dm2 != dm3 {
156        return 0.0;
157    }
158    if dm1 == 0 && dm2 == 0 && dm3 == 0 {
159        let j1 = dj1 / 2;
160        let j2 = dj2 / 2;
161        let j3 = dj3 / 2;
162        let j = j1 + j2 + j3;
163        let g = j / 2;
164        return phase(g - j3) as f64 * (binomial(g, j3) * binomial(j3, g - j1)) as f64
165            / ((binomial(j + 1, dj3 + 1) * binomial(dj3, j - dj1)) as f64).sqrt();
166    }
167    let j = (dj1 + dj2 + dj3) / 2;
168    let jm1 = j - dj1;
169    let jm2 = j - dj2;
170    let jm3 = j - dj3;
171    let j1mm1 = (dj1 as i64 - dm1) as u64 / 2;
172    let j2mm2 = (dj2 as i64 - dm2) as u64 / 2;
173    let j3mm3 = (dj3 as i64 - dm3) as u64 / 2;
174    let j2pm2 = (dj2 as i64 + dm2) as u64 / 2;
175    let a = ((binomial(dj1, jm2) * binomial(dj2, jm3)) as f64
176        / (binomial(j + 1, jm3)
177            * binomial(dj1, j1mm1)
178            * binomial(dj2, j2mm2)
179            * binomial(dj3, j3mm3)) as f64)
180        .sqrt();
181    let mut b: i64 = 0;
182    let k_min = const_imax(
183        0,
184        const_imax(j1mm1 as i64 - jm2 as i64, j2pm2 as i64 - jm1 as i64),
185    ) as u64;
186    let k_max = const_umin(jm3, const_umin(j1mm1, j2pm2));
187    for z in k_min..=k_max {
188        b = -b + (binomial(jm3, z) * binomial(jm2, j1mm1 - z) * binomial(jm1, j2pm2 - z)) as i64;
189    }
190    a * (phase(k_max) * b) as f64
191}
192
193/// Computes the Wigner 3-j symbol
194///
195/// $`\begin{pmatrix} j_1 & j_2 & j_3 \\ m_1 & m_2 & m_3 \end{pmatrix}`$
196///
197/// using the doubled-quantum-number convention.
198///
199/// # Parameters
200///
201/// All angular momenta and projections are passed as strongly typed quantum numbers.
202///
203/// # Returns
204///
205/// Returns the Wigner 3-j symbol as `f64`.
206///
207/// The function returns `0.0` when any selection rule fails:
208/// - `(j, m)` is invalid for any input pair
209/// - $`j_1`$, $`j_2`$, and $`j_3`$ violate the triangle condition
210/// - $`m_1 + m_2 + m_3 \ne 0`$
211///
212/// # Notes
213///
214/// - The implementation uses a finite alternating sum over binomial factors.
215///
216/// # Examples
217///
218/// Integer angular momentum:
219///
220/// ```rust
221/// # use laddu_physics::{j, m};
222/// # use laddu_physics::math::wigner_3j;
223/// let w = wigner_3j(j!(1), m!(1), j!(1), m!(-1), j!(1), m!(0)); // (1 1 1; 1 -1 0)
224/// assert_eq!(w, f64::sqrt(1.0 / 6.0))
225/// ```
226///
227/// Half-integer angular momentum:
228///
229/// ```rust
230/// # use laddu_physics::{j, m};
231/// # use laddu_physics::math::wigner_3j;
232/// let w = wigner_3j(j!(1/2), m!(1/2), j!(1/2), m!(-1/2), j!(1), m!(0)); // (1/2 1/2 1; 1/2 -1/2 0)
233/// assert_eq!(w, f64::sqrt(1.0 / 6.0))
234/// ```
235///
236/// # Convention
237///
238/// This symbol is related to the Clebsch–Gordon coefficient by
239///
240/// $`\begin{pmatrix} j_1 & j_2 & j_3 \\ m_1 & m_2 & m_3 \end{pmatrix} = ((-1)^{j_1 - j_2 - m_3} / \sqrt{2j_3 + 1}) \langle j_1 m_1; j_2 m_2 | j_3, -m_3\rangle`$
241pub fn wigner_3j(j1: J, m1: M, j2: J, m2: M, j3: J, m3: M) -> f64 {
242    wigner_3j_doubled(
243        j1.doubled() as u64,
244        j2.doubled() as u64,
245        j3.doubled() as u64,
246        m1.doubled() as i64,
247        m2.doubled() as i64,
248        m3.doubled() as i64,
249    )
250}
251
252fn wigner_3j_doubled(dj1: u64, dj2: u64, dj3: u64, dm1: i64, dm2: i64, dm3: i64) -> f64 {
253    if !(check_jm(dj1 as i64, dm1) && check_jm(dj2 as i64, dm2) && check_jm(dj3 as i64, dm3)) {
254        return 0.0;
255    }
256    if !check_coupling(dj1 as i64, dj2 as i64, dj3 as i64) {
257        return 0.0;
258    }
259    if dm1 + dm2 + dm3 != 0 {
260        return 0.0;
261    }
262    let j = (dj1 + dj2 + dj3) / 2;
263    let jm1 = j - dj1;
264    let jm2 = j - dj2;
265    let jm3 = j - dj3;
266    let j1mm1 = (dj1 as i64 - dm1) as u64 / 2;
267    let j2mm2 = (dj2 as i64 - dm2) as u64 / 2;
268    let j3mm3 = (dj3 as i64 - dm3) as u64 / 2;
269    let j1pm1 = (dj1 as i64 + dm1) as u64 / 2;
270    let a = ((binomial(dj1, jm2) * binomial(dj2, jm1)) as f64
271        / ((j + 1)
272            * binomial(j, jm3)
273            * binomial(dj1, j1mm1)
274            * binomial(dj2, j2mm2)
275            * binomial(dj3, j3mm3)) as f64)
276        .sqrt();
277    let mut b: i64 = 0;
278    let k_min = const_imax(
279        0,
280        const_imax(j1pm1 as i64 - jm2 as i64, j2mm2 as i64 - jm1 as i64),
281    ) as u64;
282    let k_max = const_umin(jm3, const_umin(j1pm1, j2mm2));
283    for z in k_min..=k_max {
284        b = -b + (binomial(jm3, z) * binomial(jm2, j1pm1 - z) * binomial(jm1, j2mm2 - z)) as i64;
285    }
286    a * (phase(dj1 + (dj3 as i64 + dm3) as u64 / 2 + k_max) * b) as f64
287}
288
289/// Precomputed helper for Wigner rotation matrix elements.
290///
291/// Stores all coefficients needed for repeated evaluation of
292///
293/// $`d^j_{m' m}(\beta)`$ and $`D^j_{m' m}(\alpha,\beta,\gamma)`$.
294///
295/// # Definitions
296///
297/// $`D^j_{m' m}(\alpha,\beta,\gamma) = e^{-i m' \alpha} d^j_{m' m}(\beta) e^{-i m \gamma}`$
298///
299/// # Notes
300///
301/// Designed for reuse across many angle evaluations.
302#[derive(Copy, Clone, Serialize, Deserialize)]
303pub struct WignerDMatrix {
304    dj: i64,    // 2 * j
305    dmp: i64,   // 2 * m'
306    dm: i64,    // 2 * m
307    jpm: i64,   // j + m
308    jmmp: i64,  // j - m'
309    delta: i64, // m' - m
310    s_min: i64,
311    s_max: i64,
312}
313impl WignerDMatrix {
314    /// Constructs a Wigner small-$`d`$/full-$`D`$ matrix element helper for fixed
315    /// quantum numbers $`j`$, $`m'`$, and $`m`$.
316    ///
317    /// All angular momenta and projections are passed as strongly typed quantum numbers.
318    ///
319    /// The constructed value precomputes the combinatorial factors and
320    /// summation bounds needed for repeated evaluation of:
321    /// - the reduced Wigner matrix element $`d^j_{m' m}(\beta)`$, and
322    /// - the full Wigner matrix element $`D^j_{m' m}(\alpha,\beta,\gamma)`$.
323    ///
324    /// # Errors
325    ///
326    /// This method will return an error result if:
327    /// - `|m'| > j`
328    /// - `|m| > j`
329    /// - `j` and `m'` do not have matching integer/half-integer parity
330    /// - `j` and `m` do not have matching integer/half-integer parity
331    /// - the internally derived summation bounds are inconsistent
332    ///
333    /// # Notes
334    ///
335    /// The returned struct is intended for reuse when evaluating the same
336    /// matrix element for many angles. This avoids recomputing factorial-based
337    /// prefactors on every call.
338    ///
339    /// # Examples
340    ///
341    /// Integer angular momentum:
342    ///
343    /// ```rust
344    /// # use laddu_physics::{j, m};
345    /// # use laddu_physics::math::WignerDMatrix;
346    /// let w = WignerDMatrix::new(j!(1), m!(1), m!(0)).unwrap(); // j = 1, m' = 1, m = 0
347    /// ```
348    ///
349    /// Half-integer angular momentum:
350    ///
351    /// ```rust
352    /// # use laddu_physics::{j, m};
353    /// # use laddu_physics::math::WignerDMatrix;
354    /// let w = WignerDMatrix::new(j!(1/2), m!(1/2), m!(-1/2)).unwrap(); // j = 1/2, m' = 1/2, m = -1/2
355    /// ```
356    pub fn new(
357        j: impl TryInto<J>,
358        mp: impl TryInto<M>,
359        m: impl TryInto<M>,
360    ) -> LadduPhysicsResult<Self> {
361        let j = j
362            .try_into()
363            .map_err(|_| LadduPhysicsError::ConversionError("J"))?;
364        let mp = mp
365            .try_into()
366            .map_err(|_| LadduPhysicsError::ConversionError("M"))?;
367        let m = m
368            .try_into()
369            .map_err(|_| LadduPhysicsError::ConversionError("M"))?;
370        Self::new_doubled(j.doubled() as u64, mp.doubled() as i64, m.doubled() as i64)
371    }
372
373    fn new_doubled(dj: u64, dmp: i64, dm: i64) -> LadduPhysicsResult<Self> {
374        let dj = dj as i64;
375        if dmp.abs() > dj {
376            return Err(LadduPhysicsError::invalid_relation(format!(
377                "|m'| <= j, got 2*j = {dj}, 2*m' = {dmp}"
378            )));
379        }
380        if dm.abs() > dj {
381            return Err(LadduPhysicsError::invalid_relation(format!(
382                "|m| <= j, got 2*j = {dj}, 2*m = {dm}"
383            )));
384        }
385        if !check_parity(dj, dmp) {
386            return Err(LadduPhysicsError::invalid_relation(format!(
387                "j and m' must have the same integer/half-integer parity, got 2*j = {dj}, 2*m' = {dmp}"
388            )));
389        }
390        if !check_parity(dj, dm) {
391            return Err(LadduPhysicsError::invalid_relation(format!(
392                "j and m must have the same integer/half-integer parity, got 2*j = {dj}, 2*m = {dm}"
393            )));
394        }
395        let jmmp = (dj - dmp) / 2;
396        let jpm = (dj + dm) / 2;
397        let delta = (dmp - dm) / 2;
398        let s_min = 0.max(-delta);
399        let s_max = jpm.min(jmmp);
400        assert!(
401            s_min <= s_max,
402            "summation bounds are incorrect (this shouldn't happen)!"
403        );
404        Ok(Self {
405            dj,
406            dmp,
407            dm,
408            jpm,
409            jmmp,
410            delta,
411            s_min,
412            s_max,
413        })
414    }
415    /// Evaluates the reduced Wigner small-$`d`$ matrix element $`d^j_{m' m}(\beta)`$.
416    ///
417    /// The quantum numbers $`j`$, $`m'`$, and $`m`$ are those fixed when the
418    /// [`WignerDMatrix`] was constructed.
419    ///
420    /// # Parameters
421    ///
422    /// - `beta`: expression for the middle Euler angle $`\beta`$, in radians
423    ///
424    /// # Returns
425    ///
426    /// Returns an expression graph for the real-valued reduced Wigner matrix
427    /// element $`d^j_{m' m}(\beta)`$.
428    ///
429    /// # Notes
430    ///
431    /// This method builds the standard finite sum in powers of
432    /// $`\cos(\beta/2)`$ and $`\sin(\beta/2)`$.
433    ///
434    /// # Examples
435    ///
436    /// ```rust
437    /// # use laddu_physics::math::WignerDMatrix;
438    /// # use laddu_expr::event_scalar;
439    /// # use laddu_physics::{j, m};
440    /// let w = WignerDMatrix::new(j!(1), m!(1), m!(0)).unwrap(); // j = 1, m' = 1, m = 0
441    /// let expr = w.d(event_scalar("beta"));
442    /// ```
443    pub fn d(&self, beta: impl Into<Expr>) -> Expr {
444        let beta = beta.into();
445        let half_beta = 0.5 * beta;
446        let ch = half_beta.cos();
447        let sh = half_beta.sin();
448        let mut sum: Expr = 0.0.into();
449
450        for term in self.small_d_terms() {
451            let mut expr: Expr = term.coefficient.into();
452            if term.cos_power != 0 {
453                expr *= ch.powi(term.cos_power);
454            }
455            if term.sin_power != 0 {
456                expr *= sh.powi(term.sin_power);
457            }
458            sum += expr;
459        }
460
461        sum
462    }
463
464    /// Evaluates the full Wigner $`D`$ matrix element $`D^j_{m' m}(\alpha,\beta,\gamma)`$.
465    ///
466    /// The implemented convention is
467    ///
468    /// $`D^j_{m' m}(\alpha,\beta,\gamma) = e^{-\imath m' \alpha} d^j_{m' m}(\beta) e^{-\imath m \gamma}`$,
469    ///
470    /// # Parameters
471    ///
472    /// - `alpha`: expression for the first Euler angle $`\alpha`$, in radians
473    /// - `beta`: expression for the middle Euler angle $`\beta`$, in radians
474    /// - `gamma`: expression for the third Euler angle $`\gamma`$, in radians
475    ///
476    /// # Returns
477    ///
478    /// Returns an expression graph for the complex Wigner $`D`$ matrix element.
479    ///
480    /// # Notes
481    ///
482    /// Since $`d^j_{m' m}(\beta)`$ is real for real $`\beta`$, the complex phase comes
483    /// entirely from the $`\alpha`$ and $`\gamma`$ dependence.
484    ///
485    /// # Examples
486    ///
487    /// ```rust
488    /// # use laddu_physics::math::WignerDMatrix;
489    /// # use laddu_expr::event_scalar;
490    /// # use laddu_physics::{j, m};
491    /// let w = WignerDMatrix::new(j!(1), m!(1), m!(0)).unwrap(); // j = 1, m' = 1, m = 0
492    /// let expr = w.D(event_scalar("alpha"), event_scalar("beta"), event_scalar("gamma"));
493    /// ```
494    #[allow(non_snake_case)]
495    pub fn D(&self, alpha: impl Into<Expr>, beta: impl Into<Expr>, gamma: impl Into<Expr>) -> Expr {
496        let alpha = alpha.into();
497        let gamma = gamma.into();
498        let phase = -0.5 * (self.dmp as f64 * alpha + self.dm as f64 * gamma);
499        cis(phase) * self.d(beta)
500    }
501
502    fn small_d_terms(&self) -> Vec<WignerDTerm> {
503        let j_plus_mp = (self.dj + self.dmp) / 2;
504        let j_minus_m = (self.dj - self.dm) / 2;
505        let mut ln_factorial = vec![0.0; self.dj as usize + 1];
506        for i in 1..=self.dj as usize {
507            ln_factorial[i] = ln_factorial[i - 1] + (i as f64).ln();
508        }
509        let ln_prefactor = 0.5
510            * (ln_factorial[j_plus_mp as usize]
511                + ln_factorial[self.jmmp as usize]
512                + ln_factorial[self.jpm as usize]
513                + ln_factorial[j_minus_m as usize]);
514
515        (self.s_min..=self.s_max)
516            .map(|s| {
517                let denom_ln = ln_factorial[(self.jpm - s) as usize]
518                    + ln_factorial[s as usize]
519                    + ln_factorial[(self.delta + s) as usize]
520                    + ln_factorial[(self.jmmp - s) as usize];
521                let sign = if ((s + self.delta) & 1) == 0 {
522                    1.0
523                } else {
524                    -1.0
525                };
526                WignerDTerm {
527                    coefficient: sign * (ln_prefactor - denom_ln).exp(),
528                    cos_power: (self.dj - self.delta - 2 * s) as i32,
529                    sin_power: (self.delta + 2 * s) as i32,
530                }
531            })
532            .collect()
533    }
534}
535
536#[derive(Copy, Clone, Debug, PartialEq)]
537struct WignerDTerm {
538    coefficient: f64,
539    cos_power: i32,
540    sin_power: i32,
541}
542
543#[cfg(test)]
544mod tests {
545    use std::f64::consts::{FRAC_1_SQRT_2, FRAC_PI_2};
546
547    use approx::assert_relative_eq;
548    use laddu_compile::CompiledModel;
549    use laddu_runtime::CpuBackend;
550    use num::complex::Complex64;
551
552    use super::*;
553    use crate::{j, m};
554
555    fn evaluate(expr: Expr) -> Complex64 {
556        let model = CompiledModel::from_expr(&expr).unwrap();
557        let params = model.params().default_values();
558        CpuBackend.prepare(&model).evaluate(&params).unwrap()
559    }
560
561    fn assert_complex_relative_eq(actual: Complex64, expected: Complex64) {
562        assert_relative_eq!(actual.re, expected.re);
563        assert_relative_eq!(actual.im, expected.im);
564    }
565
566    #[test]
567    fn test_phase() {
568        assert_eq!(phase(0), 1);
569        assert_eq!(phase(1), -1);
570        assert_eq!(phase(2), 1);
571        assert_eq!(phase(3), -1);
572    }
573
574    #[test]
575    fn singlet_triplet_for_two_spin_half() {
576        // <1/2,1/2; 1/2,1/2 | 1,1> = 1
577        assert_relative_eq!(clebsch_gordan_doubled(1, 1, 2, 1, 1, 2), 1.0);
578
579        // <1/2,1/2; 1/2,-1/2 | 1,0> = 1/sqrt(2)
580        assert_relative_eq!(clebsch_gordan_doubled(1, 1, 2, 1, -1, 0), FRAC_1_SQRT_2);
581
582        // <1/2,-1/2; 1/2,1/2 | 1,0> = 1/sqrt(2)
583        assert_relative_eq!(clebsch_gordan_doubled(1, 1, 2, -1, 1, 0), FRAC_1_SQRT_2);
584
585        // <1/2,1/2; 1/2,-1/2 | 0,0> = 1/sqrt(2)
586        assert_relative_eq!(clebsch_gordan_doubled(1, 1, 0, 1, -1, 0), FRAC_1_SQRT_2);
587
588        // <1/2,-1/2; 1/2,1/2 | 0,0> = -1/sqrt(2)
589        assert_relative_eq!(clebsch_gordan_doubled(1, 1, 0, -1, 1, 0), -FRAC_1_SQRT_2);
590    }
591
592    #[test]
593    fn typed_clebsch_gordan_matches_doubled_helper() {
594        assert_relative_eq!(
595            clebsch_gordan(j!(1 / 2), m!(1 / 2), j!(1 / 2), m!(-1 / 2), j!(1), m!(0)),
596            clebsch_gordan_doubled(1, 1, 2, 1, -1, 0)
597        );
598        assert_eq!(
599            clebsch_gordan(j!(1 / 2), m!(1 / 2), j!(1 / 2), m!(1 / 2), j!(1), m!(0)),
600            0.0
601        );
602    }
603
604    #[test]
605    fn highest_weight_state_is_one() {
606        // <1,1; 1,1 | 2,2> = 1
607        // doubled notation: j1=j2=1 -> dj1=dj2=2, m1=m2=1 -> dm1=dm2=2
608        assert_relative_eq!(clebsch_gordan_doubled(2, 2, 4, 2, 2, 4), 1.0);
609    }
610
611    #[test]
612    fn known_spin_one_couplings() {
613        // <1,1; 1,0 | 2,1> = 1/sqrt(2)
614        assert_relative_eq!(clebsch_gordan_doubled(2, 2, 4, 2, 0, 2), FRAC_1_SQRT_2);
615
616        // <1,0; 1,0 | 2,0> = sqrt(2/3)
617        assert_relative_eq!(
618            clebsch_gordan_doubled(2, 2, 4, 0, 0, 0),
619            (2.0 / 3.0_f64).sqrt()
620        );
621
622        // <1,0; 1,0 | 0,0> = -1/sqrt(3)
623        assert_relative_eq!(
624            clebsch_gordan_doubled(2, 2, 0, 0, 0, 0),
625            -1.0 / 3.0_f64.sqrt()
626        );
627    }
628
629    #[test]
630    fn zero_when_m_sum_fails() {
631        // dm1 + dm2 != dm3
632        assert_eq!(clebsch_gordan_doubled(1, 1, 2, 1, 1, 0), 0.0);
633    }
634
635    #[test]
636    fn zero_when_triangle_rule_fails() {
637        // 1/2 + 1/2 cannot couple to j=2
638        assert_eq!(clebsch_gordan_doubled(1, 1, 4, 1, 1, 2), 0.0);
639    }
640
641    #[test]
642    fn zero_when_m_out_of_range() {
643        // For dj=1 (j=1/2), dm must be ±1 only
644        assert_eq!(clebsch_gordan_doubled(1, 1, 2, 3, -1, 2), 0.0);
645    }
646
647    #[test]
648    fn normalization_for_fixed_jm() {
649        // For j1=j2=1/2 and total J=1, M=0:
650        // |<+,-|1,0>|^2 + |<-,+|1,0>|^2 = 1
651        let c1 = clebsch_gordan_doubled(1, 1, 2, 1, -1, 0);
652        let c2 = clebsch_gordan_doubled(1, 1, 2, -1, 1, 0);
653        assert_relative_eq!(c1 * c1 + c2 * c2, 1.0);
654    }
655
656    #[test]
657    fn normalization_for_singlet() {
658        // For j1=j2=1/2 and total J=0, M=0:
659        // |<+,-|0,0>|^2 + |<-,+|0,0>|^2 = 1
660        let c1 = clebsch_gordan_doubled(1, 1, 0, 1, -1, 0);
661        let c2 = clebsch_gordan_doubled(1, 1, 0, -1, 1, 0);
662        assert_relative_eq!(c1 * c1 + c2 * c2, 1.0);
663    }
664
665    #[test]
666    fn two_spin_half_cases() {
667        // (1/2 1/2 1 ; 1/2 1/2 -1) = -1/sqrt(3)
668        assert_relative_eq!(wigner_3j_doubled(1, 1, 2, 1, 1, -2), -1.0 / 3.0_f64.sqrt());
669
670        // (1/2 1/2 0 ; 1/2 -1/2 0) = 1/sqrt(2)
671        assert_relative_eq!(wigner_3j_doubled(1, 1, 0, 1, -1, 0), FRAC_1_SQRT_2);
672
673        // (1/2 1/2 0 ; -1/2 1/2 0) = -1/sqrt(2)
674        assert_relative_eq!(wigner_3j_doubled(1, 1, 0, -1, 1, 0), -FRAC_1_SQRT_2);
675    }
676
677    #[test]
678    fn spin_one_cases() {
679        // (1 1 0 ; 0 0 0) = -1/sqrt(3)
680        assert_relative_eq!(wigner_3j_doubled(2, 2, 0, 0, 0, 0), -1.0 / 3.0_f64.sqrt());
681
682        // (1 1 2 ; 1 -1 0) = 1/sqrt(30)
683        assert_relative_eq!(wigner_3j_doubled(2, 2, 4, 2, -2, 0), 1.0 / 30.0_f64.sqrt());
684
685        // (1 1 2 ; 0 0 0) = sqrt(2/15)
686        assert_relative_eq!(wigner_3j_doubled(2, 2, 4, 0, 0, 0), (2.0 / 15.0_f64).sqrt());
687    }
688
689    #[test]
690    fn selection_rule_failures_return_zero() {
691        // m1 + m2 + m3 != 0
692        assert_eq!(wigner_3j_doubled(1, 1, 0, 1, -1, 1), 0.0);
693
694        // triangle rule fails: 1/2 + 1/2 cannot couple to 2
695        assert_eq!(wigner_3j_doubled(1, 1, 4, 1, -1, 0), 0.0);
696
697        // invalid m for j = 1/2
698        assert_eq!(wigner_3j_doubled(1, 1, 0, 3, -1, -2), 0.0);
699    }
700
701    #[test]
702    fn odd_j_sum_with_all_zero_ms_vanishes() {
703        // For integer j's, (j1 j2 j3; 0 0 0) vanishes if j1+j2+j3 is odd.
704        // Here 1+1+1 = 3 is odd.
705        assert_eq!(wigner_3j_doubled(2, 2, 2, 0, 0, 0), 0.0);
706    }
707
708    #[test]
709    fn column_swap_symmetry_even_case() {
710        // Swapping first two columns gives factor (-1)^(j1+j2+j3).
711        // Here 1+1+2 = 4 is even, so unchanged.
712        let a = wigner_3j_doubled(2, 2, 4, 2, -2, 0);
713        let b = wigner_3j_doubled(2, 2, 4, -2, 2, 0);
714        assert_relative_eq!(a, b);
715    }
716
717    #[test]
718    fn column_swap_symmetry_odd_case() {
719        // Here 1/2 + 1/2 + 0 = 1 is odd, so swap picks up a minus sign.
720        let a = wigner_3j_doubled(1, 1, 0, 1, -1, 0);
721        let b = wigner_3j_doubled(1, 1, 0, -1, 1, 0);
722        assert_relative_eq!(a, -b);
723    }
724
725    #[test]
726    fn sign_flip_symmetry() {
727        // (j1 j2 j3; -m1 -m2 -m3) = (-1)^(j1+j2+j3) (j1 j2 j3; m1 m2 m3)
728        // For j1=j2=1/2, j3=0, total is odd => minus sign.
729        let a = wigner_3j_doubled(1, 1, 0, 1, -1, 0);
730        let b = wigner_3j_doubled(1, 1, 0, -1, 1, 0);
731        assert_relative_eq!(b, -a);
732
733        // For j1=j2=1, j3=2, total is even => same sign.
734        let c = wigner_3j_doubled(2, 2, 4, 2, -2, 0);
735        let d = wigner_3j_doubled(2, 2, 4, -2, 2, 0);
736        assert_relative_eq!(d, c);
737    }
738
739    #[test]
740    fn typed_wigner_3j_matches_doubled_helper() {
741        assert_relative_eq!(
742            wigner_3j(
743                crate::j!(1 / 2),
744                crate::m!(1 / 2),
745                crate::j!(1 / 2),
746                crate::m!(-1 / 2),
747                crate::j!(1),
748                crate::m!(0)
749            ),
750            wigner_3j_doubled(1, 1, 2, 1, -1, 0)
751        );
752        assert_eq!(
753            wigner_3j(
754                crate::j!(1 / 2),
755                crate::m!(1 / 2),
756                crate::j!(1 / 2),
757                crate::m!(1 / 2),
758                crate::j!(1),
759                crate::m!(0)
760            ),
761            0.0
762        );
763    }
764
765    #[test]
766    fn relation_to_clebsch_gordon_examples() {
767        // Using:
768        // (j1 j2 j3; m1 m2 -m3) = (-1)^(j1-j2+m3) / sqrt(2j3+1) * <j1 m1 j2 m2 | j3 m3>
769        //
770        // In doubled notation the phase exponent becomes (dj1 - dj2 + dm3)/2.
771
772        // <1/2,1/2; 1/2,-1/2 | 0,0> = 1/sqrt(2)
773        // => (1/2 1/2 0 ; 1/2 -1/2 0) = 1/sqrt(2)
774        let cg = clebsch_gordan_doubled(1, 1, 0, 1, -1, 0);
775        let w3j = wigner_3j_doubled(1, 1, 0, 1, -1, 0);
776        assert_relative_eq!(w3j, cg);
777
778        // <1,1; 1,-1 | 2,0> = 1/sqrt(6)
779        // => (1 1 2 ; 1 -1 0) = 1/sqrt(30)
780        let cg = clebsch_gordan_doubled(2, 2, 4, 2, -2, 0);
781        let expected = cg / 5.0_f64.sqrt(); // sqrt(2j3+1)=sqrt(5)
782        let w3j = wigner_3j_doubled(2, 2, 4, 2, -2, 0);
783        assert_relative_eq!(w3j, expected);
784    }
785
786    #[test]
787    fn construct_integer_case() {
788        let _ = WignerDMatrix::new(j!(1), m!(1), m!(0)).unwrap();
789    }
790
791    #[test]
792    fn construct_half_integer_case() {
793        let _ = WignerDMatrix::new(j!(1 / 2), m!(1 / 2), m!(-1 / 2)).unwrap();
794    }
795
796    #[test]
797    fn invalid_wigner_d_quantum_numbers_error() {
798        assert!(WignerDMatrix::new(j!(1), m!(2), m!(0)).is_err());
799        assert!(WignerDMatrix::new(j!(1), m!(0), m!(2)).is_err());
800        assert!(WignerDMatrix::new(j!(1), m!(1 / 2), m!(0)).is_err());
801        assert!(WignerDMatrix::new(j!(1), m!(0), m!(1 / 2)).is_err());
802    }
803
804    #[test]
805    fn small_d_matches_known_numerical_values() {
806        let beta = 1.1;
807        let cb = f64::cos(beta);
808        let sb = f64::sin(beta);
809
810        let w_11 = WignerDMatrix::new(j!(1), m!(1), m!(1)).unwrap();
811        let w_10 = WignerDMatrix::new(j!(1), m!(1), m!(0)).unwrap();
812        let w_1m1 = WignerDMatrix::new(j!(1), m!(1), m!(-1)).unwrap();
813        let w_00 = WignerDMatrix::new(j!(1), m!(0), m!(0)).unwrap();
814
815        assert_complex_relative_eq(evaluate(w_11.d(beta)), Complex64::from(0.5 * (1.0 + cb)));
816        assert_complex_relative_eq(evaluate(w_10.d(beta)), Complex64::from(-FRAC_1_SQRT_2 * sb));
817        assert_complex_relative_eq(evaluate(w_1m1.d(beta)), Complex64::from(0.5 * (1.0 - cb)));
818        assert_complex_relative_eq(evaluate(w_00.d(beta)), Complex64::from(cb));
819        assert_complex_relative_eq(evaluate(w_10.d(FRAC_PI_2)), Complex64::from(-FRAC_1_SQRT_2));
820    }
821
822    #[test]
823    fn full_d_matches_phase_definition_numerically() {
824        let alpha = 0.31;
825        let beta = 0.82;
826        let gamma = -0.47;
827        let w = WignerDMatrix::new(j!(3 / 2), m!(1 / 2), m!(-1 / 2)).unwrap();
828
829        let d = evaluate(w.d(beta));
830        let expected = Complex64::cis(-0.5 * (alpha - gamma)) * d;
831
832        assert_complex_relative_eq(evaluate(w.D(alpha, beta, gamma)), expected);
833    }
834
835    #[test]
836    fn small_d_builds_regular_expression_graph() {
837        use laddu_expr::{ExprNode, UnaryOp, event_scalar};
838
839        let w = WignerDMatrix::new(j!(1), m!(1), m!(0)).unwrap();
840        let graph = w.d(event_scalar("beta")).to_graph();
841
842        assert!(graph.nodes().iter().any(|node| matches!(
843            node,
844            ExprNode::Unary {
845                op: UnaryOp::Sin | UnaryOp::Cos | UnaryOp::PowI(_),
846                ..
847            }
848        )));
849    }
850
851    #[test]
852    fn full_d_builds_regular_expression_graph() {
853        use laddu_expr::{ExprNode, event_scalar};
854
855        let w = WignerDMatrix::new(j!(3 / 2), m!(1 / 2), m!(-1 / 2)).unwrap();
856        let graph = w
857            .D(
858                event_scalar("alpha"),
859                event_scalar("beta"),
860                event_scalar("gamma"),
861            )
862            .to_graph();
863
864        assert!(
865            graph
866                .nodes()
867                .iter()
868                .any(|node| matches!(node, ExprNode::ComplexConst(_)))
869        );
870        assert!(graph.nodes().iter().all(|node| !matches!(
871            node,
872            ExprNode::Solve { .. } | ExprNode::MatMul { .. } | ExprNode::MatVec { .. }
873        )));
874    }
875}