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

Method train

src/thundersvm/model/svc.cpp:41–145  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

39}
40
41void SVC::train(const DataSet &dataset, SvmParam param) {
42 DataSet dataset_ = dataset;
43 dataset_.group_classes();
44 model_setup(dataset_, param);
45
46 vector<SyncArray<float_type>> alpha(n_binary_models);
47 vector<bool> is_sv(dataset_.n_instances(), false);
48
49 int k = 0;
50 for (int i = 0; i < n_classes; ++i) {
51 for (int j = i + 1; j < n_classes; ++j) {
52 train_binary(dataset_, i, j, alpha[k], rho.host_data()[k]);
53 vector<int> original_index = dataset_.original_index(i, j);
54 CHECK_EQ(original_index.size(), alpha[k].size());
55 const float_type *alpha_data = alpha[k].host_data();
56 for (int l = 0; l < alpha[k].size(); ++l) {
57 is_sv[original_index[l]] = is_sv[original_index[l]] || (alpha_data[l] != 0);
58 }
59 k++;
60 }
61 }
62
63 for (int i = 0; i < dataset_.n_classes(); ++i) {
64 vector<int> original_index = dataset_.original_index(i);
65 DataSet::node2d i_instances = dataset_.instances(i);
66 int *n_sv_data = n_sv.host_data();
67 for (int j = 0; j < i_instances.size(); ++j) {
68 if (is_sv[original_index[j]]) {
69 n_sv_data[i]++;
70 sv.push_back(i_instances[j]);
71 sv_indices.push_back(original_index[j]);
72 }
73 }
74 }
75
76 n_total_sv = sv.size();
77 LOG(INFO) << "#total unique sv = " << n_total_sv;
78 coef.resize((n_classes - 1) * n_total_sv);
79
80 vector<int> sv_start(1, 0);
81 const int *n_sv_data = n_sv.host_data();
82 for (int i = 1; i < n_classes; ++i) {
83 sv_start.push_back(sv_start[i - 1] + n_sv_data[i - 1]);
84 }
85
86 k = 0;
87 float_type *coef_data = coef.host_data();
88 for (int i = 0; i < n_classes; ++i) {
89 for (int j = i + 1; j < n_classes; ++j) {
90 const float_type *alpha_data = alpha[k].host_data();
91 vector<int> original_index = dataset_.original_index(i, j);
92 int ci = dataset_.count()[i];
93 int cj = dataset_.count()[j];
94 int m = sv_start[i];
95 for (int l = 0; l < ci; ++l) {
96 if (is_sv[original_index[l]]) {
97 coef_data[(j - 1) * n_total_sv + m++] = alpha_data[l];
98 }

Callers 12

TESTFunction · 0.45
train_RFunction · 0.45
sparse_model_scikitFunction · 0.45
dense_model_scikitFunction · 0.45
thundersvm_train_subFunction · 0.45
mainFunction · 0.45
probability_trainMethod · 0.45
cross_validationMethod · 0.45

Calls 10

group_classesMethod · 0.80
original_indexMethod · 0.80
n_classesMethod · 0.80
instancesMethod · 0.80
resizeMethod · 0.80
is_zero_basedMethod · 0.80
n_instancesMethod · 0.45
host_dataMethod · 0.45
sizeMethod · 0.45
n_featuresMethod · 0.45

Tested by 5

TESTFunction · 0.36