| 412 | } |
| 413 | |
| 414 | void 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) { |
nothing calls this directly
no test coverage detected