11/*
2- * SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION.
2+ * SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved .
33 * SPDX-License-Identifier: Apache-2.0
44 */
55
@@ -40,6 +40,7 @@ cuvs::preprocessing::pca::params to_cpp_params(const cuvsPcaParams& c_params)
4040 return cpp_params;
4141}
4242
43+ template <typename LayoutT>
4344void _fit (cuvsResources_t res,
4445 const cuvsPcaParams& params,
4546 DLManagedTensor* input_tensor,
@@ -54,12 +55,12 @@ void _fit(cuvsResources_t res,
5455 auto res_ptr = reinterpret_cast <raft::resources*>(res);
5556 auto cpp_params = to_cpp_params (params);
5657
57- using matrix_type = raft::device_matrix_view<float , int64_t , raft::col_major >;
58+ using matrix_type = raft::device_matrix_view<float , int64_t , LayoutT >;
5859 using vector_type = raft::device_vector_view<float , int64_t >;
5960 using scalar_type = raft::device_scalar_view<float , int64_t >;
6061
61- auto input = cuvs::core::from_dlpack<matrix_type>(input_tensor);
62- auto components = cuvs::core::from_dlpack<matrix_type>(components_tensor);
62+ auto input = cuvs::core::from_dlpack<matrix_type>(input_tensor);
63+ auto components = cuvs::core::from_dlpack<matrix_type>(components_tensor);
6364 auto explained_var = cuvs::core::from_dlpack<vector_type>(explained_var_tensor);
6465 auto explained_var_ratio = cuvs::core::from_dlpack<vector_type>(explained_var_ratio_tensor);
6566 auto singular_vals = cuvs::core::from_dlpack<vector_type>(singular_vals_tensor);
@@ -78,6 +79,7 @@ void _fit(cuvsResources_t res,
7879 flip_signs_based_on_U);
7980}
8081
82+ template <typename LayoutT>
8183void _fit_transform (cuvsResources_t res,
8284 const cuvsPcaParams& params,
8385 DLManagedTensor* input_tensor,
@@ -93,13 +95,13 @@ void _fit_transform(cuvsResources_t res,
9395 auto res_ptr = reinterpret_cast <raft::resources*>(res);
9496 auto cpp_params = to_cpp_params (params);
9597
96- using matrix_type = raft::device_matrix_view<float , int64_t , raft::col_major >;
98+ using matrix_type = raft::device_matrix_view<float , int64_t , LayoutT >;
9799 using vector_type = raft::device_vector_view<float , int64_t >;
98100 using scalar_type = raft::device_scalar_view<float , int64_t >;
99101
100- auto input = cuvs::core::from_dlpack<matrix_type>(input_tensor);
101- auto trans_input = cuvs::core::from_dlpack<matrix_type>(trans_input_tensor);
102- auto components = cuvs::core::from_dlpack<matrix_type>(components_tensor);
102+ auto input = cuvs::core::from_dlpack<matrix_type>(input_tensor);
103+ auto trans_input = cuvs::core::from_dlpack<matrix_type>(trans_input_tensor);
104+ auto components = cuvs::core::from_dlpack<matrix_type>(components_tensor);
103105 auto explained_var = cuvs::core::from_dlpack<vector_type>(explained_var_tensor);
104106 auto explained_var_ratio = cuvs::core::from_dlpack<vector_type>(explained_var_ratio_tensor);
105107 auto singular_vals = cuvs::core::from_dlpack<vector_type>(singular_vals_tensor);
@@ -119,6 +121,7 @@ void _fit_transform(cuvsResources_t res,
119121 flip_signs_based_on_U);
120122}
121123
124+ template <typename LayoutT>
122125void _transform (cuvsResources_t res,
123126 const cuvsPcaParams& params,
124127 DLManagedTensor* input_tensor,
@@ -130,7 +133,7 @@ void _transform(cuvsResources_t res,
130133 auto res_ptr = reinterpret_cast <raft::resources*>(res);
131134 auto cpp_params = to_cpp_params (params);
132135
133- using matrix_type = raft::device_matrix_view<float , int64_t , raft::col_major >;
136+ using matrix_type = raft::device_matrix_view<float , int64_t , LayoutT >;
134137 using vector_type = raft::device_vector_view<float , int64_t >;
135138
136139 auto input = cuvs::core::from_dlpack<matrix_type>(input_tensor);
@@ -143,6 +146,7 @@ void _transform(cuvsResources_t res,
143146 *res_ptr, cpp_params, input, components, singular_vals, mu, trans_input);
144147}
145148
149+ template <typename LayoutT>
146150void _inverse_transform (cuvsResources_t res,
147151 const cuvsPcaParams& params,
148152 DLManagedTensor* trans_input_tensor,
@@ -154,7 +158,7 @@ void _inverse_transform(cuvsResources_t res,
154158 auto res_ptr = reinterpret_cast <raft::resources*>(res);
155159 auto cpp_params = to_cpp_params (params);
156160
157- using matrix_type = raft::device_matrix_view<float , int64_t , raft::col_major >;
161+ using matrix_type = raft::device_matrix_view<float , int64_t , LayoutT >;
158162 using vector_type = raft::device_vector_view<float , int64_t >;
159163
160164 auto trans_input = cuvs::core::from_dlpack<matrix_type>(trans_input_tensor);
@@ -205,19 +209,32 @@ extern "C" cuvsError_t cuvsPcaFit(cuvsResources_t res,
205209 " PCA input must be float32 (kDLFloat, 32 bits)" );
206210 RAFT_EXPECTS (cuvs::core::is_dlpack_device_compatible (input->dl_tensor ),
207211 " PCA input must be device-accessible memory" );
208- RAFT_EXPECTS (cuvs::core::is_f_contiguous (input),
209- " PCA input must be col-major (Fortran-contiguous)" );
210-
211- _fit (res,
212- *params,
213- input,
214- components,
215- explained_var,
216- explained_var_ratio,
217- singular_vals,
218- mu,
219- noise_vars,
220- flip_signs_based_on_U);
212+
213+ if (cuvs::core::is_f_contiguous (input)) {
214+ _fit<raft::col_major>(res,
215+ *params,
216+ input,
217+ components,
218+ explained_var,
219+ explained_var_ratio,
220+ singular_vals,
221+ mu,
222+ noise_vars,
223+ flip_signs_based_on_U);
224+ } else if (cuvs::core::is_c_contiguous (input)) {
225+ _fit<raft::row_major>(res,
226+ *params,
227+ input,
228+ components,
229+ explained_var,
230+ explained_var_ratio,
231+ singular_vals,
232+ mu,
233+ noise_vars,
234+ flip_signs_based_on_U);
235+ } else {
236+ RAFT_FAIL (" PCA input must be contiguous (C- or F-order)" );
237+ }
221238 });
222239}
223240
@@ -239,20 +256,34 @@ extern "C" cuvsError_t cuvsPcaFitTransform(cuvsResources_t res,
239256 " PCA input must be float32 (kDLFloat, 32 bits)" );
240257 RAFT_EXPECTS (cuvs::core::is_dlpack_device_compatible (input->dl_tensor ),
241258 " PCA input must be device-accessible memory" );
242- RAFT_EXPECTS (cuvs::core::is_f_contiguous (input),
243- " PCA input must be col-major (Fortran-contiguous)" );
244-
245- _fit_transform (res,
246- *params,
247- input,
248- trans_input,
249- components,
250- explained_var,
251- explained_var_ratio,
252- singular_vals,
253- mu,
254- noise_vars,
255- flip_signs_based_on_U);
259+
260+ if (cuvs::core::is_f_contiguous (input)) {
261+ _fit_transform<raft::col_major>(res,
262+ *params,
263+ input,
264+ trans_input,
265+ components,
266+ explained_var,
267+ explained_var_ratio,
268+ singular_vals,
269+ mu,
270+ noise_vars,
271+ flip_signs_based_on_U);
272+ } else if (cuvs::core::is_c_contiguous (input)) {
273+ _fit_transform<raft::row_major>(res,
274+ *params,
275+ input,
276+ trans_input,
277+ components,
278+ explained_var,
279+ explained_var_ratio,
280+ singular_vals,
281+ mu,
282+ noise_vars,
283+ flip_signs_based_on_U);
284+ } else {
285+ RAFT_FAIL (" PCA input must be contiguous (C- or F-order)" );
286+ }
256287 });
257288}
258289
@@ -270,10 +301,14 @@ extern "C" cuvsError_t cuvsPcaTransform(cuvsResources_t res,
270301 " PCA input must be float32 (kDLFloat, 32 bits)" );
271302 RAFT_EXPECTS (cuvs::core::is_dlpack_device_compatible (input->dl_tensor ),
272303 " PCA input must be device-accessible memory" );
273- RAFT_EXPECTS (cuvs::core::is_f_contiguous (input),
274- " PCA input must be col-major (Fortran-contiguous)" );
275304
276- _transform (res, *params, input, components, singular_vals, mu, trans_input);
305+ if (cuvs::core::is_f_contiguous (input)) {
306+ _transform<raft::col_major>(res, *params, input, components, singular_vals, mu, trans_input);
307+ } else if (cuvs::core::is_c_contiguous (input)) {
308+ _transform<raft::row_major>(res, *params, input, components, singular_vals, mu, trans_input);
309+ } else {
310+ RAFT_FAIL (" PCA input must be contiguous (C- or F-order)" );
311+ }
277312 });
278313}
279314
@@ -291,9 +326,15 @@ extern "C" cuvsError_t cuvsPcaInverseTransform(cuvsResources_t res,
291326 " PCA trans_input must be float32 (kDLFloat, 32 bits)" );
292327 RAFT_EXPECTS (cuvs::core::is_dlpack_device_compatible (trans_input->dl_tensor ),
293328 " PCA trans_input must be device-accessible memory" );
294- RAFT_EXPECTS (cuvs::core::is_f_contiguous (trans_input),
295- " PCA trans_input must be col-major (Fortran-contiguous)" );
296329
297- _inverse_transform (res, *params, trans_input, components, singular_vals, mu, output);
330+ if (cuvs::core::is_f_contiguous (trans_input)) {
331+ _inverse_transform<raft::col_major>(
332+ res, *params, trans_input, components, singular_vals, mu, output);
333+ } else if (cuvs::core::is_c_contiguous (trans_input)) {
334+ _inverse_transform<raft::row_major>(
335+ res, *params, trans_input, components, singular_vals, mu, output);
336+ } else {
337+ RAFT_FAIL (" PCA trans_input must be contiguous (C- or F-order)" );
338+ }
298339 });
299340}
0 commit comments