pub fn one_hot_encode(
categorical: &Array2<String>,
column_names: Option<&[&str]>,
) -> Result<(Array2<f64>, Vec<String>), DatasetError>Expand description
One-hot encode a matrix of categorical string features.
The mixed-type loaders (adult, titanic, bank_marketing, abalone,
kddcup99, palmer_penguins) and the all-categorical ones (mushroom,
car_evaluation) keep their categorical values in
ColumnData::String columns. Read each
column by name with as_string, then stack the columns into an
Array2<String>. No numeric model can consume strings directly. This function
expands each column into one indicator column per level it takes. A row gets
1.0 in the column for its own level, and 0.0 everywhere else.
This function sorts levels within a column, so the output layout depends only
on the values present. The returned names identify the columns as
<column>=<level>, using column_names when supplied, and column_0,
column_1, and so on otherwise.
This widens the matrix by however many distinct levels the data holds. That is
harmless for mushroom (22 columns become 117). Before running it on
kddcup99’s service column (70 levels over millions of rows), check the
resulting width.
§Parameters
categorical- The categorical matrix, shape(n_samples, n_features).column_names- Optional names for the source columns, used to build the output names. Must have one entry per column when supplied.
§Returns
(Array2<f64>, Vec<String>)- The indicator matrix, shape(n_samples, total_levels), and one name per output column.
§Errors
DatasetError::ValidationError- Returns this whencategoricalhas no rows or no columns.DatasetError::LengthMismatch- Returns this whencolumn_namesisSomebut its length does not match the column count.
§Example
use dataset_ml::preprocessing::one_hot_encode;
use ndarray::array;
let categorical = array![
["male".to_string(), "S".to_string()],
["female".to_string(), "C".to_string()],
["male".to_string(), "C".to_string()],
];
let (encoded, names) = one_hot_encode(&categorical, Some(&["sex", "port"])).unwrap();
assert_eq!(names, vec!["sex=female", "sex=male", "port=C", "port=S"]);
assert_eq!(encoded.row(0).to_vec(), vec![0.0, 1.0, 0.0, 1.0]); // male, S