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

Method probability_train

src/thundersvm/model/svc.cpp:414–479  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

412}
413
414void SVC::probability_train(const DataSet &dataset) {
415 SvmParam param_no_prob = param;
416 param_no_prob.probability = 0;
417
418 vector<float_type> dec_predict_all(dataset.n_instances() * n_binary_models);
419
420 //cross-validation dec_values
421 int n_fold = 5;
422 for (int k = 0; k < n_fold; ++k) {
423 SvmModel *temp_model = new SVC();
424 DataSet::node2d x_train, x_test;
425 vector<float_type> y_train, y_test;
426 vector<int> test_idx;
427 for (int i = 0; i < dataset.n_classes(); ++i) {
428 int fold_test_count = dataset.count()[i] / n_fold;
429 vector<int> class_idx = dataset.original_index(i);
430 auto idx_begin = class_idx.begin() + fold_test_count * k;
431 auto idx_end = idx_begin;
432 if (k == n_fold - 1) {
433 idx_end = class_idx.end();
434 } else {
435 while (idx_end != class_idx.end() && idx_end - idx_begin < fold_test_count) idx_end++;
436 }
437 for (int j: vector<int>(idx_begin, idx_end)) {
438 x_test.push_back(dataset.instances()[j]);
439 y_test.push_back(dataset.y()[j]);
440 test_idx.push_back(j);
441 }
442 class_idx.erase(idx_begin, idx_end);
443 for (int j:class_idx) {
444 x_train.push_back(dataset.instances()[j]);
445 y_train.push_back(dataset.y()[j]);
446 }
447 }
448 DataSet train_dataset(x_train, dataset.n_features(), y_train);
449 temp_model->train(train_dataset, param_no_prob);
450 SyncArray<float_type> dec_predict(x_test.size() * n_binary_models);
451 temp_model->predict_dec_values(x_test, dec_predict, 1000);
452 float_type *dec_predict_data = dec_predict.host_data();
453 for (int i = 0; i < x_test.size(); ++i) {
454 memcpy(&dec_predict_all[test_idx[i] * n_binary_models], &dec_predict_data[i * n_binary_models],
455 sizeof(float_type) * n_binary_models);
456 }
457 delete temp_model;
458 }
459 int k = 0;
460 for (int i = 0; i < n_classes; ++i) {
461 for (int j = i + 1; j < n_classes; ++j) {
462 vector<int> ori_idx;
463 vector<int> y;
464 vector<float_type> dec_values_subproblem;
465 ori_idx = dataset.original_index(i);
466 for (int l = 0; l < dataset.count()[i]; ++l) {
467 y.push_back(+1);
468 dec_values_subproblem.push_back(dec_predict_all[ori_idx[l] * n_binary_models + k]);
469 }
470 ori_idx = dataset.original_index(j);
471 for (int l = 0; l < dataset.count()[j]; ++l) {

Callers

nothing calls this directly

Calls 13

sigmoidTrainFunction · 0.85
n_classesMethod · 0.80
original_indexMethod · 0.80
endMethod · 0.80
instancesMethod · 0.80
predict_dec_valuesMethod · 0.80
dataMethod · 0.80
n_instancesMethod · 0.45
beginMethod · 0.45
n_featuresMethod · 0.45
trainMethod · 0.45
sizeMethod · 0.45

Tested by

no test coverage detected