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

Function TestAbsoluteErrorLeaf

tests/cpp/objective/test_regression_obj.cc:346–400  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

344}
345
346void 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
402void 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,

Callers 2

TESTFunction · 0.85
TESTFunction · 0.85

Calls 15

AllFunction · 0.85
beginFunction · 0.85
endFunction · 0.85
MakeTensorViewFunction · 0.85
ReshapeMethod · 0.80
HostSpanMethod · 0.80
ExpandNodeMethod · 0.80
GetNumLeavesMethod · 0.80
ConfigureMethod · 0.45
SizeMethod · 0.45
SliceMethod · 0.45
HostViewMethod · 0.45

Tested by

no test coverage detected