MCPcopy Create free account
hub / github.com/dmlc/xgboost / TestLearnerSerialization

Function TestLearnerSerialization

tests/cpp/test_serialization.cc:150–290  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

148void CompareJsonModels(Json l, Json r) { CompareJSON(std::move(l), std::move(r)); }
149
150void TestLearnerSerialization(Args args, FeatureMap const& fmap, std::shared_ptr<DMatrix> p_dmat) {
151 for (auto& batch : p_dmat->GetBatches<SparsePage>()) {
152 batch.data.HostVector();
153 batch.offset.HostVector();
154 }
155
156 int32_t constexpr kIters = 2;
157
158 common::TemporaryDirectory tempdir;
159 std::string const fname = tempdir.Str() + "/model";
160
161 std::vector<std::string> dumped_0;
162 std::string model_at_kiter;
163
164 // Train for kIters.
165 {
166 std::unique_ptr<dmlc::Stream> fo(dmlc::Stream::Create(fname.c_str(), "w"));
167 std::unique_ptr<Learner> learner{Learner::Create({p_dmat})};
168 learner->SetParams(args);
169 for (int32_t iter = 0; iter < kIters; ++iter) {
170 learner->UpdateOneIter(iter, p_dmat);
171 }
172 dumped_0 = learner->DumpModel(fmap, true, "json");
173 learner->Save(fo.get());
174
175 common::MemoryBufferStream mem_out(&model_at_kiter);
176 learner->Save(&mem_out);
177 }
178
179 // Assert dumped model is same after loading
180 std::vector<std::string> dumped_1;
181 {
182 std::unique_ptr<dmlc::Stream> fi(dmlc::Stream::Create(fname.c_str(), "r"));
183 std::unique_ptr<Learner> learner{Learner::Create({p_dmat})};
184 learner->Load(fi.get());
185 learner->Configure();
186 dumped_1 = learner->DumpModel(fmap, true, "json");
187 }
188 ASSERT_EQ(dumped_0, dumped_1);
189
190 std::string model_at_2kiter;
191
192 // Test training continuation with data from host
193 {
194 std::string continued_model;
195 {
196 // Continue the previous training with another kIters
197 std::unique_ptr<dmlc::Stream> fi(dmlc::Stream::Create(fname.c_str(), "r"));
198 std::unique_ptr<Learner> learner{Learner::Create({p_dmat})};
199 learner->Load(fi.get());
200 learner->Configure();
201
202 // verify the loaded model doesn't change.
203 std::string serialised_model_tmp;
204 common::MemoryBufferStream mem_out(&serialised_model_tmp);
205 learner->Save(&mem_out);
206 ASSERT_EQ(model_at_kiter, serialised_model_tmp);
207

Callers 1

TEST_FFunction · 0.85

Calls 13

CompareJSONFunction · 0.85
c_strMethod · 0.80
SetParamsMethod · 0.80
UpdateOneIterMethod · 0.80
SetParamMethod · 0.80
StrMethod · 0.45
DumpModelMethod · 0.45
SaveMethod · 0.45
getMethod · 0.45
LoadMethod · 0.45
ConfigureMethod · 0.45
SetDeviceMethod · 0.45

Tested by

no test coverage detected