Skip to main content

legume_numeric/matrix/
traits.rs

1use crate::matrix::common_io::{Delimiter, ReadLinesOut};
2#[cfg(feature = "tensor")]
3use candle_core::{Device, Tensor};
4use num_traits::Float;
5
6/// Trait for running statistics operations
7///
8/// Provides a common interface for both dense (ndarray-based) and
9/// sparse running statistics implementations.
10pub trait RunningStatOps<T>
11where
12    T: Float,
13{
14    type Output;
15
16    fn clear(&mut self);
17    fn count_positives(&self) -> Self::Output;
18    fn sum(&self) -> Self::Output;
19    fn mean(&self) -> Self::Output;
20    fn variance(&self) -> Self::Output;
21    fn std(&self) -> Self::Output;
22}
23
24/// some linear algebra routines
25pub trait RandomizedAlgs {
26    type InMat;
27    type OutMat;
28    type DVec;
29    type Scalar;
30
31    /// randomized singular value decomposition
32    /// # input
33    /// * `X`: `n x d` matrix
34    /// # output
35    /// * `U`: `n x k`
36    /// * `D`: `k x 1`
37    /// * `V`: `d x k`
38    fn rsvd(&self, max_rank: usize) -> anyhow::Result<(Self::OutMat, Self::DVec, Self::OutMat)>;
39}
40
41/// Convert to and from the vector of triplets
42pub trait MatTriplets {
43    type Mat;
44    type Scalar;
45
46    fn from_nonzero_triplets<I>(
47        nrow: usize,
48        ncol: usize,
49        triplets: &[(I, I, Self::Scalar)],
50    ) -> anyhow::Result<Self::Mat>
51    where
52        I: TryInto<usize> + Copy,
53        <I as TryInto<usize>>::Error: std::fmt::Debug;
54
55    fn to_nonzero_triplets(&self) -> anyhow::Result<NRowNColTriplets<Self::Scalar>>;
56}
57
58pub struct NRowNColTriplets<Scalar> {
59    pub nrow: usize,
60    pub ncol: usize,
61    pub triplets: Vec<(usize, usize, Scalar)>,
62}
63
64#[cfg(feature = "tensor")]
65/// Reading off from `Tensor`
66pub trait ConvertMatOps {
67    type Mat;
68    type Scalar;
69
70    fn from_tensor(_: &Tensor) -> anyhow::Result<Self::Mat>;
71    fn to_tensor(&self, dev: &Device) -> anyhow::Result<Tensor>;
72}
73
74/// normalize, sum_to_one, scale, and centre columns
75pub trait MatOps {
76    type Mat;
77    type Scalar;
78
79    /// make each column sum to 1
80    fn sum_to_one_columns_inplace(&mut self);
81    /// make each column sum to 1
82    fn sum_to_one_columns(&self) -> Self::Mat;
83
84    /// make each row sum to 1
85    fn sum_to_one_rows_inplace(&mut self);
86    /// make each row sum to 1
87    fn sum_to_one_rows(&self) -> Self::Mat;
88
89    /// normalize logits after taking exp `(log-sum-exp)`
90    fn normalize_exp_logits_columns_inplace(&mut self);
91    /// normalize logits after taking exp `(log-sum-exp)`
92    fn normalize_exp_logits_columns(&self) -> Self::Mat;
93
94    /// column-wise log-softmax: subtract each column's log-sum-exp so the
95    /// `exp` of each column sums to 1. Returns log-probabilities (unlike
96    /// [`Self::normalize_exp_logits_columns`], which returns probabilities).
97    fn log_softmax_columns_inplace(&mut self);
98    /// column-wise log-softmax (see [`Self::log_softmax_columns_inplace`])
99    fn log_softmax_columns(&self) -> Self::Mat;
100
101    /// vector norm for each column
102    fn normalize_columns_inplace(&mut self);
103    /// vector norm for each column
104    fn normalize_columns(&self) -> Self::Mat;
105
106    /// standardization for each column
107    fn scale_columns_inplace(&mut self);
108    /// standardization for each column
109    fn scale_columns(&self) -> Self::Mat;
110
111    /// standardization for each row
112    fn scale_rows_inplace(&mut self);
113    /// standardization for each row
114    fn scale_rows(&self) -> Self::Mat;
115
116    /// centering for each column
117    fn centre_columns_inplace(&mut self);
118    /// centering for each column
119    fn centre_columns(&self) -> Self::Mat;
120}
121
122pub trait AdjustByDivisionOp<Other, Scalar> {
123    /// Adjust each column with the column of the matching batch index
124    ///
125    /// Assume: `Y[g] ~ Poisson(X[g] * λ)`
126    /// (1) Estimate the λ parameter by taking overall ratio, namely,
127    /// `λ = Σ Y[g] / Σ X[g]`
128    ///
129    /// (2) Take the residual (in the log space)
130    /// `ln Y[g] - ln (λ X[g])` or `Y[g]/λX[g]` if `X[g] > 0`
131    /// otherwise, do nothing
132    fn adjust_by_division_of_selected_inplace(&mut self, denom_db: &Other, batches: &[usize]);
133
134    /// adjust each column with the corresponding column of the denom
135    ///
136    /// Assume: `Y[g] ~ Poisson(X[g] * λ)`
137    /// (1) Estimate the λ parameter by taking overall ratio, namely,
138    /// `λ = Σ Y[g] / Σ X[g]`
139    ///
140    /// (2) Take the residual (in the log space)
141    /// `ln Y[g] - ln (λ X[g])` or `Y[g]/λX[g]` if `X[g] > 0`
142    /// otherwise, do nothing
143    fn adjust_by_division_inplace(&mut self, denom: &Other);
144}
145
146pub trait MatElemOps {
147    type Mat;
148    type Scalar;
149    fn log1p_inplace(&mut self);
150    fn log1p(&self) -> Self::Mat;
151}
152
153#[cfg(feature = "tensor")]
154/// Elementwise chains fused into ONE pass, because candle's CPU backend runs them
155/// one core at a time.
156///
157/// Only matmul reaches `gemm`, which candle drives with `Parallelism::Rayon`;
158/// `unary_map` and `binary_map` are plain serial iterators, and the vectorized
159/// `f32_vec` path is `#[cfg(feature = "mkl" / "accelerate")]` — SIMD, still one
160/// core. So a loop whose matmuls scale across every core stalls on the
161/// elementwise ops between them, and any chain over a large tensor is worth
162/// collapsing into a single rayon pass.
163///
164/// Implemented for `Tensor` on CPU only. Off CPU the device's own kernels are
165/// already parallel, and each method falls back to the op chain it stands in for
166/// — same numbers either way, which the tests assert bitwise.
167pub trait FusedTensorOps: Sized {
168    /// `exp(min(self + offset, ceiling))`, i.e. the Poisson rate from a linear
169    /// predictor with the overflow guard `exp` needs (f32 overflows at 88).
170    ///
171    /// Replaces `self.broadcast_add(offset)?.minimum(ceiling)?.exp()`. `self` is
172    /// `[N, F]`; `offset` is any shape that chain broadcasts against it — `[N, F]`,
173    /// `[1, F]` or `[N, 1]`.
174    ///
175    /// # The receiver must be unaliased
176    ///
177    /// On CPU this overwrites `self`'s storage and hands it back, so the whole
178    /// chain costs one buffer instead of three. `Tensor` is an `Arc`, so taking
179    /// `self` by value does **not** prove exclusivity — `x.reshape(..)` yields a
180    /// contiguous tensor sharing `x`'s storage and would pass every guard here,
181    /// silently overwriting `x`. Pass a freshly computed tensor (a `matmul`
182    /// result), never a `clone`, `narrow` or `reshape` of one still in use.
183    ///
184    /// Back-prop is unsupported by construction (candle's in-place custom ops
185    /// carry no backward), which is why the callers take their gradients in
186    /// closed form.
187    ///
188    /// Deliberately single-offset. A chain carrying **two** offsets — a `[N, 1]`
189    /// column and a `[1, F]` row, as the joint velocity solver in
190    /// `graph-embedding-util` does — needs its own method rather than a caller
191    /// pre-broadcasting one of them, which would cost the full-size op this
192    /// exists to remove.
193    fn clamped_exp_add_inplace(self, offset: &Tensor, ceiling: f64) -> anyhow::Result<Self>;
194}
195
196/// TF-IDF (Term Frequency–Inverse Document Frequency) transformation
197///
198/// A numerical statistic reflecting how important a word (term) is to a document
199/// in a collection or corpus. (Wikipedia)
200///
201/// Treats the matrix as a term-document matrix where:
202/// - Rows are "terms" (e.g., genes, words)
203/// - Columns are "documents" (e.g., cell types, text documents)
204///
205/// **TF-IDF(t, d) = TF(t, d) × IDF(t)**
206///
207/// where:
208/// - TF(t, d) = term frequency of term t in document d (matrix values)
209/// - IDF(t) = log(N / df(t)) = inverse document frequency
210/// - N = total number of documents (columns)
211/// - df(t) = document frequency = number of documents containing term t
212///
213/// Terms appearing in many documents get lower weight; terms specific to few
214/// documents get higher weight.
215pub trait TfIdfOps {
216    type Mat;
217
218    /// Apply TF-IDF transformation
219    ///
220    /// IDF(t) = log(N / (df(t) + 1)) where df(t) = number of non-zero entries in row t
221    fn tfidf(&self) -> Self::Mat;
222
223    /// Apply TF-IDF followed by L2 column normalization
224    ///
225    /// Useful for cosine similarity comparisons between documents (columns)
226    fn tfidf_normalize_columns(&self) -> Self::Mat;
227}
228
229/// Operations to sample random matrices, only works for
230/// `nalgebra::DMatrix` and `ndarray::Array2`
231pub trait SampleOps {
232    type Mat;
233    type Scalar;
234
235    /// Sample a matrix from a uniform distribution `U(0,1)`.
236    ///
237    /// Unseeded: draws fresh entropy each call. For reproducible output use
238    /// [`SampleOps::runif_seeded`].
239    fn runif(dd: usize, nn: usize) -> Self::Mat;
240
241    /// Sample a matrix from a normal distribution `N(0,1)`.
242    ///
243    /// Unseeded: draws fresh entropy each call. For reproducible output use
244    /// [`SampleOps::rnorm_seeded`].
245    fn rnorm(dd: usize, nn: usize) -> Self::Mat;
246
247    /// Sample a matrix from a gamma distribution with `param` is
248    /// `(shape α, scale θ)`
249    ///
250    /// $$f(x|\alpha,\theta) = \frac{\theta^{-\alpha}}{\Gamma(\alpha)} x^{\alpha - 1} e^{-x/\theta}$$
251    ///
252    /// Note: `rate = 1/scale` or $\beta = 1/\theta$
253    ///
254    /// Unseeded: draws fresh entropy each call. For reproducible output use
255    /// [`SampleOps::rgamma_seeded`].
256    fn rgamma(dd: usize, nn: usize, param: (f32, f32)) -> Self::Mat;
257
258    /// Seeded, thread-order-independent `U(0,1)` sample. Byte-identical across
259    /// runs, thread counts, and machines for a fixed `seed`. See
260    /// [`crate::matrix::rand_util`].
261    fn runif_seeded(dd: usize, nn: usize, seed: u64) -> Self::Mat;
262
263    /// Seeded, thread-order-independent `N(0,1)` sample. Byte-identical across
264    /// runs, thread counts, and machines for a fixed `seed`. See
265    /// [`crate::matrix::rand_util`].
266    fn rnorm_seeded(dd: usize, nn: usize, seed: u64) -> Self::Mat;
267
268    /// Seeded, thread-order-independent gamma sample (`param = (shape α, scale θ)`).
269    /// Byte-identical across runs, thread counts, and machines for a fixed
270    /// `seed`. See [`crate::matrix::rand_util`].
271    fn rgamma_seeded(dd: usize, nn: usize, param: (f32, f32), seed: u64) -> Self::Mat;
272}
273
274pub trait DistanceOps {
275    type Scalar;
276    type Other;
277
278    /// A vector of Euclidean distances between sources and targets `other`
279    ///
280    /// * `other`: other data matrix
281    fn euclidean_distance(
282        &self,
283        other: &Self::Other,
284    ) -> anyhow::Result<Vec<(usize, usize, Self::Scalar)>>;
285
286    /// A vector of Euclidean distances between sources and targets `other`
287    ///
288    /// * `other`: other data matrix
289    /// * `select_columns_in_other`: specific columns
290    fn euclidean_distance_on_select_columns(
291        &self,
292        other: &Self::Other,
293        select_columns_in_other: &[usize],
294    ) -> anyhow::Result<Vec<(usize, usize, Self::Scalar)>>;
295}
296
297pub trait EncodingOps
298where
299    Self: Sized,
300{
301    type Mat;
302    type Scalar;
303
304    /// Sinusoidal Positional Encoding
305    /// * `emb_dim` - embedding dimension, say `d`
306    /// * returns each column's embedding results (row x 2d)
307    ///
308    /// for each element r of each column c:
309    ///  ret[r, 2i] = sin(x[r,c]/10000^(2i/d))
310    ///  ret[r, 2i + 1] = cos(x[r,c]/10000^(2i/d))
311    /// where i in [0, d/2-1]
312    fn positional_embedding_columns(&self, emb_dim: usize) -> anyhow::Result<Self::Mat>;
313}
314
315/// Operations that involves multiple types
316pub trait CompositeOps {
317    type Scalar;
318    type Mat;
319    type Other;
320
321    /// `self[:,col] += other[:,col]`
322    /// * `other`: `CscMatrix`
323    /// * `col`: column index
324    fn add_assign_column(&mut self, other: &Self::Other, col: usize);
325
326    /// `self += other`
327    /// * `other`: `CscMatrix`
328    fn add_assign(&mut self, other: &Self::Other);
329}
330
331/// Read and write matrices from and to files
332pub trait IoOps {
333    type Scalar;
334    type Mat;
335
336    fn read_file_delim(
337        file_path: &str,
338        delim: impl Into<Delimiter>,
339        skip: Option<usize>,
340    ) -> anyhow::Result<Self::Mat>;
341
342    /// Read the data matrix with row and column names
343    ///
344    /// * `file_path`: data file name
345    /// * `delim`: delimiter (`char` vector or string)
346    /// * `header_row`: header line (0-based); `None` will find no header
347    /// * `row_name_column_index`: column index (0-based) corresponds to row name
348    /// * `select_column_indices`: column indices (0-based) to include
349    /// * `select_column_names`: column names to include
350    ///
351    fn read_data(
352        file_path: &str,
353        delim: impl Into<Delimiter>,
354        header_row: Option<usize>,
355        row_name_column_index: Option<usize>,
356        select_column_indices: Option<&[usize]>,
357        select_column_names: Option<&[Box<str>]>,
358    ) -> anyhow::Result<MatWithNames<Self::Mat>>;
359
360    /// Read the data matrix with row and column names
361    ///
362    /// * `file_path`: data file name
363    /// * `delim`: delimiter (`char` vector or string)
364    /// * `header_row`: header line (0-based); `None` will find no header
365    /// * `header_column`: column index (0-based) corresponds to row name
366    ///
367    fn read_data_with_names(
368        file_path: &str,
369        delim: impl Into<Delimiter>,
370        header_row: Option<usize>,
371        header_column: Option<usize>,
372    ) -> anyhow::Result<MatWithNames<Self::Mat>> {
373        Self::read_data(file_path, delim, header_row, header_column, None, None)
374    }
375
376    #[allow(clippy::type_complexity)]
377    fn read_data_vec_with_indices_names(
378        file_path: &str,
379        delim: impl Into<Delimiter>,
380        header_line: Option<usize>,
381        row_name_index: Option<usize>,
382        column_indices: Option<&[usize]>,
383        column_names: Option<&[Box<str>]>,
384    ) -> anyhow::Result<(Vec<Box<str>>, Vec<Box<str>>, Vec<Self::Scalar>)>
385    where
386        Self::Scalar: std::str::FromStr,
387        <Self::Scalar as std::str::FromStr>::Err: std::fmt::Debug,
388    {
389        let hdr_line = match header_line {
390            Some(skip) => skip as i64,
391            None => -1, // no skipping
392        };
393
394        let ReadLinesOut { mut lines, header } =
395            crate::matrix::common_io::read_lines_of_words_delim(file_path, delim, hdr_line)?;
396
397        // A blank line tokenizes to one empty field, not to zero fields, so it
398        // must be dropped here or it caps the width check below at 1 and then
399        // breaks the value loop. Dropping it up front fixes both at once.
400        lines.retain(|w| !(w.len() == 1 && w[0].is_empty()));
401
402        let data_width = lines.iter().map(|w| w.len()).min().unwrap_or(header.len());
403        // R's write.table omits a name for the row-label column, so the header
404        // is one field short of the data rows and every header position names
405        // the data column one to its RIGHT. Detect that shape once; both the
406        // name matching and the naming lookup below shift through it.
407        let header_offset = usize::from(
408            !header.is_empty() && header.len() + 1 == data_width && row_name_index == Some(0),
409        );
410
411        let mut relevant_indices: Vec<usize> = vec![];
412
413        let indices_given = column_indices.is_some_and(|ix| !ix.is_empty());
414        if let Some(indices) = column_indices {
415            relevant_indices.extend(indices.iter().copied());
416        }
417
418        // Explicit indices OVERRIDE names, as the callers' help documents; a
419        // union would quietly widen the selection with every default name that
420        // happens to be present in the header.
421        if !indices_given {
422            if let Some(names) = column_names {
423                // The tokenizer has already unquoted both sides.
424                let name_indices: Vec<usize> = header
425                    .iter()
426                    .enumerate()
427                    .filter_map(|(i, name)| {
428                        if names.iter().any(|n| n == name) {
429                            Some(i + header_offset)
430                        } else {
431                            None
432                        }
433                    })
434                    .collect();
435                relevant_indices.extend(name_indices);
436            }
437        }
438
439        // Neither selector given: take EVERY column except the row-name one.
440        // Without this the selection stays empty and the reader silently returns
441        // a 0-column matrix, so `read_data(.., None, None)` — the form used by
442        // `senna`'s `read_mat` and `data-beans-sim`'s topic-file loader — could
443        // never read a delimited file at all.
444        if column_indices.is_none() && column_names.is_none() {
445            let n_col = header
446                .len()
447                .max(lines.first().map_or(0, |words| words.len()));
448            relevant_indices.extend((0..n_col).filter(|j| Some(*j) != row_name_index));
449        }
450
451        relevant_indices.sort_unstable();
452        relevant_indices.dedup();
453
454        // Every subscript below is checked here first: the row-name column, the
455        // selected data columns, and the header lookup that names them. Width is
456        // the NARROWEST data row, ignoring blank ones, because a ragged file
457        // otherwise passes this and panics later in the value loop. The header
458        // is checked separately, since it can be one field short of the data
459        // rows when a writer omits a name for the row-label column.
460        let n_columns = data_width;
461        let mut to_check: Vec<usize> = relevant_indices.clone();
462        to_check.extend(row_name_index);
463        if let Some(&bad) = to_check.iter().find(|&&j| j >= n_columns) {
464            return Err(anyhow::anyhow!(
465                "column index {bad} is out of range: the file has {n_columns} column(s). \
466                 Name the columns to read, or pass indices within range."
467            ));
468        }
469
470        let row_names: Vec<Box<str>> = match row_name_index {
471            // Unquoted, for the same reason the header is: a fully-quoted csv
472            // would otherwise yield row names carrying their quotes, which then
473            // match nothing when joined against a matrix's own names.
474            Some(row_name_index) => lines
475                .iter()
476                .map(|words| words[row_name_index].clone())
477                .collect(),
478            _ => (0..lines.len())
479                .map(|x| x.to_string().into_boxed_str())
480                .collect(),
481        };
482
483        // Indices can come from a caller's fallback rather than from a name
484        // match, so they are not guaranteed to exist in this file. Say which
485        // column was asked for and how many the file has, instead of panicking
486        // on the subscript several lines later.
487        let column_names: Vec<Box<str>> = if header.is_empty() {
488            relevant_indices
489                .iter()
490                .map(|x| x.to_string().into_boxed_str())
491                .collect()
492        } else {
493            relevant_indices
494                .iter()
495                // A header can be narrower than the data rows; fall back to the
496                // position rather than panicking on a name that was never written.
497                .map(|&j| {
498                    j.checked_sub(header_offset)
499                        .and_then(|k| header.get(k))
500                        .cloned()
501                        .unwrap_or_else(|| j.to_string().into_boxed_str())
502                })
503                .collect()
504        };
505
506        let data: Vec<Vec<Self::Scalar>> = lines
507            .iter()
508            .map(|words| {
509                relevant_indices
510                    .iter()
511                    .map(|&i| words[i].parse::<Self::Scalar>().expect("failed to parse"))
512                    .collect()
513            })
514            .collect();
515
516        let data = data.into_iter().flatten().collect::<Vec<_>>();
517
518        Ok((row_names, column_names, data))
519    }
520
521    /// Read a `tsv` file while skipping until the header row
522    fn from_tsv(tsv_file: &str, skip: Option<usize>) -> anyhow::Result<Self::Mat> {
523        Self::read_file_delim(tsv_file, "\t", skip)
524    }
525
526    /// Read a `csv` file while skipping until the header row
527    fn from_csv(csv_file: &str, skip: Option<usize>) -> anyhow::Result<Self::Mat> {
528        Self::read_file_delim(csv_file, ",", skip)
529    }
530
531    /// write the matrix down to a file with delimiter
532    /// * `file_path`: output file path
533    /// * `delim`: separation character or string
534    fn write_file_delim(&self, file: &str, delim: &str) -> anyhow::Result<()>;
535
536    /// write the matrix down to a tsv file
537    /// * `file_path`: output file path
538    fn to_tsv(&self, tsv_file: &str) -> anyhow::Result<()> {
539        self.write_file_delim(tsv_file, "\t")
540    }
541
542    /// write the matrix down to a csv file
543    /// * `file_path`: output file path
544    fn to_csv(&self, csv_file: &str) -> anyhow::Result<()> {
545        self.write_file_delim(csv_file, ",")
546    }
547
548    /// write the matrix down to parquet with full control over naming
549    /// * `file_path`: output file path
550    /// * `row_names`: Tuple of (optional row_names, optional row_column_name)
551    ///   - `(None, None)`: use numeric row names `[0, n)` with "row" column name
552    ///   - `(None, Some("cell_pair"))`: use numeric row names with "cell_pair" column name
553    ///   - `(Some(names), None)`: use provided names with "row" column name
554    ///   - `(Some(names), Some("gene"))`: use provided names with "gene" column name
555    /// * `column_names`: if `None`, just add `[0, n)` numbers.
556    fn to_parquet_with_names(
557        &self,
558        file_path: &str,
559        row_names: (Option<&[Box<str>]>, Option<&str>),
560        column_names: Option<&[Box<str>]>,
561    ) -> anyhow::Result<()>;
562
563    /// write the matrix down to parquet with default names
564    /// * `file_path`: output file path
565    ///   Uses numeric row/column names and default "row" column name
566    fn to_parquet(&self, file_path: &str) -> anyhow::Result<()> {
567        self.to_parquet_with_names(file_path, (None, None), None)
568    }
569
570    /// Read a real-valued numeric matrix with the default row
571    /// index(0) and all the other available columns.
572    /// Assumes column 0 contains row names.
573    ///
574    fn from_parquet(file_path: &str) -> anyhow::Result<MatWithNames<Self::Mat>> {
575        Self::from_parquet_with_indices(file_path, Some(0), None)
576    }
577
578    /// Read a real-valued numeric matrix treating all columns as data.
579    /// Row names will be generated as "0", "1", "2", ...
580    ///
581    fn from_parquet_no_row_names(file_path: &str) -> anyhow::Result<MatWithNames<Self::Mat>> {
582        Self::from_parquet_with_indices(file_path, None, None)
583    }
584
585    /// Read a real-valued numeric matrix from the parquet file. We
586    /// can specify row name index. We can specify the row name column
587    /// index and desired column indices.
588    /// * `row_name_index`: column index (0-based) corresponds to row name
589    fn from_parquet_with_row_names(
590        file_path: &str,
591        row_name_index: Option<usize>,
592    ) -> anyhow::Result<MatWithNames<Self::Mat>> {
593        Self::from_parquet_with_indices_names(file_path, row_name_index, None, None)
594    }
595
596    /// Read a real-valued numeric matrix from the parquet file. We
597    /// can specify row name index. We can specify the row name column
598    /// index and desired column indices.
599    /// * `row_name_index`: column index (0-based) corresponds to row name
600    /// * `column_indices`: column indices (0-based) to include
601    fn from_parquet_with_indices(
602        file_path: &str,
603        row_name_index: Option<usize>,
604        column_indices: Option<&[usize]>,
605    ) -> anyhow::Result<MatWithNames<Self::Mat>> {
606        Self::from_parquet_with_indices_names(file_path, row_name_index, column_indices, None)
607    }
608
609    /// Read a real-valued numeric matrix from the parquet file.  We
610    /// can specify row name index.  We can specify the row name
611    /// column index and desired column names.
612    /// * `row_name_index`: column index (0-based) corresponds to row name
613    /// * `column_names`: column names to include
614    fn from_parquet_with_names(
615        file_path: &str,
616        row_name_index: Option<usize>,
617        column_names: Option<&[Box<str>]>,
618    ) -> anyhow::Result<MatWithNames<Self::Mat>> {
619        Self::from_parquet_with_indices_names(file_path, row_name_index, None, column_names)
620    }
621
622    /// Read a real-valued numeric matrix from the parquet file.  We
623    /// can specify row name index.  We can specify the row name
624    /// column index and desired column indices and names.
625    /// * `row_name_index`: column index (0-based) corresponds to row name
626    /// * `column_indices`: column indices (0-based) to include
627    /// * `column_names`: column names to include
628    fn from_parquet_with_indices_names(
629        file_path: &str,
630        row_name_index: Option<usize>,
631        column_indices: Option<&[usize]>,
632        column_names: Option<&[Box<str>]>,
633    ) -> anyhow::Result<MatWithNames<Self::Mat>>;
634}
635
636/// intput data matrix `mat` with `rows` and `cols`
637pub struct MatWithNames<M> {
638    pub rows: Vec<Box<str>>,
639    pub cols: Vec<Box<str>>,
640    pub mat: M,
641}
642
643pub trait MeltOps {
644    type Scalar;
645    type Mat;
646    /// melt a matrix with indices
647    fn melt_with_indexes(&self) -> (Vec<Self::Scalar>, Vec<Vec<usize>>);
648    /// melt a matrix
649    fn melt(&self) -> Vec<Self::Scalar>;
650    /// Melt multiple matrices/tensors together in a single traversal for cache efficiency.
651    /// All inputs must have the same dimensions.
652    /// Returns (values for each input, indices for each dimension).
653    fn melt_many_with_indexes(&self, others: &[&Self])
654        -> (Vec<Vec<Self::Scalar>>, Vec<Vec<usize>>);
655}
656
657#[cfg(feature = "tensor")]
658pub trait CandleDataLoaderOps {
659    type Scalar;
660    type Mat;
661    // /// unify transpose
662    // fn transpose(&self) -> Self::Mat;
663    /// take each row vector as a sample
664    fn rows_to_tensor_vec(&self) -> Vec<Tensor>;
665
666    /// Return (nrows, ncols) dimensions
667    fn data_shape(&self) -> (usize, usize);
668
669    /// Extract row i as Vec<f32>.
670    ///
671    /// WARNING: default creates ALL row tensors then picks one — O(N*D) for O(D) work.
672    /// Implementors should override this.
673    fn row_to_f32_vec(&self, i: usize) -> Vec<f32> {
674        let t = &self.rows_to_tensor_vec()[i];
675        t.flatten_all().unwrap().to_vec1::<f32>().unwrap()
676    }
677}