@@ -142,7 +142,7 @@ inline std::enable_if_t<std::is_floating_point_v<MathT>> predict_core(
142142 &beta,
143143 distances.data (),
144144 n_clusters,
145- stream);
145+ stream. get () );
146146
147147 auto distances_const_view = raft::make_device_matrix_view<const MathT, IdxT, raft::row_major>(
148148 distances.data (), n_rows, n_clusters);
@@ -286,13 +286,29 @@ void calc_centers_and_sizes(const raft::resources& handle,
286286
287287 // Apply mapping only when the data and math types are different.
288288 if constexpr (std::is_same_v<T, MathT>) {
289- raft::linalg::reduce_rows_by_key (
290- dataset, dim, labels, nullptr , n_rows, dim, n_clusters, centers, stream, reset_counters);
289+ raft::linalg::reduce_rows_by_key (dataset,
290+ dim,
291+ labels,
292+ nullptr ,
293+ n_rows,
294+ dim,
295+ n_clusters,
296+ centers,
297+ stream.get (),
298+ reset_counters);
291299 } else {
292300 // todo(lsugy): use iterator from KV output of fusedL2NN
293301 thrust::transform_iterator<MappingOpT, const T*> mapping_itr (dataset, mapping_op);
294- raft::linalg::reduce_rows_by_key (
295- mapping_itr, dim, labels, nullptr , n_rows, dim, n_clusters, centers, stream, reset_counters);
302+ raft::linalg::reduce_rows_by_key (mapping_itr,
303+ dim,
304+ labels,
305+ nullptr ,
306+ n_rows,
307+ dim,
308+ n_clusters,
309+ centers,
310+ stream.get (),
311+ reset_counters);
296312 }
297313
298314 // Compute weight of each cluster
@@ -689,38 +705,39 @@ auto adjust_centers(const raft::resources& handle,
689705 search_count.set_value_to_zero_async (stream);
690706 const dim3 grid_dim (raft::ceildiv (n_clusters, static_cast <IdxT>(kBlockDimY )), 1 , 1 );
691707 adjust_centers_random_donor_kernel<kBlockDimY >
692- <<<grid_dim, block_dim, 0 , stream>>> (centers,
693- n_clusters,
694- dim,
695- dataset,
696- n_rows,
697- labels,
698- cluster_sizes,
699- lower_threshold,
700- static_cast <IdxT>(n_rows / n_clusters),
701- centroid_offset,
702- ofst,
703- search_count.data (),
704- update_count.data (),
705- mapping_op);
708+ <<<grid_dim, block_dim, 0 , stream.get() >>> (centers,
709+ n_clusters,
710+ dim,
711+ dataset,
712+ n_rows,
713+ labels,
714+ cluster_sizes,
715+ lower_threshold,
716+ static_cast <IdxT>(n_rows / n_clusters),
717+ centroid_offset,
718+ ofst,
719+ search_count.data (),
720+ update_count.data (),
721+ mapping_op);
706722 return update_count.value (stream) > 0 ; // NB: rmm scalar performs the sync
707723 }
708724
709725 raft::update_device (receiver_clusters.data (), host_receiver_clusters.data (), n_pairs, stream);
710726 raft::update_device (donor_clusters.data (), host_donor_clusters.data (), n_pairs, stream);
711727 const dim3 grid_dim (raft::ceildiv (n_pairs, static_cast <IdxT>(kBlockDimY )), 1 , 1 );
712- adjust_centers_kernel<kBlockDimY ><<<grid_dim, block_dim, 0 , stream>>> (centers,
713- n_pairs,
714- dim,
715- dataset,
716- n_rows,
717- labels,
718- receiver_clusters.data (),
719- donor_clusters.data (),
720- centroid_offset,
721- ofst,
722- update_count.data (),
723- mapping_op);
728+ adjust_centers_kernel<kBlockDimY >
729+ <<<grid_dim, block_dim, 0 , stream.get()>>> (centers,
730+ n_pairs,
731+ dim,
732+ dataset,
733+ n_rows,
734+ labels,
735+ receiver_clusters.data (),
736+ donor_clusters.data (),
737+ centroid_offset,
738+ ofst,
739+ update_count.data (),
740+ mapping_op);
724741 auto n_updates = update_count.value (stream); // NB: rmm scalar performs the sync
725742 RAFT_EXPECTS (n_updates == n_pairs, " Balanced k-means failed to update all adjusted centers" );
726743 return n_updates > 0 ;
@@ -1068,7 +1085,7 @@ auto build_fine_clusters(const raft::resources& handle,
10681085 }
10691086
10701087 thrust::transform_iterator<MappingOpT, const T*> mapping_itr (dataset_mptr, mapping_op);
1071- raft::matrix::gather (mapping_itr, dim, n_rows, mc_trainset_ids, k, mc_trainset, stream);
1088+ raft::matrix::gather (mapping_itr, dim, n_rows, mc_trainset_ids, k, mc_trainset, stream. get () );
10721089 if (params.metric == cuvs::distance::DistanceType::L2Expanded ||
10731090 params.metric == cuvs::distance::DistanceType::L2SqrtExpanded ||
10741091 params.metric == cuvs::distance::DistanceType::CosineExpanded) {
0 commit comments