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

Function Train

catboost/libs/train_lib/cross_validation.cpp:240–303  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

238}
239
240void Train(
241 const NCatboostOptions::TCatBoostOptions& catboostOption,
242 const TString& trainDir,
243 const TMaybe<TCustomObjectiveDescriptor>& objectiveDescriptor,
244 const TMaybe<TCustomMetricDescriptor>& evalMetricDescriptor,
245 const TLabelConverter& labelConverter,
246 const TVector<THolder<IMetric>>& metrics,
247 bool isErrorTrackerActive,
248 ITrainingCallbacks* trainingCallbacks,
249 TFoldContext* foldContext,
250 IModelTrainer* modelTrainer,
251 NPar::ILocalExecutor* localExecutor
252) {
253 TTrainModelInternalOptions internalOptions;
254 internalOptions.CalcMetricsOnly = !foldContext->FullModel.Defined();
255 internalOptions.ForceCalcEvalMetricOnEveryIteration = isErrorTrackerActive;
256
257 auto foldOutputOptions = foldContext->OutputOptions;
258 foldOutputOptions.SetTrainDir(trainDir);
259 if (foldContext->FullModel.Defined()) {
260 // TrainModel saves model either to memory pointed by dstModel, or to ResultModelPath
261 foldOutputOptions.ResultModelPath = NCatboostOptions::TOption<TString>("result_model_file", "model");
262 }
263 TMetricsAndTimeLeftHistory metricsAndTimeHistory;
264 const auto defaultCustomCallbacks = MakeHolder<TCustomCallbacks>(Nothing());
265 modelTrainer->TrainModel(
266 internalOptions,
267 catboostOption,
268 foldOutputOptions,
269 objectiveDescriptor,
270 evalMetricDescriptor,
271 foldContext->TrainingData,
272 /*precomputedSingleOnlineCtrDataForSingleFold*/ Nothing(),
273 labelConverter,
274 trainingCallbacks,
275 defaultCustomCallbacks.Get(),
276 /*initModel*/ Nothing(),
277 THolder<TLearnProgress>(),
278 /*initModelApplyCompatiblePools*/ TDataProviders(),
279 localExecutor,
280 /*rand*/ Nothing(),
281 foldContext->FullModel.Defined() ? foldContext->FullModel.Get() : nullptr,
282 TVector<TEvalResult*>{&foldContext->LastUpdateEvalResult},
283 &metricsAndTimeHistory,
284 (foldContext->TaskType == ETaskType::CPU) ? &foldContext->LearnProgress : nullptr
285 );
286 if (foldContext->FullModel.Defined()) {
287 TFileOutput modelFile(JoinFsPaths(trainDir, foldContext->OutputOptions.ResultModelPath.Get()));
288 foldContext->FullModel->Save(&modelFile);
289 }
290 const auto skipMetricOnTrain = GetSkipMetricOnTrain(metrics);
291 for (const auto& trainMetrics : metricsAndTimeHistory.LearnMetricsHistory) {
292 foldContext->MetricValuesOnTrain.emplace_back(GetMetricValues(metrics, skipMetricOnTrain, trainMetrics));
293 }
294 for (const auto& testMetrics : metricsAndTimeHistory.TestMetricsHistory) {
295 CB_ENSURE(testMetrics.size() <= 1, "Expect only one test dataset");
296 if (!testMetrics.empty()) {
297 foldContext->MetricValuesOnTest.emplace_back(GetMetricValues(metrics, /*skipMetric*/{}, testMetrics[0]));

Callers 1

CrossValidateFunction · 0.70

Calls 12

NothingFunction · 0.85
JoinFsPathsFunction · 0.85
GetSkipMetricOnTrainFunction · 0.85
GetMetricValuesFunction · 0.85
SetTrainDirMethod · 0.80
DefinedMethod · 0.45
TrainModelMethod · 0.45
GetMethod · 0.45
SaveMethod · 0.45
emplace_backMethod · 0.45
sizeMethod · 0.45
emptyMethod · 0.45

Tested by

no test coverage detected