| 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, |
no test coverage detected