Write history metrics to loggers, error trackers and get info from per iteration metric based callback.
| 246 | |
| 247 | // Write history metrics to loggers, error trackers and get info from per iteration metric based callback. |
| 248 | static void ProcessHistoryMetrics( |
| 249 | const TTrainingDataProviders& data, |
| 250 | const TLearnContext& ctx, |
| 251 | ITrainingCallbacks* trainingCallbacks, |
| 252 | TMetricsData* metricsData, |
| 253 | TLoggingData* loggingData, |
| 254 | bool* continueTraining) { |
| 255 | |
| 256 | loggingData->LearnToken = GetTrainModelLearnToken(); |
| 257 | loggingData->TestTokens = GetTrainModelTestTokens(data.Test.ysize()); |
| 258 | |
| 259 | if (ctx.OutputOptions.AllowWriteFiles()) { |
| 260 | InitializeFileLoggers( |
| 261 | ctx.Params, |
| 262 | ctx.Files, |
| 263 | GetConstPointers(metricsData->Metrics), |
| 264 | loggingData->LearnToken, |
| 265 | loggingData->TestTokens, |
| 266 | ctx.OutputOptions.GetMetricPeriod(), |
| 267 | &loggingData->Logger); |
| 268 | } |
| 269 | |
| 270 | const TVector<TTimeInfo>& timeHistory = ctx.LearnProgress->MetricsAndTimeHistory.TimeHistory; |
| 271 | const TVector<TVector<THashMap<TString, double>>>& testMetricsHistory = |
| 272 | ctx.LearnProgress->MetricsAndTimeHistory.TestMetricsHistory; |
| 273 | |
| 274 | const bool useBestModel = ctx.OutputOptions.ShrinkModelToBestIteration(); |
| 275 | *continueTraining = true; |
| 276 | for (int iter : xrange(ctx.LearnProgress->GetCurrentTrainingIterationCount())) { |
| 277 | if (iter < testMetricsHistory.ysize() && ShouldCalcErrorTrackerMetric(iter, *metricsData, ctx) && metricsData->ErrorTracker) { |
| 278 | const TString& errorTrackerMetricDescription = metricsData->Metrics[metricsData->ErrorTrackerMetricIdx]->GetDescription(); |
| 279 | const double error = testMetricsHistory[iter].back().at(errorTrackerMetricDescription); |
| 280 | metricsData->ErrorTracker->AddError(error, iter); |
| 281 | if (useBestModel && iter + 1 >= ctx.OutputOptions.BestModelMinTrees) { |
| 282 | metricsData->BestModelMinTreesTracker->AddError(error, iter); |
| 283 | } |
| 284 | } |
| 285 | |
| 286 | Log(iter, |
| 287 | GetMetricsDescription(metricsData->Metrics), |
| 288 | ctx.LearnProgress->MetricsAndTimeHistory.LearnMetricsHistory, |
| 289 | testMetricsHistory, |
| 290 | metricsData->ErrorTracker ? TMaybe<double>(metricsData->ErrorTracker->GetBestError()) : Nothing(), |
| 291 | metricsData->ErrorTracker ? TMaybe<int>(metricsData->ErrorTracker->GetBestIteration()) : Nothing(), |
| 292 | TProfileResults(timeHistory[iter].PassedTime, timeHistory[iter].RemainingTime), |
| 293 | loggingData->LearnToken, |
| 294 | loggingData->TestTokens, |
| 295 | ShouldCalcAllMetrics(iter, *metricsData, ctx), |
| 296 | &loggingData->Logger |
| 297 | ); |
| 298 | |
| 299 | *continueTraining = trainingCallbacks->IsContinueTraining(ctx.LearnProgress->MetricsAndTimeHistory); |
| 300 | } |
| 301 | |
| 302 | AddConsoleLogger( |
| 303 | loggingData->LearnToken, |
| 304 | loggingData->TestTokens, |
| 305 | /*hasTrain=*/true, |
no test coverage detected