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

Method TrainModel

catboost/cuda/train_lib/train.cpp:295–495  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

293 class TGPUModelTrainer: public IModelTrainer {
294 public:
295 void TrainModel(
296 const TTrainModelInternalOptions& internalOptions,
297 const NCatboostOptions::TCatBoostOptions& catboostOptions,
298 const NCatboostOptions::TOutputFilesOptions& outputOptions,
299 const TMaybe<TCustomObjectiveDescriptor>& objectiveDescriptor,
300 const TMaybe<TCustomMetricDescriptor>& evalMetricDescriptor,
301 TTrainingDataProviders trainingData,
302 TMaybe<NCB::TPrecomputedOnlineCtrData> precomputedSingleOnlineCtrDataForSingleFold,
303 const TLabelConverter& labelConverter,
304 ITrainingCallbacks* trainingCallbacks,
305 ICustomCallbacks* /*customCallbacks*/,
306 TMaybe<TFullModel*> initModel,
307 THolder<TLearnProgress> initLearnProgress,
308 NCB::TDataProviders initModelApplyCompatiblePools,
309 NPar::ILocalExecutor* localExecutor,
310 const TMaybe<TRestorableFastRng64*> rand,
311 TFullModel* dstModel,
312 const TVector<TEvalResult*>& evalResultPtrs,
313 TMetricsAndTimeLeftHistory* metricsAndTimeHistory,
314 THolder<TLearnProgress>* dstLearnProgress) const override {
315
316 Y_UNUSED(rand);
317 CB_ENSURE(trainingData.Test.size() <= 1, "Multiple eval sets not supported for GPU");
318 CB_ENSURE(!precomputedSingleOnlineCtrDataForSingleFold,
319 "Precomputed online CTR data for GPU is not yet supported");
320 CB_ENSURE(
321 evalResultPtrs.empty() || (evalResultPtrs.size() == trainingData.Test.size()),
322 "Need test dataset to evaluate resulting model");
323 CB_ENSURE(!initModel && !initLearnProgress, "Training continuation for GPU is not yet supported");
324 Y_UNUSED(initModelApplyCompatiblePools);
325 CB_ENSURE_INTERNAL(!dstLearnProgress, "Returning learn progress for GPU is not yet supported");
326
327 NCatboostOptions::TCatBoostOptions updatedCatboostOptions(catboostOptions);
328
329 bool saveFinalCtrsInModel = !internalOptions.CalcMetricsOnly &&
330 (outputOptions.GetFinalCtrComputationMode() == EFinalCtrComputationMode::Default) &&
331 (trainingData.Learn->ObjectsData->GetQuantizedFeaturesInfo()
332 ->CalcMaxCategoricalFeaturesUniqueValuesCountOnLearn()
333 > updatedCatboostOptions.CatFeatureParams.Get().OneHotMaxSize.Get());
334
335 auto quantizedFeaturesInfo = trainingData.Learn->ObjectsData->GetQuantizedFeaturesInfo();
336 TVector<TExclusiveFeaturesBundle> exclusiveBundlesCopy;
337 const auto lossFunction = catboostOptions.LossFunctionDescription->LossFunction;
338 // TODO(kirillovs): check and enable on pairwise losses
339 if (!IsGpuPlainDocParallelOnlyMode(lossFunction)) {
340 exclusiveBundlesCopy.assign(
341 trainingData.Learn->ObjectsData->GetExclusiveFeatureBundlesMetaData().begin(),
342 trainingData.Learn->ObjectsData->GetExclusiveFeatureBundlesMetaData().end()
343 );
344 }
345 ui32 objectsCount = trainingData.Learn->GetObjectCount();
346 if (!trainingData.Test.empty()) {
347 objectsCount += trainingData.Test[0]->GetObjectCount();
348 }
349 TBinarizedFeaturesManager featuresManager(updatedCatboostOptions.CatFeatureParams,
350 trainingData.FeatureEstimators,
351 *trainingData.Learn->MetaInfo.FeaturesLayout,
352 exclusiveBundlesCopy,

Callers 2

TuneHyperparamsTrainTestFunction · 0.45
TrainModelImplFunction · 0.45

Tested by

no test coverage detected