| 267 | }; |
| 268 | |
| 269 | static void UpdateLearningRate(ui32 learnObjectCount, bool useBestModel, NCatboostOptions::TCatBoostOptions* catBoostOptions) { |
| 270 | const bool boostFromAverage = catBoostOptions->BoostingOptions->BoostFromAverage.Get(); |
| 271 | auto& learningRate = catBoostOptions->BoostingOptions->LearningRate; |
| 272 | const int iterationCount = catBoostOptions->BoostingOptions->IterationCount; |
| 273 | const auto lossFunction = catBoostOptions->LossFunctionDescription->GetLossFunction(); |
| 274 | const auto taskType = catBoostOptions->GetTaskType(); |
| 275 | |
| 276 | if ( |
| 277 | learningRate.NotSet() && |
| 278 | catBoostOptions->ObliviousTreeOptions->LeavesEstimationMethod.NotSet() && |
| 279 | catBoostOptions->ObliviousTreeOptions->LeavesEstimationIterations.NotSet() && |
| 280 | catBoostOptions->ObliviousTreeOptions->L2Reg.NotSet() |
| 281 | ) { |
| 282 | TAutoLRParamsGuesser lrGuesser; |
| 283 | if (lrGuesser.NeedToUpdate(taskType, lossFunction, useBestModel, boostFromAverage)) { |
| 284 | learningRate = lrGuesser.GetLearningRate(taskType, lossFunction, useBestModel, boostFromAverage, iterationCount, learnObjectCount); |
| 285 | CATBOOST_NOTICE_LOG << "Learning rate set to " << learningRate << Endl; |
| 286 | } |
| 287 | } |
| 288 | } |
| 289 | |
| 290 | static void UpdateLeavesEstimationIterations( |
| 291 | const NCB::TDataMetaInfo& trainDataMetaInfo, |
no test coverage detected