A C Implementation of IRLS Algorithm for Solving Logit/Logistic Regression
#include <iostream>
#include <gsl/gsl_blas.h>
#include <gsl/gsl_linalg.h>
#include <gsl/gsl_matrix.h>
#include <cmath>
#include <algorithm>
inline double linkinv(double t){
return 1-1/(1+exp(t));
}
inline double gprime(double t){
return 1/(exp(t/2)+exp(-t/2))/(exp(t/2)+exp(-t/2));
}
inline double gprime_sq(double t){
return gprime(t)*gprime(t);
}
inline double varg(double t){
return t*(1-t);
}
double l2norm(gsl_vector *x){
double sum=0;
for(size_t i=0;i<x->size;++i){
sum += gsl_vector_get(x,i)*gsl_vector_get(x,i);
}
return sqrt(sum);
}
void print_vec(gsl_vector *x, size_t n, std::string x_lab){
std::cout << x_lab << ":\t";
for(size_t i=0; i<std::min(n,x->size); ++i){
std::cout << gsl_vector_get(x,i)<<",";
}
std::cout << std::endl;
}
void print_mat(gsl_matrix *x, size_t n1, size_t n2, std::string x_lab){
std::cout << x_lab << ":\t";
for(size_t i=0; i<std::min(n1,x->size1); ++i){
for(size_t j=0; j<std::min(n2,x->size2); ++j){
std::cout << gsl_matrix_get(x,i,j)<<",";
}
std::cout << std::endl;
}
}
int logistic_regression(gsl_matrix *orig_x, gsl_vector *y, gsl_vector **beta, gsl_vector **yhat, size_t max_iter, double eps, bool *converge){
gsl_set_error_handler_off();
size_t p = orig_x->size2;
size_t n = orig_x->size1;
if (n!=y->size){
return 1;
}
gsl_matrix *x = gsl_matrix_alloc(n,p);
gsl_matrix_memcpy(x,orig_x);
*beta = gsl_vector_alloc(p);
*yhat = gsl_vector_alloc(n);
gsl_vector *t = gsl_vector_alloc(n);
gsl_vector_set_zero(t);
gsl_vector *work_vec = gsl_vector_alloc(p);
gsl_matrix *work_mat = gsl_matrix_alloc(p,p);
gsl_matrix *v = gsl_matrix_alloc(p,p);
gsl_vector *d = gsl_vector_alloc(p);
gsl_linalg_SV_decomp_mod(x,work_mat,v,d,work_vec);
gsl_vector *z = gsl_vector_alloc(n);
gsl_matrix *W = gsl_matrix_alloc(n,n);
gsl_matrix_set_zero(W);
gsl_vector *s = gsl_vector_alloc(p);
gsl_vector_set_zero(s);
gsl_vector *s_old = gsl_vector_alloc(p);
gsl_matrix *UWU = gsl_matrix_alloc(p,p);
gsl_matrix *WU = gsl_matrix_alloc(n,p);
gsl_vector *Wz = gsl_vector_alloc(n);
gsl_vector *UWz = gsl_vector_alloc(p);
gsl_matrix *SVD_V = gsl_matrix_alloc(p,p);
gsl_vector *SVD_S = gsl_vector_alloc(p);
for(size_t step=0; step < max_iter; ++step){
double ti;
for(size_t i=0;i<n;++i){
ti = gsl_vector_get(t,i);
gsl_vector_set(z,i,ti+(gsl_vector_get(y,i)-linkinv(ti))/gprime(ti));
gsl_matrix_set(W,i,i,gprime_sq(ti)/varg(linkinv(ti)));
}
gsl_vector_memcpy(s_old,s);
gsl_matrix_set_zero(WU);
gsl_matrix_set_zero(UWU);
gsl_blas_dgemm(CblasNoTrans,CblasNoTrans,1.0,W,x,0.0,WU);
gsl_blas_dgemm(CblasTrans,CblasNoTrans,1.0,x,WU,0.0,UWU);
int chol_s = gsl_linalg_cholesky_decomp(UWU);
gsl_blas_dgemv(CblasNoTrans,1.0,W,z,0.0,Wz);
gsl_blas_dgemv(CblasTrans,1.0,x,Wz,0.0,UWz);
if(chol_s){
gsl_linalg_SV_decomp_jacobi (UWU,SVD_V,SVD_S);
gsl_linalg_SV_solve(UWU, SVD_V, SVD_S, UWz, s);
}else{
gsl_linalg_cholesky_solve(UWU,UWz,s);
}
gsl_blas_dgemv(CblasNoTrans,1.0,x,s,0.0,t);
gsl_vector_sub (s_old, s);
if (l2norm(s_old)<eps){
break;
}
}
gsl_matrix_free(work_mat);
gsl_matrix_free(W);
gsl_matrix_free(UWU);
gsl_matrix_free(WU);
gsl_matrix_free(SVD_V);
gsl_vector_free(work_vec);
gsl_vector_free(s);
gsl_vector_free(z);
gsl_vector_free(s_old);
gsl_vector_free(Wz);
gsl_vector_free(UWz);
gsl_vector_free(SVD_S);
int status = gsl_linalg_SV_solve(x,v,d,t,*beta);
gsl_matrix_free(x);
gsl_matrix_free(v);
gsl_vector_free(d);
gsl_vector_free(t);
if(status){
return 1;
}
gsl_blas_dgemv(CblasNoTrans,1.0,orig_x,*beta,0.0,*yhat);
for (size_t i=0;i<(*yhat)->size;++i){
gsl_vector_set(*yhat,i,linkinv(gsl_vector_get(*yhat,i)));
}
return 0;
}