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

Function CrossValidate

catboost/libs/train_lib/cross_validation.cpp:343–554  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

341
342
343void CrossValidate(
344 NJson::TJsonValue plainJsonParams,
345 NCB::TQuantizedFeaturesInfoPtr quantizedFeaturesInfo,
346 const TMaybe<TCustomObjectiveDescriptor>& objectiveDescriptor,
347 const TMaybe<TCustomMetricDescriptor>& evalMetricDescriptor,
348 TLabelConverter& labelConverter,
349 NCB::TDataProviderPtr data,
350 const TCrossValidationParams& cvParams,
351 NPar::ILocalExecutor* localExecutor,
352 TVector<TCVResult>* results,
353 bool isAlreadyShuffled) {
354
355 cvParams.Check();
356
357 NJson::TJsonValue jsonParams;
358 NJson::TJsonValue outputJsonParams;
359 ConvertIgnoredFeaturesFromStringToIndices(data.Get()->MetaInfo, &plainJsonParams);
360 NCatboostOptions::PlainJsonToOptions(plainJsonParams, &jsonParams, &outputJsonParams);
361 ConvertParamsToCanonicalFormat(data.Get()->MetaInfo, &jsonParams);
362 NCatboostOptions::TCatBoostOptions catBoostOptions(NCatboostOptions::LoadOptions(jsonParams));
363 NCatboostOptions::TOutputFilesOptions outputFileOptions;
364 outputFileOptions.Load(outputJsonParams);
365
366 if (catBoostOptions.DataProcessingOptions->ClassLabels->empty()) {
367 catBoostOptions.DataProcessingOptions->ClassLabels = data->MetaInfo.ClassLabels;
368 }
369 ui32 approxDimension = GetApproxDimension(catBoostOptions,
370 labelConverter,
371 data->RawTargetData.GetTargetDimension());
372
373 if (IsYetiRankLossFunction(catBoostOptions.LossFunctionDescription.Get().LossFunction)) {
374 // Can't use standard UpdateYetiRankEvalMetric because for raw data TargetStats might not be available
375 UpdateYetiRankEvalMetric(data, localExecutor, &catBoostOptions);
376 }
377
378 UpdateSampleRateOption(data->ObjectsData->GetObjectCount(), &catBoostOptions);
379
380 InitializeEvalMetricIfNotSet(catBoostOptions.MetricOptions->ObjectiveMetric,
381 &catBoostOptions.MetricOptions->EvalMetric);
382
383 UpdateMetricPeriodOption(catBoostOptions, &outputFileOptions);
384
385 TVector<THolder<IMetric>> metrics = CreateMetrics(
386 catBoostOptions.MetricOptions,
387 evalMetricDescriptor,
388 approxDimension,
389 data->MetaInfo.HasWeights
390 );
391
392 CheckMetrics(metrics, catBoostOptions.LossFunctionDescription.Get().GetLossFunction());
393 CheckCrossValidationOptions(data, metrics, catBoostOptions, outputFileOptions, cvParams);
394
395 const ui64 cpuUsedRamLimit =
396 ParseMemorySizeDescription(catBoostOptions.SystemOptions->CpuUsedRamLimit.Get());
397
398 TRestorableFastRng64 rand(cvParams.PartitionRandSeed);
399 if (cvParams.Shuffle && !isAlreadyShuffled) {
400 auto objectsGroupingSubset = NCB::Shuffle(data->ObjectsGrouping, 1, &rand);

Callers 4

TuneHyperparamsCVFunction · 0.85
GridSearchFunction · 0.85
RandomizedSearchFunction · 0.85
CatBoostCV_RFunction · 0.85

Tested by

no test coverage detected