Skip to main content

rusolver/
distributed.rs

1// SPDX-License-Identifier: Apache-2.0
2//! Row-partitioned HOST FP64 preconditioned CG. Matrix rows stay partitioned;
3//! search vectors are all-gathered, scalars globally reduced. No device/RDMA claim.
4use crate::{SolverError,CgOptions,IterativeStatus};
5use crate::numerics::{checked,dot,finite,sum_iter,zeros};
6use crate::sparse_direct::validate_csr;
7/// Ordered blocking communicator. All ranks call the same solve in the same order.
8/// all_gather_fixed returns one equally-sized buffer PER rank, in rank order.
9/// abort MUST wake pending peers (or transport must have a finite timeout).
10pub trait SolverCommunicator {
11    fn rank(&self)->usize;
12    fn size(&self)->usize;
13    fn all_gather_fixed(&self,phase:u64,local:&[f64])->Result<Vec<Vec<f64>>,SolverError>;
14    fn abort(&self,reason:&str);
15}
16pub struct LocalCommunicator;
17impl SolverCommunicator for LocalCommunicator{
18    fn rank(&self)->usize{0}fn size(&self)->usize{1}
19    fn all_gather_fixed(&self,_:u64,local:&[f64])->Result<Vec<Vec<f64>>,SolverError>{Ok(vec![local.to_vec()])}
20    fn abort(&self,_:&str){}
21}
22#[derive(Clone,Copy,Debug,PartialEq,Eq)]pub struct RowPartition{pub global_rows:usize,pub start:usize,pub rows:usize}
23impl RowPartition{
24    pub fn balanced(n:usize,rank:usize,size:usize)->Result<Self,SolverError>{
25        if n==0||size==0||rank>=size||n>u32::MAX as usize||size>u32::MAX as usize{return Err(SolverError::Shape("partition/rank"));}
26        let q=n/size;let rem=n%size;let rows=q+usize::from(rank<rem);
27        let start=rank*q+rank.min(rem);Ok(Self{global_rows:n,start,rows})
28    }
29}
30pub struct DistributedCsr<'a>{pub partition:RowPartition,offsets:&'a[usize],columns:&'a[usize],values:&'a[f64]}
31impl<'a>DistributedCsr<'a>{
32    pub fn new(partition:RowPartition,offsets:&'a[usize],columns:&'a[usize],values:&'a[f64])->Result<Self,SolverError>{
33        if partition.start.checked_add(partition.rows).is_none_or(|end|end>partition.global_rows){return Err(SolverError::Shape("partition range"));}
34        validate_csr(partition.rows,partition.global_rows,offsets,columns,values)?;Ok(Self{partition,offsets,columns,values})
35    }
36    fn apply(&self,x:&[f64])->Result<Vec<f64>,SolverError>{
37        if x.len()!=self.partition.global_rows{return Err(SolverError::Shape("distributed vector length"));}let mut y=zeros(self.partition.rows)?;
38        for i in 0..y.len(){y[i]=sum_iter((self.offsets[i]..self.offsets[i+1]).map(|p|self.values[p]*x[self.columns[p]]))?;}Ok(y)
39    }
40    fn diagonal(&self)->Result<Vec<f64>,SolverError>{
41        let mut d=zeros(self.partition.rows)?;for i in 0..d.len(){d[i]=sum_iter((self.offsets[i]..self.offsets[i+1]).filter(|&p|self.columns[p]==i+self.partition.start).map(|p|self.values[p]))?;
42            if d[i]<=0.0{return Err(SolverError::NotPositiveDefinite{index:i+self.partition.start,pivot:d[i]});}}
43        Ok(d)
44    }
45}
46#[derive(Clone,Debug)]pub struct DistributedCgReport{
47    pub local_solution:Vec<f64>,pub partition:RowPartition,pub iterations:usize,
48    pub residual_norm:f64,pub rhs_norm:f64,pub status:IterativeStatus,pub collective_calls:u64,
49}
50struct Exchange<'a,C:?Sized>{comm:&'a C,phase:u64}
51impl<C:SolverCommunicator+?Sized>Exchange<'_,C>{
52    fn gather(&mut self,x:&[f64])->Result<Vec<Vec<f64>>,SolverError>{
53        self.phase=self.phase.checked_add(1).ok_or(SolverError::SizeOverflow)?;
54        let all=self.comm.all_gather_fixed(self.phase,x)?;
55        if all.len()!=self.comm.size()||all.iter().any(|v|v.len()!=x.len()){return Err(SolverError::Communication("all-gather shape mismatch".into()));}
56        for v in &all{finite(v)?;}Ok(all)
57    }
58    fn sum(&mut self,x:f64)->Result<f64,SolverError>{checked(x,"local reduction")?;sum_iter(self.gather(&[x])?.iter().map(|v|v[0]))}
59    fn norm(&mut self,x:&[f64])->Result<f64,SolverError>{
60        let local=x.iter().fold(0.0f64,|a,b|a.max(b.abs()));
61        let scale=self.gather(&[local])?.iter().fold(0.0f64,|a,b|a.max(b[0]));
62        if scale==0.0{return Ok(0.0);}let s=sum_iter(x.iter().map(|&x|(x/scale)*(x/scale)))?;
63        checked(scale*self.sum(s)?.sqrt(),"distributed norm")
64    }
65    fn vector(&mut self,local:&[f64],n:usize)->Result<Vec<f64>,SolverError>{
66        let width=n.div_ceil(self.comm.size());let mut padded=zeros(width)?;padded[..local.len()].copy_from_slice(local);
67        let all=self.gather(&padded)?;let mut full=zeros(n)?;
68        for(rank,v)in all.iter().enumerate(){let p=RowPartition::balanced(n,rank,self.comm.size())?;full[p.start..p.start+p.rows].copy_from_slice(&v[..p.rows]);}Ok(full)
69    }
70}
71/// Solve SPD A x=b with balanced contiguous row partitioning. optional Jacobi is
72/// local diagonal scaling. No checkpoint/restart inside a solve: an error aborts
73/// the communicator; create a fresh communicator before retrying.
74pub fn distributed_cg<C:SolverCommunicator+?Sized>(comm:&C,a:&DistributedCsr<'_>,b:&[f64],
75initial:Option<&[f64]>,jacobi:bool,options:CgOptions)->Result<DistributedCgReport,SolverError>{
76    let result=solve(comm,a,b,initial,jacobi,options);
77    if let Err(ref e)=result{comm.abort(&e.to_string());}result
78}
79fn solve<C:SolverCommunicator+?Sized>(comm:&C,a:&DistributedCsr<'_>,b:&[f64],initial:Option<&[f64]>,jacobi:bool,options:CgOptions)->Result<DistributedCgReport,SolverError>{
80    let partition=RowPartition::balanced(a.partition.global_rows,comm.rank(),comm.size())?;
81    if partition!=a.partition||b.len()!=partition.rows||initial.is_some_and(|x|x.len()!=b.len()){return Err(SolverError::Shape("distributed local partition/RHS"));}
82    finite(b)?;options.tolerance.validate()?;
83    if options.residual_recompute_interval==0||options.residual_recompute_interval>u32::MAX as usize||options.max_iterations>u32::MAX as usize||(options.tolerance.absolute==0.0&&options.tolerance.relative==0.0){return Err(SolverError::InvalidOption("distributed CG options"));}
84    let mut net=Exchange{comm,phase:0};
85    let config=[partition.global_rows as f64,options.max_iterations as f64,options.residual_recompute_interval as f64,
86        options.tolerance.absolute,options.tolerance.relative,if jacobi{1.0}else{0.0},if initial.is_some(){1.0}else{0.0}];
87    if net.gather(&config)?.iter().any(|v|v.as_slice()!=config){return Err(SolverError::Communication("ranks disagree on solve options".into()));}
88    let diag=if jacobi{a.diagonal()?}else{vec![1.0;b.len()]};
89    let mut x=initial.map_or_else(||vec![0.0;b.len()],|v|v.to_vec());finite(&x)?;
90    let n=partition.global_rows;let rhs_norm=net.norm(b)?;let target=options.tolerance.threshold(rhs_norm)?;
91    let ax=a.apply(&net.vector(&x,n)?)?;let mut r:Vec<f64>=b.iter().zip(&ax).map(|(&b,&ax)|b-ax).collect();finite(&r)?;
92    let mut residual=net.norm(&r)?;
93    let mut iterations=0;let mut status=IterativeStatus::MaxIterations;
94    let mut z=zeros(b.len())?;for i in 0..z.len(){z[i]=checked(r[i]/diag[i],"distributed preconditioner")?;}
95    let mut p=z.clone();let mut rho=net.sum(dot(&r,&z)?)?;
96    if residual<=target{status=IterativeStatus::Converged;}
97    else{for it in 1..=options.max_iterations{
98        if rho<=0.0{return Err(SolverError::Breakdown("distributed nonpositive residual/preconditioner"));}
99        let ap=a.apply(&net.vector(&p,n)?)?;let curvature=net.sum(dot(&p,&ap)?)?;
100        if curvature<=0.0{return Err(SolverError::Breakdown("distributed CG requires SPD operator"));}
101        let alpha=checked(rho/curvature,"distributed alpha")?;
102        for i in 0..x.len(){x[i]=checked(alpha.mul_add(p[i],x[i]),"distributed iterate")?;r[i]=checked((-alpha).mul_add(ap[i],r[i]),"distributed residual")?;}
103        residual=net.norm(&r)?;let replace=it%options.residual_recompute_interval==0||residual<=target||it==options.max_iterations;
104        if replace{let ax=a.apply(&net.vector(&x,n)?)?;for i in 0..r.len(){r[i]=checked(b[i]-ax[i],"distributed true residual")?;}residual=net.norm(&r)?;}
105        iterations=it;if residual<=target{status=IterativeStatus::Converged;break;}
106        if it==options.max_iterations{break;}
107        for i in 0..z.len(){z[i]=checked(r[i]/diag[i],"distributed Jacobi")?;}
108        let next=net.sum(dot(&r,&z)?)?;if next<=0.0{return Err(SolverError::Breakdown("distributed CG rho underflow"));}
109        let beta=checked(next/rho,"distributed beta")?;
110        for i in 0..p.len(){p[i]=if replace{z[i]}else{checked(beta.mul_add(p[i],z[i]),"distributed direction")?};}rho=next;
111    }}
112    Ok(DistributedCgReport{local_solution:x,partition,iterations,residual_norm:residual,rhs_norm,status,collective_calls:net.phase})
113}
114
115/// Reuses the EXISTING GXCL transport (CPU/TCP path), not a second socket stack.
116/// A dedicated session is required; don't interleave unrelated collectives.
117#[cfg(feature="collective")]
118pub struct RucclCommunicator<'a>{session:&'a dyn ruccl::rank::RankTransport}
119#[cfg(feature="collective")]
120impl<'a>RucclCommunicator<'a>{
121    pub fn new(session:&'a dyn ruccl::rank::RankTransport)->Result<Self,SolverError>{
122        use ruccl::rank::CollectiveTransport;
123        if !matches!(session.transport(),CollectiveTransport::TcpHostStaged|CollectiveTransport::TcpPeer){return Err(SolverError::InvalidOption("solver adapter currently supports GXCL TCP transports only"));}
124        Ok(Self{session})
125    }
126}
127#[cfg(feature="collective")]
128impl SolverCommunicator for RucclCommunicator<'_>{
129    fn rank(&self)->usize{self.session.rank()as usize}fn size(&self)->usize{self.session.world_size()as usize}
130    fn all_gather_fixed(&self,phase:u64,local:&[f64])->Result<Vec<Vec<f64>>,SolverError>{
131        use ruccl::rank::{Opcode,ElementType,ANY_RANK};
132        let mut words=Vec::with_capacity(local.len()+2);words.push((phase>>32)as f64);words.push((phase as u32)as f64);words.extend_from_slice(local);
133        let bytes:Vec<u8>=words.iter().flat_map(|x|x.to_le_bytes()).collect();
134        let response=self.session.exchange(Opcode::AllGather,ElementType::F64,ANY_RANK,words.len()as u64,bytes)
135            .map_err(|e|SolverError::Communication(e.to_string()))?;
136        let width=words.len().checked_mul(8).ok_or(SolverError::SizeOverflow)?;
137        if response.payload.len()!=width.checked_mul(self.size()).ok_or(SolverError::SizeOverflow)?{return Err(SolverError::Communication("GXCL allgather byte count".into()));}
138        let mut out=Vec::new();for rank_bytes in response.payload.chunks_exact(width){
139            let mut v=Vec::with_capacity(words.len());for b in rank_bytes.chunks_exact(8){let mut a=[0;8];a.copy_from_slice(b);v.push(f64::from_le_bytes(a));}
140            if v[..2]!=words[..2]{return Err(SolverError::Communication("solver collective phase mismatch".into()));}out.push(v[2..].to_vec());
141        }Ok(out)
142    }
143    fn abort(&self,reason:&str){let _=self.session.abort(reason);}
144}