1#[allow(unused_imports)]
32use pyo3::prelude::*;
33
34mod clustering;
36mod datasets;
37mod ensemble;
38mod linear;
39mod metrics;
40mod model_selection;
41mod naive_bayes;
42mod neural_network;
43#[allow(dead_code)]
48mod preprocessing;
49mod tree;
50mod utils;
51
52pub use clustering::*;
54pub use ensemble::*;
55pub use linear::*;
56pub use metrics::*;
57pub use model_selection::*;
58pub use naive_bayes::*;
59pub use neural_network::*;
60pub use preprocessing::*;
61pub use tree::*;
62pub use utils::*;
63
64#[pymodule]
66fn _sklears(m: &Bound<'_, PyModule>) -> PyResult<()> {
67 m.add("__version__", env!("CARGO_PKG_VERSION"))?;
69 m.add(
70 "__doc__",
71 "High-performance machine learning library with scikit-learn compatibility",
72 )?;
73
74 m.add_class::<linear::PyLinearRegression>()?;
76 m.add_class::<linear::PyRidge>()?;
77 m.add_class::<linear::PyLasso>()?;
78 m.add_class::<linear::PyElasticNet>()?;
79 m.add_class::<linear::PyBayesianRidge>()?;
80 m.add_class::<linear::PyARDRegression>()?;
81 m.add_class::<linear::PyLogisticRegression>()?;
82
83 m.add_class::<ensemble::PyGradientBoostingClassifier>()?;
85 m.add_class::<ensemble::PyGradientBoostingRegressor>()?;
86 m.add_class::<ensemble::PyAdaBoostClassifier>()?;
87 m.add_class::<ensemble::PyVotingClassifier>()?;
88 m.add_class::<ensemble::PyBaggingClassifier>()?;
89
90 m.add_class::<neural_network::PyMLPClassifier>()?;
92 m.add_class::<neural_network::PyMLPRegressor>()?;
93
94 m.add_class::<naive_bayes::PyGaussianNB>()?;
102 m.add_class::<naive_bayes::PyMultinomialNB>()?;
103 m.add_class::<naive_bayes::PyBernoulliNB>()?;
104 m.add_class::<naive_bayes::PyComplementNB>()?;
105
106 m.add_class::<clustering::PyKMeans>()?;
108 m.add_class::<clustering::PyDBSCAN>()?;
109
110 m.add_class::<preprocessing::PyStandardScaler>()?;
112 m.add_class::<preprocessing::PyMinMaxScaler>()?;
113 m.add_class::<preprocessing::PyLabelEncoder>()?;
114
115 m.add_function(wrap_pyfunction!(metrics::mean_squared_error, m)?)?;
117 m.add_function(wrap_pyfunction!(metrics::mean_absolute_error, m)?)?;
118 m.add_function(wrap_pyfunction!(metrics::r2_score, m)?)?;
119 m.add_function(wrap_pyfunction!(metrics::mean_squared_log_error, m)?)?;
120 m.add_function(wrap_pyfunction!(metrics::median_absolute_error, m)?)?;
121
122 m.add_function(wrap_pyfunction!(metrics::accuracy_score, m)?)?;
124 m.add_function(wrap_pyfunction!(metrics::precision_score, m)?)?;
125 m.add_function(wrap_pyfunction!(metrics::recall_score, m)?)?;
126 m.add_function(wrap_pyfunction!(metrics::f1_score, m)?)?;
127 m.add_function(wrap_pyfunction!(metrics::confusion_matrix, m)?)?;
128 m.add_function(wrap_pyfunction!(metrics::classification_report, m)?)?;
129
130 m.add_function(wrap_pyfunction!(model_selection::train_test_split, m)?)?;
132 m.add_class::<model_selection::PyKFold>()?;
133
134 datasets::register_dataset_functions(m)?;
136
137 m.add_function(wrap_pyfunction!(utils::get_version, m)?)?;
139 m.add_function(wrap_pyfunction!(utils::get_build_info, m)?)?;
140 m.add_function(wrap_pyfunction!(utils::get_hardware_info, m)?)?;
141 m.add_function(wrap_pyfunction!(utils::benchmark_basic_operations, m)?)?;
142
143 Ok(())
144}