Skip to content

Commit

Permalink
BUG: remove sample parameter from pca call to mean (#5980)
Browse files Browse the repository at this point in the history
This fixes a call to `raft::stats::mean()` deactivating the sample parameter in the pca code. 

CC @dantegd

Authors:
  - Malte Förster (https://github.com/mfoerste4)

Approvers:
  - Bradley Dice (https://github.com/bdice)
  - Dante Gama Dessavre (https://github.com/dantegd)
  - Corey J. Nolet (https://github.com/cjnolet)

URL: #5980
  • Loading branch information
mfoerste4 authored Jul 26, 2024
1 parent d1e97b0 commit 3e73853
Show file tree
Hide file tree
Showing 2 changed files with 5 additions and 5 deletions.
2 changes: 1 addition & 1 deletion cpp/src/pca/pca.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -108,7 +108,7 @@ void pcaFit(const raft::handle_t& handle,
auto n_components = prms.n_components;
if (n_components > prms.n_cols) n_components = prms.n_cols;

raft::stats::mean(mu, input, prms.n_cols, prms.n_rows, true, false, stream);
raft::stats::mean(mu, input, prms.n_cols, prms.n_rows, false, false, stream);

auto len = prms.n_cols * prms.n_cols;
rmm::device_uvector<math_t> cov(len, stream);
Expand Down
8 changes: 4 additions & 4 deletions cpp/src/tsvd/tsvd.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -271,21 +271,21 @@ void tsvdFitTransform(const raft::handle_t& handle,

rmm::device_uvector<math_t> mu_trans(prms.n_components, stream);
raft::stats::mean(
mu_trans.data(), trans_input, prms.n_components, prms.n_rows, true, false, stream);
mu_trans.data(), trans_input, prms.n_components, prms.n_rows, false, false, stream);
raft::stats::vars(explained_var,
trans_input,
mu_trans.data(),
prms.n_components,
prms.n_rows,
true,
false,
false,
stream);

rmm::device_uvector<math_t> mu(prms.n_cols, stream);
rmm::device_uvector<math_t> vars(prms.n_cols, stream);

raft::stats::mean(mu.data(), input, prms.n_cols, prms.n_rows, true, false, stream);
raft::stats::vars(vars.data(), input, mu.data(), prms.n_cols, prms.n_rows, true, false, stream);
raft::stats::mean(mu.data(), input, prms.n_cols, prms.n_rows, false, false, stream);
raft::stats::vars(vars.data(), input, mu.data(), prms.n_cols, prms.n_rows, false, false, stream);

rmm::device_scalar<math_t> total_vars(stream);
raft::stats::sum(total_vars.data(), vars.data(), std::size_t(1), prms.n_cols, false, stream);
Expand Down

0 comments on commit 3e73853

Please sign in to comment.