| 344 | } |
| 345 | |
| 346 | void TestAbsoluteErrorLeaf(const Context* ctx) { |
| 347 | bst_target_t constexpr kTargets = 3, kRows = 16; |
| 348 | std::unique_ptr<ObjFunction> obj{ObjFunction::Create("reg:absoluteerror", ctx)}; |
| 349 | obj->Configure({}); |
| 350 | |
| 351 | MetaInfo info; |
| 352 | info.num_row_ = kRows; |
| 353 | info.labels.Reshape(16, kTargets); |
| 354 | HostDeviceVector<float> predt(info.labels.Size()); |
| 355 | |
| 356 | for (bst_target_t t{0}; t < kTargets; ++t) { |
| 357 | auto h_labels = info.labels.HostView().Slice(linalg::All(), t); |
| 358 | std::iota(linalg::begin(h_labels), linalg::end(h_labels), .0f); |
| 359 | |
| 360 | auto h_predt = |
| 361 | linalg::MakeTensorView(ctx, predt.HostSpan(), kRows, kTargets).Slice(linalg::All(), t); |
| 362 | for (size_t i = 0; i < h_predt.Size(); ++i) { |
| 363 | h_predt(i) = h_labels(i) + i; |
| 364 | } |
| 365 | |
| 366 | HostDeviceVector<bst_node_t> position(h_labels.Size(), 0); |
| 367 | auto& h_position = position.HostVector(); |
| 368 | for (int32_t i = 0; i < 3; ++i) { |
| 369 | h_position[i] = ~i; // negation for sampled nodes. |
| 370 | } |
| 371 | for (size_t i = 3; i < 8; ++i) { |
| 372 | h_position[i] = 3; |
| 373 | } |
| 374 | // empty leaf for node 4 |
| 375 | for (size_t i = 8; i < 13; ++i) { |
| 376 | h_position[i] = 5; |
| 377 | } |
| 378 | for (size_t i = 13; i < h_labels.Size(); ++i) { |
| 379 | h_position[i] = 6; |
| 380 | } |
| 381 | |
| 382 | RegTree tree; |
| 383 | tree.ExpandNode(0, /*split_index=*/1, 2, true, 0.0f, 2.f, 3.f, 4.f, 2.f, 1.f, 1.f); |
| 384 | tree.ExpandNode(1, /*split_index=*/1, 2, true, 0.0f, 2.f, 3.f, 4.f, 2.f, 1.f, 1.f); |
| 385 | tree.ExpandNode(2, /*split_index=*/1, 2, true, 0.0f, 2.f, 3.f, 4.f, 2.f, 1.f, 1.f); |
| 386 | ASSERT_EQ(tree.GetNumLeaves(), 4); |
| 387 | |
| 388 | auto empty_leaf = tree[4].LeafValue(); |
| 389 | |
| 390 | tree::TrainParam param; |
| 391 | param.Init(Args{}); |
| 392 | auto lr = param.learning_rate; |
| 393 | |
| 394 | obj->UpdateTreeLeaf(position, info, lr, predt, t, &tree); |
| 395 | ASSERT_EQ(tree[3].LeafValue(), -5.0f * lr); |
| 396 | ASSERT_EQ(tree[4].LeafValue(), empty_leaf * lr); |
| 397 | ASSERT_EQ(tree[5].LeafValue(), -10.0f * lr); |
| 398 | ASSERT_EQ(tree[6].LeafValue(), -14.0f * lr); |
| 399 | } |
| 400 | } |
| 401 | |
| 402 | void TestVectorLeafObj(Context const* ctx, std::string name, Args const& args, bst_idx_t n_samples, |
| 403 | bst_idx_t n_target_labels, std::vector<float> const& sol_left, |
no test coverage detected