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

Function UpdateUndefinedRandomSeed

catboost/private/libs/algo/preprocess.cpp:52–89  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

50}
51
52void UpdateUndefinedRandomSeed(
53 ETaskType taskType,
54 const NCatboostOptions::TOutputFilesOptions& outputOptions,
55 NJson::TJsonValue* updatedJsonParams,
56 std::function<void(TIFStream*, TString&)> paramsLoader) {
57
58 const TString snapshotFilename = TOutputFiles::AlignFilePath(
59 outputOptions.GetTrainDir(),
60 outputOptions.GetSnapshotFilename(),
61 /*namePrefix=*/ ""
62 );
63 if (outputOptions.SaveSnapshot() && NFs::Exists(snapshotFilename)) {
64 TString serializedTrainParams;
65 NJson::TJsonValue restoredJsonParams;
66 try {
67 TProgressHelper(ToString(taskType)).CheckedLoad(
68 snapshotFilename,
69 [&](TIFStream* inputStream) {
70 paramsLoader(inputStream, serializedTrainParams);
71 }
72 );
73 ReadJsonTree(serializedTrainParams, &restoredJsonParams);
74 CB_ENSURE(restoredJsonParams.Has("random_seed"), "Snapshot is broken.");
75 } catch (const TCatBoostException&) {
76 throw;
77 } catch (...) {
78 CATBOOST_WARNING_LOG << "Can't load progress from snapshot file: " << snapshotFilename <<
79 " Exception: " << CurrentExceptionMessage() << Endl;
80 return;
81 }
82
83 if (!(*updatedJsonParams)["flat_params"].Has("random_seed") &&
84 !restoredJsonParams["flat_params"].Has("random_seed"))
85 {
86 (*updatedJsonParams)["random_seed"] = restoredJsonParams["random_seed"];
87 }
88 }
89}
90
91void UpdateUndefinedClassLabels(
92 const TVector<NJson::TJsonValue>& classLabels,

Callers 2

TrainModelFunction · 0.85
ModelBasedEvalFunction · 0.85

Calls 7

TProgressHelperClass · 0.85
CurrentExceptionMessageFunction · 0.85
SaveSnapshotMethod · 0.80
CheckedLoadMethod · 0.80
ToStringFunction · 0.50
ReadJsonTreeFunction · 0.50
HasMethod · 0.45

Tested by

no test coverage detected