MCPcopy Create free account
hub / github.com/catboost/catboost / ProcessHistoryMetrics

Function ProcessHistoryMetrics

catboost/libs/train_lib/train_model.cpp:248–310  ·  view source on GitHub ↗

Write history metrics to loggers, error trackers and get info from per iteration metric based callback.

Source from the content-addressed store, hash-verified

246
247// Write history metrics to loggers, error trackers and get info from per iteration metric based callback.
248static 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,

Callers 1

TrainFunction · 0.85

Calls 15

GetTrainModelLearnTokenFunction · 0.85
GetTrainModelTestTokensFunction · 0.85
InitializeFileLoggersFunction · 0.85
GetConstPointersFunction · 0.85
xrangeFunction · 0.85
GetMetricsDescriptionFunction · 0.85
NothingFunction · 0.85
ShouldCalcAllMetricsFunction · 0.85
AddConsoleLoggerFunction · 0.85
AllowWriteFilesMethod · 0.80
GetMetricPeriodMethod · 0.80

Tested by

no test coverage detected