MCPcopy Create free account
hub / github.com/Xtra-Computing/thundersvm / train

Method train

src/thundersvm/model/nusvr.cpp:7–41  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

5#include <thundersvm/solver/nusmosolver.h>
6
7void NuSVR::train(const DataSet &dataset, SvmParam param) {
8 model_setup(dataset, param);
9 int n_instances = dataset.n_instances();
10
11 //duplicate instances
12 DataSet::node2d instances_2(dataset.instances());
13 instances_2.insert(instances_2.end(), dataset.instances().begin(), dataset.instances().end());
14
15 KernelMatrix kernelMatrix(instances_2, param);
16
17 SyncArray<float_type> f_val(n_instances * 2);
18 SyncArray<int> y(n_instances * 2);
19
20 SyncArray<float_type> alpha_2(n_instances * 2);
21 float_type *f_val_data = f_val.host_data();
22 int *y_data = y.host_data();
23 float_type *alpha_2_data = alpha_2.host_data();
24 float_type sum = param.C * param.nu * n_instances / 2;
25 for (int i = 0; i < n_instances; ++i) {
26 alpha_2_data[i] = alpha_2_data[i + n_instances] = min(sum, param.C);
27 sum -= alpha_2_data[i];
28 f_val_data[i] = f_val_data[i + n_instances] = -dataset.y()[i];
29 y_data[i] = +1;
30 y_data[i + n_instances] = -1;
31 }
32
33 int ws_size = get_working_set_size(n_instances * 2, kernelMatrix.n_features());
34 NuSMOSolver solver(true);
35 solver.solve(kernelMatrix, y, alpha_2, rho.host_data()[0], f_val, param.epsilon, param.C, param.C, ws_size, max_iter);
36 save_svr_coef(alpha_2, dataset.instances());
37
38 if(param.kernel_type == SvmParam::LINEAR){
39 compute_linear_coef_single_model(dataset.n_features(), dataset.is_zero_based());
40 }
41}
42
43void NuSVR::model_setup(const DataSet &dataset, SvmParam &param) {
44 SVR::model_setup(dataset, param);

Callers

nothing calls this directly

Calls 9

minFunction · 0.85
instancesMethod · 0.80
endMethod · 0.80
solveMethod · 0.80
is_zero_basedMethod · 0.80
n_instancesMethod · 0.45
beginMethod · 0.45
host_dataMethod · 0.45
n_featuresMethod · 0.45

Tested by

no test coverage detected