use crate::gpu_kernels::sae_rowjet::SaeRowJetPath;
pub const ARROW_REDUCTION_LEAF_ROWS: usize = gam_linalg::pairwise_reduce::BASE_CHUNK;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ArrowCurvature {
GaussNewton,
ExactNewton,
}
impl ArrowCurvature {
#[inline]
pub fn residual_scale(self) -> f64 {
match self {
Self::GaussNewton => 0.0,
Self::ExactNewton => 1.0,
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct ArrowBlocks {
pub n_rows: usize,
pub q: usize,
pub n_beta: usize,
pub g_t: Vec<f64>,
pub h_tt: Vec<f64>,
pub h_tb: Vec<f64>,
pub g_beta: Vec<f64>,
pub h_bb: Vec<f64>,
}
impl ArrowBlocks {
}
#[derive(Debug, Clone, PartialEq)]
pub struct ArrowDirection {
pub t: Vec<f64>,
pub beta: Vec<f64>,
}
#[derive(Debug, Clone, PartialEq)]
pub struct ArrowScore {
pub n_rows: usize,
pub q: usize,
pub n_beta: usize,
pub g_t: Vec<f64>,
pub g_beta: Vec<f64>,
}
pub struct ResidentRowJetHandle {
n_atoms: usize,
q: usize,
p: usize,
n_beta: usize,
inv_tau: f64,
capacity_rows: usize,
path: SaeRowJetPath,
}
impl std::fmt::Debug for ResidentRowJetHandle {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ResidentRowJetHandle")
.field("n_atoms", &self.n_atoms)
.field("q", &self.q)
.field("p", &self.p)
.field("n_beta", &self.n_beta)
.field("inv_tau", &self.inv_tau)
.field("capacity_rows", &self.capacity_rows)
.field("path", &self.path)
.finish()
}
}
impl ResidentRowJetHandle {
#[inline]
pub fn deterministic(&self) -> bool {
true
}
#[inline]
pub fn path(&self) -> SaeRowJetPath {
self.path
}
}
pub const RESIDENT_ARROW_KERNEL_SOURCE: &str = r#"
__device__ __forceinline__ double row_mean(
const double* z, const int* active, const double* decoded,
int row, int k, int p, int c)
{
double mean=0.0;
for(int a=0;a<k;++a){
if(active[row*k+a]) mean=__dadd_rn(mean, __dmul_rn(z[row*k+a], decoded[(row*k+a)*p+c]));
}
return mean;
}
__device__ __forceinline__ double channel_first(
const double* z, const int* active, const int* kind, const int* atom,
const double* decoded, const double* d1, const double* sqrt_w,
double inv_tau, int k, int q, int p, int row, int slot, int c, double mean)
{
int a=atom[row*q+slot];
double root=sqrt_w[row];
if(kind[row*q+slot]==0){
double component = active[row*k+a] ? decoded[(row*k+a)*p+c] : 0.0;
double centered = component - mean;
double coefficient = __dmul_rn(root, __dmul_rn(inv_tau, z[row*k+a]));
return __dmul_rn(coefficient, centered);
}
if(!active[row*k+a]) return 0.0;
double coefficient=__dmul_rn(z[row*k+a], root);
return __dmul_rn(coefficient, d1[(row*q+slot)*p+c]);
}
__device__ __forceinline__ double channel_second(
const double* z, const int* active, const int* kind, const int* atom,
const double* decoded, const double* d1, const double* d2,
const double* sqrt_w, double inv_tau, int k, int q, int p,
int row, int slot_a, int slot_b, int c, double mean)
{
int ka=kind[row*q+slot_a], kb=kind[row*q+slot_b];
int aa=atom[row*q+slot_a], ab=atom[row*q+slot_b];
double root=sqrt_w[row];
if(ka==0 && kb==0){
double component_a = active[row*k+aa] ? decoded[(row*k+aa)*p+c] : 0.0;
double component_b = active[row*k+ab] ? decoded[(row*k+ab)*p+c] : 0.0;
double centered_a = component_a - mean;
double centered_b = component_b - mean;
double za=z[row*k+aa], zb=z[row*k+ab];
double diagonal = aa==ab ? 1.0 : 0.0;
double common=__dmul_rn(__dmul_rn(inv_tau, inv_tau), za);
double coefficient_a=__dmul_rn(root, __dmul_rn(common, diagonal-zb));
double coefficient_b=__dmul_rn(root, __dmul_rn(-common, zb));
return __dadd_rn(__dmul_rn(coefficient_a, centered_a),
__dmul_rn(coefficient_b, centered_b));
}
if(ka==0 || kb==0){
int logit_atom = ka==0 ? aa : ab;
int coord_atom = ka==0 ? ab : aa;
int coord_slot = ka==0 ? slot_b : slot_a;
if(!active[row*k+coord_atom]) return 0.0;
double diagonal = coord_atom==logit_atom ? 1.0 : 0.0;
double coefficient=__dmul_rn(__dmul_rn(z[row*k+coord_atom], diagonal-z[row*k+logit_atom]), inv_tau);
coefficient=__dmul_rn(coefficient, root);
return __dmul_rn(coefficient, d1[(row*q+coord_slot)*p+c]);
}
if(aa==ab){
if(!active[row*k+aa]) return 0.0;
double coefficient=__dmul_rn(z[row*k+aa], root);
return __dmul_rn(coefficient, d2[((row*q+slot_a)*q+slot_b)*p+c]);
}
return 0.0;
}
__device__ __forceinline__ double channel_beta(
const double* z, const int* active, const int* beta_atom,
const double* beta_phi, const double* beta_output, const double* sqrt_w,
int k, int p, int nb, int row, int border, int c)
{
int a=beta_atom[border];
if(!active[row*k+a]) return 0.0;
double base=__dmul_rn(z[row*k+a], beta_phi[row*nb+border]);
base=__dmul_rn(base, sqrt_w[row]);
return __dmul_rn(base, beta_output[border*p+c]);
}
__device__ __forceinline__ double channel_mixed(
const double* z, const int* active, const int* kind, const int* atom,
const int* beta_atom, const double* beta_phi, const double* beta_first,
const double* beta_output, const double* sqrt_w, double inv_tau,
int k, int q, int p, int nb, int row, int slot, int border, int c)
{
int target=beta_atom[border];
if(!active[row*k+target]) return 0.0;
int source_atom=atom[row*q+slot];
double scalar;
if(kind[row*q+slot]==0){
double diagonal = target==source_atom ? 1.0 : 0.0;
scalar=__dmul_rn(__dmul_rn(z[row*k+target], diagonal-z[row*k+source_atom]), inv_tau);
scalar=__dmul_rn(scalar, beta_phi[row*nb+border]);
}else if(source_atom==target){
scalar=__dmul_rn(z[row*k+target], beta_first[(row*q+slot)*nb+border]);
}else{
scalar=0.0;
}
scalar=__dmul_rn(scalar, sqrt_w[row]);
return __dmul_rn(scalar, beta_output[border*p+c]);
}
extern "C" __global__ void sae_arrow_gt(
const double* z, const int* active, const int* kind, const int* atom,
const double* decoded, const double* d1, const double* sqrt_w,
const double* residual, double inv_tau, int k, int q, int p,
unsigned long long total, double* g_t)
{
unsigned long long index=(unsigned long long)blockIdx.x*blockDim.x+threadIdx.x;
if(index>=total) return;
int slot=(int)(index%(unsigned long long)q);
int row=(int)(index/(unsigned long long)q);
double acc=0.0;
for(int c=0;c<p;++c){
double mean=row_mean(z,active,decoded,row,k,p,c);
double f=channel_first(z,active,kind,atom,decoded,d1,sqrt_w,inv_tau,k,q,p,row,slot,c,mean);
acc=__dadd_rn(acc, __dmul_rn(f, residual[row*p+c]));
}
g_t[index]=acc;
}
extern "C" __global__ void sae_arrow_htt(
const double* z, const int* active, const int* kind, const int* atom,
const double* decoded, const double* d1, const double* d2,
const double* sqrt_w, const double* residual, double inv_tau,
double scale, int k, int q, int p, unsigned long long total, double* h_tt)
{
unsigned long long index=(unsigned long long)blockIdx.x*blockDim.x+threadIdx.x;
if(index>=total) return;
int slot_b=(int)(index%(unsigned long long)q);
unsigned long long rem=index/(unsigned long long)q;
int slot_a=(int)(rem%(unsigned long long)q);
int row=(int)(rem/(unsigned long long)q);
double acc=0.0;
for(int c=0;c<p;++c){
double mean=row_mean(z,active,decoded,row,k,p,c);
double fa=channel_first(z,active,kind,atom,decoded,d1,sqrt_w,inv_tau,k,q,p,row,slot_a,c,mean);
double fb=channel_first(z,active,kind,atom,decoded,d1,sqrt_w,inv_tau,k,q,p,row,slot_b,c,mean);
double s2=channel_second(z,active,kind,atom,decoded,d1,d2,sqrt_w,inv_tau,k,q,p,row,slot_a,slot_b,c,mean);
double gauss_newton=__dmul_rn(fa, fb);
double curvature=__dmul_rn(scale, __dmul_rn(residual[row*p+c], s2));
acc=__dadd_rn(acc, __dadd_rn(gauss_newton, curvature));
}
h_tt[index]=acc;
}
extern "C" __global__ void sae_arrow_htb(
const double* z, const int* active, const int* kind, const int* atom,
const int* beta_atom, const double* decoded, const double* d1,
const double* beta_phi, const double* beta_first, const double* beta_output,
const double* sqrt_w, const double* residual, double inv_tau, double scale,
int k, int q, int p, int nb, unsigned long long total, double* h_tb)
{
unsigned long long index=(unsigned long long)blockIdx.x*blockDim.x+threadIdx.x;
if(index>=total) return;
int border=(int)(index%(unsigned long long)nb);
unsigned long long rem=index/(unsigned long long)nb;
int slot=(int)(rem%(unsigned long long)q);
int row=(int)(rem/(unsigned long long)q);
double acc=0.0;
for(int c=0;c<p;++c){
double mean=row_mean(z,active,decoded,row,k,p,c);
double f=channel_first(z,active,kind,atom,decoded,d1,sqrt_w,inv_tau,k,q,p,row,slot,c,mean);
double b=channel_beta(z,active,beta_atom,beta_phi,beta_output,sqrt_w,k,p,nb,row,border,c);
double m=channel_mixed(z,active,kind,atom,beta_atom,beta_phi,beta_first,beta_output,
sqrt_w,inv_tau,k,q,p,nb,row,slot,border,c);
double gauss_newton=__dmul_rn(f, b);
double curvature=__dmul_rn(scale, __dmul_rn(residual[row*p+c], m));
acc=__dadd_rn(acc, __dadd_rn(gauss_newton, curvature));
}
h_tb[index]=acc;
}
// Per-leaf partials of the shared beta blocks. Element layout per leaf:
// [0, nb) -> g_beta
// [nb, nb + nb*nb) -> h_bb (row-major)
// Rows inside a leaf are folded in ASCENDING order, matching the host mirror.
extern "C" __global__ void sae_arrow_beta_leaf(
const double* z, const int* active, const int* beta_atom,
const double* beta_phi, const double* beta_output, const double* sqrt_w,
const double* residual, int k, int p, int nb, int leaf_rows, int n_rows,
unsigned long long total, double* partials)
{
unsigned long long index=(unsigned long long)blockIdx.x*blockDim.x+threadIdx.x;
if(index>=total) return;
int width=nb+nb*nb;
int elem=(int)(index%(unsigned long long)width);
int leaf=(int)(index/(unsigned long long)width);
int start=leaf*leaf_rows;
int end=start+leaf_rows; if(end>n_rows) end=n_rows;
double acc=0.0;
for(int row=start;row<end;++row){
double contribution=0.0;
if(elem<nb){
int i=elem;
for(int c=0;c<p;++c){
double bi=channel_beta(z,active,beta_atom,beta_phi,beta_output,sqrt_w,k,p,nb,row,i,c);
contribution=__dadd_rn(contribution, __dmul_rn(bi, residual[row*p+c]));
}
}else{
int flat=elem-nb;
int i=flat/nb, j=flat%nb;
for(int c=0;c<p;++c){
double bi=channel_beta(z,active,beta_atom,beta_phi,beta_output,sqrt_w,k,p,nb,row,i,c);
double bj=channel_beta(z,active,beta_atom,beta_phi,beta_output,sqrt_w,k,p,nb,row,j,c);
contribution=__dadd_rn(contribution, __dmul_rn(bi, bj));
}
}
acc=__dadd_rn(acc, contribution);
}
partials[index]=acc;
}
// One level of the strict binary pairing: out[node] = in[2*node] + in[2*node+1],
// with an odd tail CARRIED. Fixed pairing ⇒ association order is a pure function
// of the leaf count.
extern "C" __global__ void sae_arrow_beta_merge(
const double* in_level, int in_nodes, int width,
unsigned long long total, double* out_level)
{
unsigned long long index=(unsigned long long)blockIdx.x*blockDim.x+threadIdx.x;
if(index>=total) return;
int elem=(int)(index%(unsigned long long)width);
int node=(int)(index/(unsigned long long)width);
int left=2*node;
int right=left+1;
double value=in_level[(long long)left*width+elem];
if(right<in_nodes){
value=__dadd_rn(value, in_level[(long long)right*width+elem]);
}
out_level[index]=value;
}
"#;