| 30 | } // anonymous namespace |
| 31 | |
| 32 | TEST(Approx, Partitioner) { |
| 33 | size_t n_samples = 1024, n_features = 1, base_rowid = 0; |
| 34 | Context ctx; |
| 35 | ctx.InitAllowUnknown(Args{}); |
| 36 | CommonRowPartitioner partitioner{&ctx, n_samples, base_rowid, false}; |
| 37 | ASSERT_EQ(partitioner.base_rowid, base_rowid); |
| 38 | ASSERT_EQ(partitioner.Size(), 1); |
| 39 | ASSERT_EQ(partitioner.Partitions()[0].Size(), n_samples); |
| 40 | |
| 41 | auto const Xy = RandomDataGenerator{n_samples, n_features, 0}.GenerateDMatrix(true); |
| 42 | auto hess = GenerateHess(n_samples); |
| 43 | std::vector<CPUExpandEntry> candidates{{0, 0}}; |
| 44 | candidates.front().split.loss_chg = 0.4; |
| 45 | |
| 46 | for (auto const& page : Xy->GetBatches<GHistIndexMatrix>(&ctx, {64, hess, true})) { |
| 47 | bst_feature_t const split_ind = 0; |
| 48 | { |
| 49 | auto min_value = -std::numeric_limits<float>::infinity(); |
| 50 | RegTree tree; |
| 51 | CommonRowPartitioner partitioner{&ctx, n_samples, base_rowid, false}; |
| 52 | GetSplit(&tree, min_value, &candidates); |
| 53 | partitioner.UpdatePosition(&ctx, page, candidates, tree.HostScView()); |
| 54 | ASSERT_EQ(partitioner.Size(), 3); |
| 55 | ASSERT_EQ(partitioner[1].Size(), 0); |
| 56 | ASSERT_EQ(partitioner[2].Size(), n_samples); |
| 57 | } |
| 58 | { |
| 59 | CommonRowPartitioner partitioner{&ctx, n_samples, base_rowid, false}; |
| 60 | auto ptr = page.cut.Ptrs()[split_ind + 1]; |
| 61 | float split_value = page.cut.Values().at(ptr / 2); |
| 62 | RegTree tree; |
| 63 | GetSplit(&tree, split_value, &candidates); |
| 64 | partitioner.UpdatePosition(&ctx, page, candidates, tree.HostScView()); |
| 65 | |
| 66 | { |
| 67 | auto left_nidx = tree[RegTree::kRoot].LeftChild(); |
| 68 | auto const& elem = partitioner[left_nidx]; |
| 69 | ASSERT_LT(elem.Size(), n_samples); |
| 70 | ASSERT_GT(elem.Size(), 1); |
| 71 | for (auto& it : elem) { |
| 72 | auto value = page.cut.Values().at(page.index[it]); |
| 73 | ASSERT_LE(value, split_value); |
| 74 | } |
| 75 | } |
| 76 | { |
| 77 | auto right_nidx = tree[RegTree::kRoot].RightChild(); |
| 78 | auto const& elem = partitioner[right_nidx]; |
| 79 | for (auto& it : elem) { |
| 80 | auto value = page.cut.Values().at(page.index[it]); |
| 81 | ASSERT_GT(value, split_value) << it; |
| 82 | } |
| 83 | } |
| 84 | } |
| 85 | } |
| 86 | } |
| 87 | |
| 88 | TEST(Approx, InteractionConstraint) { |
| 89 | auto constexpr kRows = 32; |
nothing calls this directly
no test coverage detected