Skip to main content

cas_poly/
lib.rs

1//! cas 多项式层(P0-M2):有序稀疏分布式多项式,泛型于 [`Ring`] 系数。
2//!
3//! 设计见 `cas/DESIGN.md` D3:
4//! - **有序**:项按 `PolyRing` 携带的单项式序(缺省 DegRevLex)严格升序,
5//!   [`Poly::terms`] 的次序即序列化/打印/代码生成次序(D6 四方同序);
6//! - **稀疏**:零系数项不存在,指数向量平坦存储(`exps[i*nvars..]`);
7//! - **确定性**:构造即规范化(合并同类项、去零、按项序排序);
8//!   乘法经哈希合并后**必经排序出口**,内部哈希迭代序不进入输出。
9//!
10//! M2 范围:加/减/乘/幂/偏导/求值。gcd、除法、因式分解属 M3/M4/M5;
11//! ℚ 上"本原 ℤ 表示 + 容度"的去分母优化同属 M3。
12
13use cas_domain::{Rational, Ring};
14use std::cmp::Ordering;
15
16mod factor;
17mod factor_mv;
18mod gcd;
19use std::collections::HashMap;
20use std::fmt;
21use std::sync::Arc;
22
23/// 单项式序。变元位置即 `PolyRing::vars` 的下标。
24#[derive(Clone, Copy, Debug, PartialEq, Eq)]
25pub enum MonOrder {
26    /// 字典序:左起第一个非零差决定。
27    Lex,
28    /// 分次反字典序(Gröbner 友好,缺省):先比总次数;同次时右起
29    /// 第一个非零差为负者更小。
30    DegRevLex,
31}
32
33pub(crate) fn cmp_monomials(order: MonOrder, a: &[u32], b: &[u32]) -> Ordering {
34    debug_assert_eq!(a.len(), b.len());
35    match order {
36        MonOrder::Lex => {
37            for (&x, &y) in a.iter().zip(b.iter()) {
38                if x != y {
39                    return x.cmp(&y);
40                }
41            }
42            Ordering::Equal
43        }
44        MonOrder::DegRevLex => {
45            let da: u64 = a.iter().map(|&e| e as u64).sum();
46            let db: u64 = b.iter().map(|&e| e as u64).sum();
47            da.cmp(&db).then_with(|| {
48                for i in (0..a.len()).rev() {
49                    let d = a[i] as i64 - b[i] as i64;
50                    if d != 0 {
51                        return d.cmp(&0);
52                    }
53                }
54                Ordering::Equal
55            })
56        }
57    }
58}
59
60/// 多项式环身份 =(变元表, 项序)。相同身份的 Poly 才可相互运算。
61#[derive(Clone, Debug, PartialEq, Eq)]
62pub struct PolyRing {
63    pub vars: Vec<Box<str>>,
64    pub order: MonOrder,
65}
66
67impl PolyRing {
68    pub fn new<I: IntoIterator<Item: Into<Box<str>>>>(vars: I, order: MonOrder) -> Arc<Self> {
69        let vars: Vec<Box<str>> = vars.into_iter().map(Into::into).collect();
70        assert!(!vars.is_empty(), "PolyRing 至少一个变元");
71        Arc::new(PolyRing { vars, order })
72    }
73
74    pub fn nvars(&self) -> usize {
75        self.vars.len()
76    }
77}
78
79/// 稀疏有序多项式。`exps` 长度恒为 `nvars × nterms`,第 i 项指数为
80/// `exps[i*nvars..(i+1)*nvars]`;项按环的项序严格升序,系数零的项不存在。
81#[derive(Clone, PartialEq, Eq)]
82pub struct Poly<C: Ring> {
83    ring: Arc<PolyRing>,
84    exps: Vec<u32>,
85    coefs: Vec<C>,
86}
87
88impl<C: Ring> Poly<C> {
89    /// 由(指数向量, 系数)序列构造,构造即规范化:合并同类项、去零、排序。
90    pub fn from_terms<I: IntoIterator<Item = (Vec<u32>, C)>>(
91        ring: Arc<PolyRing>,
92        terms: I,
93    ) -> Self {
94        let nvars = ring.nvars();
95        let mut acc: HashMap<Vec<u32>, C> = HashMap::new();
96        for (e, c) in terms {
97            assert_eq!(e.len(), nvars, "指数向量长度须等于变元数 {nvars}");
98            if c.is_zero() {
99                continue;
100            }
101            let slot = acc.entry(e).or_insert_with(Ring::zero);
102            *slot = slot.add(&c);
103        }
104        let mut items: Vec<(Vec<u32>, C)> = acc.into_iter().filter(|(_, c)| !c.is_zero()).collect();
105        let order = ring.order;
106        items.sort_by(|(a, _), (b, _)| cmp_monomials(order, a, b));
107        let exps: Vec<u32> = items.iter().flat_map(|(e, _)| e.iter().copied()).collect();
108        let coefs: Vec<C> = items.into_iter().map(|(_, c)| c).collect();
109        Poly { ring, exps, coefs }
110    }
111
112    pub fn zero(ring: Arc<PolyRing>) -> Self {
113        Poly {
114            ring,
115            exps: Vec::new(),
116            coefs: Vec::new(),
117        }
118    }
119
120    pub fn constant(ring: Arc<PolyRing>, c: C) -> Self {
121        if c.is_zero() {
122            Self::zero(ring)
123        } else {
124            let n = ring.nvars();
125            Poly {
126                ring,
127                exps: vec![0; n],
128                coefs: vec![c],
129            }
130        }
131    }
132
133    pub fn ring(&self) -> &Arc<PolyRing> {
134        &self.ring
135    }
136
137    pub fn nterms(&self) -> usize {
138        self.coefs.len()
139    }
140
141    pub fn is_zero(&self) -> bool {
142        self.coefs.is_empty()
143    }
144
145    pub fn is_constant(&self) -> bool {
146        self.coefs.len() <= 1 && self.exps.iter().all(|&e| e == 0)
147    }
148
149    /// 总次数(零多项式为 0)。
150    pub fn degree(&self) -> u64 {
151        self.nvars_terms()
152            .map(|(e, _)| e.iter().map(|&x| x as u64).sum())
153            .max()
154            .unwrap_or(0)
155    }
156
157    /// 稳定项迭代器:次序 = 项序 = 序列化/打印/代码生成次序(D6)。
158    pub fn terms(&self) -> impl Iterator<Item = (&[u32], &C)> {
159        let n = self.ring.nvars();
160        self.coefs
161            .iter()
162            .enumerate()
163            .map(move |(i, c)| (&self.exps[i * n..(i + 1) * n], c))
164    }
165
166    fn nvars_terms(&self) -> impl Iterator<Item = (&[u32], &C)> {
167        self.terms()
168    }
169
170    fn exps_of(&self, i: usize) -> &[u32] {
171        let n = self.ring.nvars();
172        &self.exps[i * n..(i + 1) * n]
173    }
174
175    pub fn add(&self, other: &Self) -> Self {
176        self.assert_same_ring(other);
177        // 两序列均已按项序升序 → 双指针归并
178        let (mut i, mut j, mut items) = (0, 0, Vec::with_capacity(self.nterms() + other.nterms()));
179        while i < self.nterms() && j < other.nterms() {
180            match cmp_monomials(self.ring.order, self.exps_of(i), other.exps_of(j)) {
181                Ordering::Less => {
182                    items.push((self.exps_of(i).to_vec(), self.coefs[i].clone()));
183                    i += 1;
184                }
185                Ordering::Greater => {
186                    items.push((other.exps_of(j).to_vec(), other.coefs[j].clone()));
187                    j += 1;
188                }
189                Ordering::Equal => {
190                    let c = self.coefs[i].add(&other.coefs[j]);
191                    if !c.is_zero() {
192                        items.push((self.exps_of(i).to_vec(), c));
193                    }
194                    i += 1;
195                    j += 1;
196                }
197            }
198        }
199        items.extend(self.terms_from(i).map(|(e, c)| (e.to_vec(), c.clone())));
200        items.extend(other.terms_from(j).map(|(e, c)| (e.to_vec(), c.clone())));
201        Self::from_sorted(self.ring.clone(), items)
202    }
203
204    fn terms_from(&self, start: usize) -> impl Iterator<Item = (&[u32], &C)> {
205        let n = self.ring.nvars();
206        (start..self.nterms()).map(move |i| (&self.exps[i * n..(i + 1) * n], &self.coefs[i]))
207    }
208
209    /// 由**已按项序升序**的项构造(归并路径专用,跳过排序)。
210    fn from_sorted(ring: Arc<PolyRing>, items: Vec<(Vec<u32>, C)>) -> Self {
211        let exps: Vec<u32> = items.iter().flat_map(|(e, _)| e.iter().copied()).collect();
212        let coefs: Vec<C> = items.into_iter().map(|(_, c)| c).collect();
213        Poly { ring, exps, coefs }
214    }
215
216    pub fn sub(&self, other: &Self) -> Self {
217        self.add(&other.neg())
218    }
219
220    pub fn neg(&self) -> Self {
221        Poly {
222            ring: self.ring.clone(),
223            exps: self.exps.clone(),
224            coefs: self.coefs.iter().map(Ring::neg).collect(),
225        }
226    }
227
228    pub fn mul(&self, other: &Self) -> Self {
229        self.assert_same_ring(other);
230        if self.is_zero() || other.is_zero() {
231            return Self::zero(self.ring.clone());
232        }
233        let mut acc: HashMap<Vec<u32>, C> = HashMap::with_capacity(self.nterms() * other.nterms());
234        for i in 0..self.nterms() {
235            let (ea, ca) = (self.exps_of(i), &self.coefs[i]);
236            for j in 0..other.nterms() {
237                let key: Vec<u32> = ea
238                    .iter()
239                    .zip(other.exps_of(j))
240                    .map(|(&x, &y)| x + y)
241                    .collect();
242                let slot = acc.entry(key).or_insert_with(Ring::zero);
243                *slot = slot.add(&ca.mul(&other.coefs[j]));
244            }
245        }
246        let mut items: Vec<(Vec<u32>, C)> = acc.into_iter().filter(|(_, c)| !c.is_zero()).collect();
247        let order = self.ring.order;
248        // 哈希迭代不进入输出(D6 纪律):出口必经排序
249        items.sort_by(|(a, _), (b, _)| cmp_monomials(order, a, b));
250        Self::from_sorted(self.ring.clone(), items)
251    }
252
253    pub fn pow(&self, e: u32) -> Self {
254        match e {
255            0 => Self::constant(self.ring.clone(), C::one()),
256            1 => self.clone(),
257            _ => {
258                // 平方求幂
259                let mut acc = Self::constant(self.ring.clone(), C::one());
260                let mut base = self.clone();
261                let mut k = e;
262                while k > 0 {
263                    if k & 1 == 1 {
264                        acc = acc.mul(&base);
265                    }
266                    k >>= 1;
267                    if k > 0 {
268                        base = base.mul(&base);
269                    }
270                }
271                acc
272            }
273        }
274    }
275
276    /// 对第 `idx` 个变元求偏导。
277    pub fn deriv(&self, idx: usize) -> Self {
278        assert!(idx < self.ring.nvars(), "变元下标越界");
279        let items: Vec<(Vec<u32>, C)> = self
280            .terms()
281            .filter(|(e, _)| e[idx] > 0)
282            .map(|(e, c)| {
283                let mut e2 = e.to_vec();
284                e2[idx] -= 1;
285                let c2 = c.mul(&C::from_i64_coeff(e[idx] as i64));
286                (e2, c2)
287            })
288            .collect();
289        Self::from_terms(self.ring.clone(), items)
290    }
291
292    /// 在给定点求值(点分量按变元次序)。长度不符属编程错误。
293    pub fn eval(&self, vals: &[C]) -> C {
294        assert_eq!(vals.len(), self.ring.nvars(), "求值点分量数须等于变元数");
295        let mut sum = C::zero();
296        for (e, c) in self.terms() {
297            let mut m = c.clone();
298            for (j, &power) in e.iter().enumerate() {
299                m = m.mul(&vals[j].pow_u32(power));
300            }
301            sum = sum.add(&m);
302        }
303        sum
304    }
305
306    fn assert_same_ring(&self, other: &Self) {
307        assert_eq!(
308            &*self.ring, &*other.ring,
309            "多项式须属于同一 PolyRing(变元表+项序)"
310        );
311    }
312}
313
314// 系数从 u64 构造的小助手已并入 Ring::from_i64_coeff。
315
316impl Poly<Rational> {
317    /// f64 求值(系数精确转 f64 后按项求和)——数值交叉验证用。
318    pub fn eval_f64(&self, vals: &[f64]) -> f64 {
319        assert_eq!(vals.len(), self.ring.nvars());
320        let mut sum = 0.0;
321        for (e, c) in self.terms() {
322            let mut m = c.to_f64();
323            for (j, &power) in e.iter().enumerate() {
324                m *= vals[j].powi(power as i32);
325            }
326            sum += m;
327        }
328        sum
329    }
330}
331
332impl<C: Ring + fmt::Debug> fmt::Debug for Poly<C> {
333    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
334        if self.is_zero() {
335            return write!(f, "0");
336        }
337        let mut first = true;
338        for (e, c) in self.terms() {
339            if !first {
340                write!(f, " + ")?;
341            }
342            first = false;
343            write!(f, "({c:?})")?;
344            for (j, &power) in e.iter().enumerate() {
345                if power > 0 {
346                    write!(f, "*{}^{}", self.ring.vars[j], power)?;
347                }
348            }
349        }
350        Ok(())
351    }
352}
353
354#[cfg(test)]
355mod tests {
356    use super::*;
357    use cas_domain::Integer;
358
359    fn ring6() -> Arc<PolyRing> {
360        PolyRing::new(["q1", "q2", "q3", "p1", "p2", "p3"], MonOrder::DegRevLex)
361    }
362
363    fn rat(n: i64, d: i64) -> Rational {
364        Rational::from_ints(&Integer::from_i64(n), &Integer::from_i64(d)).unwrap()
365    }
366
367    fn term(_ring: &Arc<PolyRing>, e: [u32; 6], c: Rational) -> (Vec<u32>, Rational) {
368        (e.to_vec(), c)
369    }
370
371    #[test]
372    fn 构造即规范化() {
373        let ring = ring6();
374        // 乱序 + 重复同类项 + 零系数 → 合并排序去零
375        let p = Poly::from_terms(
376            ring.clone(),
377            [
378                term(&ring, [0, 0, 0, 0, 0, 0], rat(1, 2)),
379                term(&ring, [2, 0, 0, 0, 0, 0], rat(1, 3)),
380                term(&ring, [2, 0, 0, 0, 0, 0], rat(1, 6)),
381                term(&ring, [1, 0, 0, 0, 0, 0], rat(0, 1)),
382            ],
383        );
384        assert_eq!(p.nterms(), 2);
385        // grevlex:常数(次数 0)在最前
386        let (e0, c0) = p.terms().next().unwrap();
387        assert!(e0.iter().all(|&x| x == 0) && c0 == &rat(1, 2));
388        // 同类项已合并
389        let (_, c1) = p.terms().nth(1).unwrap();
390        assert_eq!(*c1, rat(1, 2));
391    }
392
393    #[test]
394    fn 加乘分配律与交换律() {
395        let ring = ring6();
396        let mk = |seed: &mut u64| -> Poly<Rational> {
397            // 确定性小随机多项式
398            let mut items = Vec::new();
399            for _ in 0..6 {
400                let r = xorshift(seed);
401                let e: Vec<u32> = (0..6).map(|k| ((r >> (k * 5)) % 3) as u32).collect();
402                let c = rat((r % 19) as i64 - 9, ((r >> 8) % 7 + 1) as i64);
403                items.push((e, c));
404            }
405            Poly::from_terms(ring.clone(), items)
406        };
407        let mut s = 12345u64;
408        for _ in 0..20 {
409            let (a, b, c) = (mk(&mut s), mk(&mut s), mk(&mut s));
410            assert!(a.add(&b) == b.add(&a), "加法交换律");
411            assert!(a.mul(&b) == b.mul(&a), "乘法交换律");
412            assert!(a.mul(&b).mul(&c) == a.mul(&b.mul(&c)), "结合律");
413            assert!(a.add(&b).mul(&c) == a.mul(&c).add(&b.mul(&c)), "分配律");
414        }
415    }
416
417    #[test]
418    fn 幂偏导与求值() {
419        let ring = ring6();
420        // (1/2 q1)^4 = 1/16 q1^4
421        let p = Poly::from_terms(ring.clone(), [term(&ring, [1, 0, 0, 0, 0, 0], rat(1, 2))]).pow(4);
422        assert_eq!(p.nterms(), 1);
423        assert_eq!(p.terms().next().unwrap().1, &rat(1, 16));
424
425        // d/dq1 (3 q1^2 p2) = 6 q1 p2
426        let f = Poly::from_terms(ring.clone(), [term(&ring, [2, 0, 0, 0, 1, 0], rat(3, 1))]);
427        let df = f.deriv(0);
428        assert_eq!(
429            df.terms().next().unwrap(),
430            (&[1, 0, 0, 0, 1, 0][..], &rat(6, 1))
431        );
432
433        // 求值一致性:(f*g)(x) == f(x)*g(x)
434        let g = Poly::from_terms(ring.clone(), [term(&ring, [0, 1, 0, 0, 0, 0], rat(-1, 3))]);
435        let pt: Vec<Rational> = [2, -3, 5, 7, 11, 13].map(|v| rat(v, 1)).to_vec();
436        let lhs = f.mul(&g).eval(&pt);
437        let rhs = f.eval(&pt).mul(&g.eval(&pt));
438        assert_eq!(lhs, rhs);
439
440        // f64 求值与精确求值一致
441        let ptf: Vec<f64> = [2.0, -3.0, 5.0, 7.0, 11.0, 13.0].to_vec();
442        let exact = f.mul(&g).eval(&pt);
443        let flt = f.mul(&g).eval_f64(&ptf);
444        assert!((exact.to_f64() - flt).abs() < 1e-12);
445    }
446
447    #[test]
448    fn 项序_grevlex约定() {
449        let a = [2u32, 0, 0, 0, 0, 0]; // q1^2,总次 2
450        let b = [0u32, 0, 0, 0, 0, 1]; // p3,总次 1
451        let c = [1u32, 0, 0, 0, 0, 1]; // q1 p3,总次 2
452        // grevlex:先比总次数(p3 < q1^2)
453        assert_eq!(cmp_monomials(MonOrder::DegRevLex, &b, &a), Ordering::Less);
454        // 同总次:右起第一个非零差为负者更小 → q1^2 < q1*p3
455        assert_eq!(cmp_monomials(MonOrder::DegRevLex, &a, &c), Ordering::Less);
456        // 字典序:左起第一个非零差为正者更大 → q1^2 > q1*p3
457        assert_eq!(cmp_monomials(MonOrder::Lex, &a, &c), Ordering::Greater);
458    }
459
460    fn xorshift(s: &mut u64) -> u64 {
461        let mut x = *s | 1;
462        x ^= x << 13;
463        x ^= x >> 7;
464        x ^= x << 17;
465        *s = x;
466        x
467    }
468}