// fitrgam.cpp// 自包含 GAM 实现循环梯度提升 决策树桩// 逻辑对应 MATLAB fitrgam// - 每个预测变量独立提升cyclic boosting// - 基学习器为决策树桩 (max_depth 1)// - 默认迭代 300 次学习率 0.1// - 预测 截距 Σ 形状函数#includeiostream#includevector#includealgorithm#includecmath#includerandom#includenumeric#includelimits// // 决策树桩 (Decision Stump)// structStump{doublethreshold0.0;doubleleft_value0.0;doubleright_value0.0;// 预测x threshold 走左子树否则走右子树doublepredict(doublex)const{return(xthreshold)?left_value:right_value;}};// 用平方误差损失训练一棵树桩Stumptrain_stump(conststd::vectordoublex,conststd::vectordoubleresidual){constintnstatic_castint(x.size());Stump s;// 1) 按 x 排序std::vectorintidx(n);std::iota(idx.begin(),idx.end(),0);std::sort(idx.begin(),idx.end(),[](inta,intb){returnx[a]x[b];});std::vectordoublexs(n),rs(n);for(inti0;in;i){xs[i]x[idx[i]];rs[i]residual[idx[i]];}// 2) 前缀和用于 O(1) 计算任意切分的 SSEstd::vectordoubleprefix_sum(n1,0.0);std::vectordoubleprefix_sq(n1,0.0);for(inti0;in;i){prefix_sum[i1]prefix_sum[i]rs[i];prefix_sq[i1]prefix_sq[i]rs[i]*rs[i];}constdoubletotal_sumprefix_sum[n];constdoubletotal_sqprefix_sq[n];// 3) 枚举所有可能的分裂点doublebest_ssestd::numeric_limitsdouble::infinity();intbest_split-1;for(inti1;in;i){if(xs[i-1]xs[i])continue;// 相同值无法分裂constintnli;constintnrn-i;constdoublesum_lprefix_sum[i];constdoublesum_rtotal_sum-sum_l;constdoublesq_lprefix_sq[i];constdoublesq_rtotal_sq-sq_l;// SSE Σy² - (Σy)²/nconstdoublesse(sq_l-sum_l*sum_l/nl)(sq_r-sum_r*sum_r/nr);if(ssebest_sse){best_ssesse;best_spliti;}}// 4) 退化情形所有 x 相同无法分裂 预测常数均值if(best_split0){constdoublemeantotal_sum/n;s.threshold0.0;s.left_valuemean;s.right_valuemean;returns;}// 5) 计算左右叶子的均值constintnlbest_split;constintnrn-best_split;constdoublesum_lprefix_sum[best_split];constdoublesum_rtotal_sum-sum_l;s.threshold0.5*(xs[best_split-1]xs[best_split]);s.left_valuesum_l/nl;s.right_valuesum_r/nr;returns;}// // GAM 模型// classGAM{public:intn_iter300;// 与 fitrgam 默认值一致doublelearning_rate0.1;// 与 fitrgam 默认值一致doubleintercept0.0;std::vectorstd::vectorStumpshapes;voidfit(conststd::vectorstd::vectordoubleX,conststd::vectordoubley,boolverbosetrue){constintnstatic_castint(y.size());constintpstatic_castint(X[0].size());// 1) 截距初始化 响应变量均值intercept0.0;for(doublev:y)interceptv;intercept/n;// 2) 形状函数初始化为 0shapes.assign(p,{});std::vectorstd::vectordoublef(p,std::vectordouble(n,0.0));// 3) 循环梯度提升for(intiter0;itern_iter;iter){for(intj0;jp;j){// 3a) 残差: r_i y_i - (intercept Σ_k f_k)当前全量预测的负梯度// 注意必须包含 f_j 自身否则增量更新 f_j lr*h 会重复累计// 同一特征信号导致循环梯度提升发散GAMBoost 标准残差定义std::vectordoubleresid(n);for(inti0;in;i){doublepartialintercept;for(intk0;kp;k){partialf[k][i];}resid[i]y[i]-partial;}// 3b) 用树桩拟合残差std::vectordoublexj(n);for(inti0;in;i)xj[i]X[i][j];Stump strain_stump(xj,resid);shapes[j].push_back(s);// 3c) 更新形状函数: f_j lr * h_j(x_j)for(inti0;in;i){f[j][i]learning_rate*s.predict(xj[i]);}}// 3d) 打印训练进度if(verbose((iter1)%300||iter0)){doublermse0.0;for(inti0;in;i){constdoubleey[i]-predict(X[i]);rmsee*e;}rmsestd::sqrt(rmse/n);std::coutIteration (iter1) | RMSE: rmse\n;}}}// 单样本预测y_pred intercept Σ_j f_j(x_j)doublepredict(conststd::vectordoublex)const{doublesintercept;for(size_t j0;jshapes.size();j){for(constStumpst:shapes[j]){slearning_rate*st.predict(x[j]);}}returns;}// 查询某个形状函数在给定点的取值doubleshape_value(intj,doublexj)const{doubles0.0;for(constStumpst:shapes[j]){slearning_rate*st.predict(xj);}returns;}};// // 主程序// intmain(){// 1. 生成模拟数据constintn200;constintp2;std::mt19937rng(42);std::normal_distributiondoublenoise(0.0,0.1);std::uniform_real_distributiondoubleunif(-1.0,1.0);std::vectorstd::vectordoubleX(n,std::vectordouble(p));std::vectordoubley(n);for(inti0;in;i){X[i][0]unif(rng);X[i][1]unif(rng);// 真实模型: y 2 3*x1 sin(x2) εy[i]2.03.0*X[i][0]std::sin(X[i][1])noise(rng);}// 2. 打印设置std::cout--- GAM Training (Gradient Boosting, Stumps) ---\n;std::coutObservations: n\n;std::coutPredictors: p\n;std::coutBoosting Iterations (M): 300\n;std::coutLearning Rate: 0.1\n\n;// 3. 训练GAM gam;gam.n_iter300;gam.learning_rate0.1;gam.fit(X,y,/*verbose*/true);// 4. 输出模型信息std::cout\n--- Training Complete ---\n;std::coutIntercept: gam.intercept\n;for(intj0;jp;j){std::coutShape Function f(j1): gam.shapes[j].size() stumps\n;}// 5. 对新查询点预测constdoubleqx10.5,qx2-0.2;std::vectordoubleqx{qx1,qx2};constdoublepf1gam.shape_value(0,qx1);constdoublepf2gam.shape_value(1,qx2);constdoublepred_ygam.predict(qx);std::cout\nPrediction for query point (x1qx1, x2qx2):\n;std::cout Intercept: gam.intercept\n;std::cout f1(qx1): pf1\n;std::cout f2(qx2): pf2\n;std::cout Predicted y: pred_y\n;constdoubletrue_y2.03.0*qx1std::sin(qx2);std::cout\n True y (noiseless): true_y\n;std::cout Error: std::abs(pred_y-true_y)\n;return0;}