diff --git a/cpp/include/raft/random/detail/make_blobs.cuh b/cpp/include/raft/random/detail/make_blobs.cuh index f119d4b39c..a5893c3a26 100644 --- a/cpp/include/raft/random/detail/make_blobs.cuh +++ b/cpp/include/raft/random/detail/make_blobs.cuh @@ -67,21 +67,22 @@ DI void get_mu_sigma(DataT& mu, cid = idx % n_rows; fid = idx / n_rows; } - IdxT center_id; + IdxT cluster_id; if (cid < n_rows) { - center_id = labels[cid]; + cluster_id = labels[cid]; } else { - center_id = 0; + cluster_id = 0; } if (fid >= n_cols) { fid = 0; } + IdxT center_id; if (row_major) { - center_id = center_id * n_cols + fid; + center_id = cluster_id * n_cols + fid; } else { - center_id += fid * n_clusters; + center_id = cluster_id + fid * n_clusters; } - sigma = cluster_std == nullptr ? cluster_std_scalar : cluster_std[cid]; + sigma = cluster_std == nullptr ? cluster_std_scalar : cluster_std[cluster_id]; mu = centers[center_id]; } diff --git a/cpp/tests/random/make_blobs.cu b/cpp/tests/random/make_blobs.cu index 662b39d6a7..338d1a33e1 100644 --- a/cpp/tests/random/make_blobs.cu +++ b/cpp/tests/random/make_blobs.cu @@ -15,6 +15,10 @@ #include +#include +#include +#include + namespace raft { namespace random { @@ -212,5 +216,101 @@ INSTANTIATE_TEST_CASE_P(MakeBlobsTests, MakeBlobsTestD_RowMajor, ::testing::Valu TEST_P(MakeBlobsTestD_ColMajor, Result) { check(); } INSTANTIATE_TEST_CASE_P(MakeBlobsTests, MakeBlobsTestD_ColMajor, ::testing::ValuesIn(inputsd_t)); +template +void check_cluster_std_by_label() +{ + // More rows than clusters exercises the per-cluster vector beyond its first two sample rows. + constexpr int n_rows = 65; + constexpr int n_cols = 2; + constexpr int n_clusters = 2; + raft::resources handle; + auto stream = resource::get_cuda_stream(handle).get(); + auto data = make_device_matrix(handle, n_rows, n_cols); + auto labels = make_device_vector(handle, n_rows); + auto control = make_device_matrix(handle, n_rows, n_cols); + auto control_labels = make_device_vector(handle, n_rows); + auto centers = make_device_matrix(handle, n_clusters, n_cols); + auto cluster_std = make_device_vector(handle, n_clusters); + + std::vector host_centers(n_clusters * n_cols, T(0)); + std::vector host_std{T(0), T(0.8)}; + raft::update_device(centers.data_handle(), host_centers.data(), host_centers.size(), stream); + raft::update_device(cluster_std.data_handle(), host_std.data(), host_std.size(), stream); + + make_blobs(handle, + data.view(), + labels.view(), + n_clusters, + std::make_optional(centers.view()), + std::make_optional(cluster_std.view()), + T(1), + false, + T(-10), + T(10), + 1234ULL, + raft::random::GenPC); + make_blobs(handle, + control.view(), + control_labels.view(), + n_clusters, + std::make_optional(centers.view()), + std::nullopt, + T(1), + false, + T(-10), + T(10), + 1234ULL, + raft::random::GenPC); + + std::vector host_data(n_rows * n_cols); + std::vector host_control(n_rows * n_cols); + std::vector host_labels(n_rows); + std::vector host_control_labels(n_rows); + raft::update_host(host_data.data(), data.data_handle(), host_data.size(), stream); + raft::update_host(host_control.data(), control.data_handle(), host_control.size(), stream); + raft::update_host(host_labels.data(), labels.data_handle(), host_labels.size(), stream); + raft::update_host( + host_control_labels.data(), control_labels.data_handle(), host_control_labels.size(), stream); + resource::sync_stream(handle); + + bool second_cluster_varies = false; + constexpr bool row_major = std::is_same::value; + for (int row = 0; row < n_rows; ++row) { + ASSERT_EQ(host_labels[row], row % n_clusters); + ASSERT_EQ(host_labels[row], host_control_labels[row]); + for (int col = 0; col < n_cols; ++col) { + auto offset = row_major ? row * n_cols + col : col * n_rows + row; + auto expected = host_control[offset] * host_std[host_labels[row]]; + if (host_labels[row] == 0) { + EXPECT_EQ(host_data[offset], T(0)); + } else { + second_cluster_varies |= host_data[offset] != T(0); + } + EXPECT_NEAR(host_data[offset], expected, T(1e-5)); + } + } + EXPECT_TRUE(second_cluster_varies); +} + +TEST(MakeBlobsClusterStd, FloatRowMajor) +{ + check_cluster_std_by_label(); +} + +TEST(MakeBlobsClusterStd, FloatColMajor) +{ + check_cluster_std_by_label(); +} + +TEST(MakeBlobsClusterStd, DoubleRowMajor) +{ + check_cluster_std_by_label(); +} + +TEST(MakeBlobsClusterStd, DoubleColMajor) +{ + check_cluster_std_by_label(); +} + } // end namespace random } // end namespace raft