| 172 | |
| 173 | namespace { |
| 174 | void TestWeightedMultiSampling(Context const* ctx) { |
| 175 | size_t constexpr kCols = 32; |
| 176 | HostDeviceVector<float> feature_weights(kCols, 0); |
| 177 | auto& h_feature_weights = feature_weights.HostVector(); |
| 178 | for (size_t i = 0; i < h_feature_weights.size(); ++i) { |
| 179 | h_feature_weights[i] = i; |
| 180 | } |
| 181 | ColumnSampler cs; |
| 182 | float bytree{0.5}, bylevel{0.5}, bynode{0.5}; |
| 183 | cs.Init(ctx, h_feature_weights.size(), feature_weights, bytree, bylevel, bynode); |
| 184 | auto feature_set = cs.GetFeatureSet(ctx, 0); |
| 185 | size_t n_sampled = kCols * bytree * bylevel * bynode; |
| 186 | ASSERT_EQ(feature_set->Size(), n_sampled); |
| 187 | feature_set = cs.GetFeatureSet(ctx, 1); |
| 188 | ASSERT_EQ(feature_set->Size(), n_sampled); |
| 189 | } |
| 190 | } // namespace |
| 191 | |
| 192 | TEST(ColumnSampler, WeightedMultiSampling) { |
no test coverage detected