| 148 | void CompareJsonModels(Json l, Json r) { CompareJSON(std::move(l), std::move(r)); } |
| 149 | |
| 150 | void 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 |
no test coverage detected