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

Function Train

catboost/libs/train_lib/train_model.cpp:376–546  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

374}
375
376static void Train(
377 const TTrainModelInternalOptions& internalOptions,
378 const TTrainingDataProviders& data,
379 ITrainingCallbacks* trainingCallbacks,
380 ICustomCallbacks* customCallbacks,
381 TLearnContext* ctx,
382 TVector<TVector<TVector<double>>>* testMultiApprox // [test][dim][docIdx], can be nullptr if not needed
383) {
384 TProfileInfo& profile = ctx->Profile;
385
386 TMetricsData metricsData;
387 InitializeAndCheckMetricData(internalOptions, data, *ctx, &metricsData);
388
389 const auto onLoadSnapshotCallback = [&] (IInputStream* in) {
390 return trainingCallbacks->OnLoadSnapshot(in);
391 };
392
393 const bool progressLoaded = ctx->TryLoadProgress(onLoadSnapshotCallback);
394
395 TLoggingData loggingData;
396 bool continueTraining;
397 ProcessHistoryMetrics(data, *ctx, trainingCallbacks, &metricsData, &loggingData, &continueTraining);
398
399 const bool useBestModel = ctx->OutputOptions.ShrinkModelToBestIteration();
400 const bool hasTest = data.GetTestSampleCount() > 0;
401 const auto& metrics = metricsData.Metrics;
402 auto& errorTracker = metricsData.ErrorTracker;
403
404 if (progressLoaded && ctx->Params.SystemOptions->IsMaster()) {
405 MapRestoreApproxFromTreeStruct(
406 (hasTest && useBestModel) ? MakeMaybe(errorTracker->GetBestIteration()) : Nothing(),
407 ctx);
408 }
409
410 if (continueTraining) {
411 InitializeSamplingStructures(data, ctx);
412 }
413
414 THPTimer timer;
415
416 const auto onSaveSnapshotCallback = [&] (IOutputStream* out) {
417 trainingCallbacks->OnSaveSnapshot(NJson::TJsonValue{}, out);
418 };
419
420 for (ui32 iter = ctx->LearnProgress->GetCurrentTrainingIterationCount();
421 continueTraining && (iter < ctx->Params.BoostingOptions->IterationCount);
422 ++iter)
423
424 {
425 if (errorTracker && errorTracker->GetIsNeedStop()) {
426 LogThatStoppingOccured(*errorTracker);
427 break;
428 }
429
430 profile.StartNextIteration();
431
432 if (timer.Passed() > ctx->OutputOptions.GetSnapshotSaveInterval()) {
433 ctx->SaveProgress(onSaveSnapshotCallback);

Callers 2

TrainModelMethod · 0.70
EvaluateFeaturesImplFunction · 0.70

Calls 15

ProcessHistoryMetricsFunction · 0.85
MakeMaybeFunction · 0.85
NothingFunction · 0.85
LogThatStoppingOccuredFunction · 0.85
TrainOneIterationFunction · 0.85
HasInvalidValuesFunction · 0.85
CalcErrorsFunction · 0.85
MapFindPtrFunction · 0.85

Tested by

no test coverage detected