| 453 | } |
| 454 | |
| 455 | void CalcErrorsDistributed( |
| 456 | const NCB::TTrainingDataProviders& trainData, |
| 457 | const TVector<THolder<IMetric>>& metrics, |
| 458 | bool calcAllMetrics, |
| 459 | bool calcErrorTrackerMetric, |
| 460 | TLearnContext* ctx) { |
| 461 | |
| 462 | Y_ASSERT(ctx->Params.SystemOptions->IsMaster()); |
| 463 | |
| 464 | // Calc additive stats in distributed manner and calc non-additive stats locally |
| 465 | |
| 466 | bool haveNonAdditiveTrainMetrics = false; |
| 467 | bool haveNonAdditiveTestMetrics = false; // TODO(akhropov): do we need per test dataset flags? |
| 468 | |
| 469 | IterateOverMetrics( |
| 470 | trainData, |
| 471 | metrics, |
| 472 | calcAllMetrics, |
| 473 | calcErrorTrackerMetric, |
| 474 | /*calcAdditiveMetrics*/ false, |
| 475 | /*calcNonAdditiveMetrics*/ true, |
| 476 | /*onLearn*/ [&] (TConstArrayRef<const IMetric*> trainMetrics) { |
| 477 | haveNonAdditiveTrainMetrics = !trainMetrics.empty(); |
| 478 | }, |
| 479 | /*onTest*/ [&] ( |
| 480 | size_t /*testIdx*/, |
| 481 | TConstArrayRef<const IMetric*> testMetrics, |
| 482 | TMaybe<int> /*filteredTrackerIdx*/ |
| 483 | ) { |
| 484 | haveNonAdditiveTestMetrics |= (!testMetrics.empty()); |
| 485 | } |
| 486 | ); |
| 487 | |
| 488 | |
| 489 | // Compute non-additive metrics locally and distributed additive stats in parallel |
| 490 | |
| 491 | TVector<std::function<void()>> tasks; |
| 492 | |
| 493 | if (haveNonAdditiveTrainMetrics || haveNonAdditiveTestMetrics) { |
| 494 | MapGetApprox( |
| 495 | trainData, |
| 496 | *ctx, |
| 497 | /*useBestModel*/ false, |
| 498 | haveNonAdditiveTrainMetrics ? &(ctx->LearnProgress->AvrgApprox) : nullptr, |
| 499 | haveNonAdditiveTestMetrics ? &(ctx->LearnProgress->TestApprox) : nullptr |
| 500 | ); |
| 501 | |
| 502 | tasks.push_back( |
| 503 | [&] () { |
| 504 | CalcErrorsLocally( |
| 505 | trainData, |
| 506 | metrics, |
| 507 | calcAllMetrics, |
| 508 | calcErrorTrackerMetric, |
| 509 | /*calcNonAdditiveMetricsOnly*/true, |
| 510 | ctx |
| 511 | ); |
| 512 | } |
no test coverage detected