From b305b9f7aca6e748256800b5a61d61f4be734c39 Mon Sep 17 00:00:00 2001 From: Luis Moreno Date: Fri, 4 Nov 2022 17:50:52 -0500 Subject: [PATCH 1/4] Handle kernel serialization --- Cargo.toml | 3 ++- src/svm/mod.rs | 43 +++++++------------------------------------ src/svm/svc.rs | 1 - src/svm/svr.rs | 1 - 4 files changed, 9 insertions(+), 39 deletions(-) diff --git a/Cargo.toml b/Cargo.toml index 0a230832..9c79af40 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -21,10 +21,11 @@ num = "0.4" rand = { version = "0.8.5", default-features = false, features = ["small_rng"] } rand_distr = { version = "0.4", optional = true } serde = { version = "1", features = ["derive"], optional = true } +typetag = { version = "0.2", optional = true } [features] default = ["serde", "datasets"] -serde = ["dep:serde"] +serde = ["dep:serde", "dep:typetag"] ndarray-bindings = ["dep:ndarray"] datasets = ["dep:rand_distr", "std"] std = ["rand/std_rng", "rand/std"] diff --git a/src/svm/mod.rs b/src/svm/mod.rs index a30fe876..55b9218d 100644 --- a/src/svm/mod.rs +++ b/src/svm/mod.rs @@ -30,8 +30,6 @@ pub mod svr; use core::fmt::Debug; -#[cfg(feature = "serde")] -use serde::ser::{SerializeStruct, Serializer}; #[cfg(feature = "serde")] use serde::{Deserialize, Serialize}; @@ -40,36 +38,17 @@ use crate::linalg::basic::arrays::{Array1, ArrayView1}; /// Defines a kernel function. /// This is a object-safe trait. -pub trait Kernel { +#[typetag::serde(tag = "type")] +pub trait Kernel: Debug { #[allow(clippy::ptr_arg)] /// Apply kernel function to x_i and x_j fn apply(&self, x_i: &Vec, x_j: &Vec) -> Result; - /// Return a serializable name - fn name(&self) -> &'static str; -} - -impl Debug for dyn Kernel { - fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result { - write!(f, "Kernel") - } -} - -#[cfg(feature = "serde")] -impl Serialize for dyn Kernel { - fn serialize(&self, serializer: S) -> Result - where - S: Serializer, - { - let mut s = serializer.serialize_struct("Kernel", 1)?; - s.serialize_field("type", &self.name())?; - s.end() - } } /// Pre-defined kernel functions #[cfg_attr(feature = "serde", derive(Serialize, Deserialize))] #[derive(Debug, Clone)] -pub struct Kernels {} +pub struct Kernels; impl Kernels { /// Return a default linear @@ -211,15 +190,14 @@ impl SigmoidKernel { } } +#[cfg_attr(feature = "serde", typetag::serde)] impl Kernel for LinearKernel { fn apply(&self, x_i: &Vec, x_j: &Vec) -> Result { Ok(x_i.dot(x_j)) } - fn name(&self) -> &'static str { - "Linear" - } } +#[cfg_attr(feature = "serde", typetag::serde)] impl Kernel for RBFKernel { fn apply(&self, x_i: &Vec, x_j: &Vec) -> Result { if self.gamma.is_none() { @@ -231,11 +209,9 @@ impl Kernel for RBFKernel { let v_diff = x_i.sub(x_j); Ok((-self.gamma.unwrap() * v_diff.mul(&v_diff).sum()).exp()) } - fn name(&self) -> &'static str { - "RBF" - } } +#[cfg_attr(feature = "serde", typetag::serde)] impl Kernel for PolynomialKernel { fn apply(&self, x_i: &Vec, x_j: &Vec) -> Result { if self.gamma.is_none() || self.coef0.is_none() || self.degree.is_none() { @@ -247,11 +223,9 @@ impl Kernel for PolynomialKernel { let dot = x_i.dot(x_j); Ok((self.gamma.unwrap() * dot + self.coef0.unwrap()).powf(self.degree.unwrap())) } - fn name(&self) -> &'static str { - "Polynomial" - } } +#[cfg_attr(feature = "serde", typetag::serde)] impl Kernel for SigmoidKernel { fn apply(&self, x_i: &Vec, x_j: &Vec) -> Result { if self.gamma.is_none() || self.coef0.is_none() { @@ -263,9 +237,6 @@ impl Kernel for SigmoidKernel { let dot = x_i.dot(x_j); Ok(self.gamma.unwrap() * dot + self.coef0.unwrap().tanh()) } - fn name(&self) -> &'static str { - "Sigmoid" - } } #[cfg(test)] diff --git a/src/svm/svc.rs b/src/svm/svc.rs index 9cb140d7..0f277136 100644 --- a/src/svm/svc.rs +++ b/src/svm/svc.rs @@ -100,7 +100,6 @@ pub struct SVCParameters>, /// Unused parameter. diff --git a/src/svm/svr.rs b/src/svm/svr.rs index 7a39a56b..c5d71620 100644 --- a/src/svm/svr.rs +++ b/src/svm/svr.rs @@ -92,7 +92,6 @@ pub struct SVRParameters { pub c: T, /// Tolerance for stopping criterion. pub tol: T, - #[cfg_attr(feature = "serde", serde(skip_deserializing))] /// The kernel function. pub kernel: Option>, } From 07c95908e13772b14cac57aa80e5483cc5b1140d Mon Sep 17 00:00:00 2001 From: Luis Moreno Date: Fri, 4 Nov 2022 18:06:48 -0500 Subject: [PATCH 2/4] Do not use typetag in WASM --- Cargo.toml | 2 ++ src/svm/mod.rs | 13 ++++++++----- src/svm/svc.rs | 4 ++++ src/svm/svr.rs | 4 ++++ 4 files changed, 18 insertions(+), 5 deletions(-) diff --git a/Cargo.toml b/Cargo.toml index 9c79af40..ad7f5476 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -21,6 +21,8 @@ num = "0.4" rand = { version = "0.8.5", default-features = false, features = ["small_rng"] } rand_distr = { version = "0.4", optional = true } serde = { version = "1", features = ["derive"], optional = true } + +[target.'cfg(not(target_arch = "wasm32"))'.dependencies] typetag = { version = "0.2", optional = true } [features] diff --git a/src/svm/mod.rs b/src/svm/mod.rs index 55b9218d..ce6602f9 100644 --- a/src/svm/mod.rs +++ b/src/svm/mod.rs @@ -38,7 +38,10 @@ use crate::linalg::basic::arrays::{Array1, ArrayView1}; /// Defines a kernel function. /// This is a object-safe trait. -#[typetag::serde(tag = "type")] +#[cfg_attr( + all(feature = "serde", not(target_arch = "wasm32")), + typetag::serde(tag = "type") +)] pub trait Kernel: Debug { #[allow(clippy::ptr_arg)] /// Apply kernel function to x_i and x_j @@ -190,14 +193,14 @@ impl SigmoidKernel { } } -#[cfg_attr(feature = "serde", typetag::serde)] +#[cfg_attr(all(feature = "serde", not(target_arch = "wasm32")), typetag::serde)] impl Kernel for LinearKernel { fn apply(&self, x_i: &Vec, x_j: &Vec) -> Result { Ok(x_i.dot(x_j)) } } -#[cfg_attr(feature = "serde", typetag::serde)] +#[cfg_attr(all(feature = "serde", not(target_arch = "wasm32")), typetag::serde)] impl Kernel for RBFKernel { fn apply(&self, x_i: &Vec, x_j: &Vec) -> Result { if self.gamma.is_none() { @@ -211,7 +214,7 @@ impl Kernel for RBFKernel { } } -#[cfg_attr(feature = "serde", typetag::serde)] +#[cfg_attr(all(feature = "serde", not(target_arch = "wasm32")), typetag::serde)] impl Kernel for PolynomialKernel { fn apply(&self, x_i: &Vec, x_j: &Vec) -> Result { if self.gamma.is_none() || self.coef0.is_none() || self.degree.is_none() { @@ -225,7 +228,7 @@ impl Kernel for PolynomialKernel { } } -#[cfg_attr(feature = "serde", typetag::serde)] +#[cfg_attr(all(feature = "serde", not(target_arch = "wasm32")), typetag::serde)] impl Kernel for SigmoidKernel { fn apply(&self, x_i: &Vec, x_j: &Vec) -> Result { if self.gamma.is_none() || self.coef0.is_none() { diff --git a/src/svm/svc.rs b/src/svm/svc.rs index 0f277136..41ab5db4 100644 --- a/src/svm/svc.rs +++ b/src/svm/svc.rs @@ -101,6 +101,10 @@ pub struct SVCParameters>, /// Unused parameter. m: PhantomData<(X, Y, TY)>, diff --git a/src/svm/svr.rs b/src/svm/svr.rs index c5d71620..c1336ca8 100644 --- a/src/svm/svr.rs +++ b/src/svm/svr.rs @@ -93,6 +93,10 @@ pub struct SVRParameters { /// Tolerance for stopping criterion. pub tol: T, /// The kernel function. + #[cfg_attr( + all(feature = "serde", target_arch = "wasm32"), + serde(skip_serializing, skip_deserializing) + )] pub kernel: Option>, } From 6021f648509efa63cf6983f513912119f5e1d642 Mon Sep 17 00:00:00 2001 From: Luis Moreno Date: Sat, 5 Nov 2022 10:14:20 -0500 Subject: [PATCH 3/4] enable tests for serialization --- src/svm/svc.rs | 7 ++++--- src/svm/svr.rs | 12 ++++-------- 2 files changed, 8 insertions(+), 11 deletions(-) diff --git a/src/svm/svc.rs b/src/svm/svc.rs index 41ab5db4..99086249 100644 --- a/src/svm/svc.rs +++ b/src/svm/svc.rs @@ -1089,7 +1089,7 @@ mod tests { wasm_bindgen_test::wasm_bindgen_test )] #[test] - #[cfg(feature = "serde")] + #[cfg(all(feature = "serde", not(target_arch = "wasm32")))] fn svc_serde() { let x = DenseMatrix::from_2d_array(&[ &[5.1, 3.5, 1.4, 0.2], @@ -1123,8 +1123,9 @@ mod tests { let svc = SVC::fit(&x, &y, ¶ms).unwrap(); // serialization - let serialized_svc = &serde_json::to_string(&svc).unwrap(); + let deserialized_svc: SVC = + serde_json::from_str(&serde_json::to_string(&svc).unwrap()).unwrap(); - println!("{:?}", serialized_svc); + assert_eq!(svc, deserialized_svc); } } diff --git a/src/svm/svr.rs b/src/svm/svr.rs index c1336ca8..90f94908 100644 --- a/src/svm/svr.rs +++ b/src/svm/svr.rs @@ -673,7 +673,7 @@ mod tests { wasm_bindgen_test::wasm_bindgen_test )] #[test] - #[cfg(feature = "serde")] + #[cfg(all(feature = "serde", not(target_arch = "wasm32")))] fn svr_serde() { let x = DenseMatrix::from_2d_array(&[ &[234.289, 235.6, 159.0, 107.608, 1947., 60.323], @@ -704,13 +704,9 @@ mod tests { let svr = SVR::fit(&x, &y, ¶ms).unwrap(); - let serialized = &serde_json::to_string(&svr).unwrap(); - - println!("{}", &serialized); - - // let deserialized_svr: SVR, LinearKernel> = - // serde_json::from_str(&serde_json::to_string(&svr).unwrap()).unwrap(); + let deserialized_svr: SVR, _> = + serde_json::from_str(&serde_json::to_string(&svr).unwrap()).unwrap(); - // assert_eq!(svr, deserialized_svr); + assert_eq!(svr, deserialized_svr); } } From fe27ce91e8a860d9cd8d2a9811818440ad0b8623 Mon Sep 17 00:00:00 2001 From: Luis Moreno Date: Tue, 8 Nov 2022 11:07:23 -0500 Subject: [PATCH 4/4] Update serde feature deps --- Cargo.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/Cargo.toml b/Cargo.toml index d3353b4a..37280fd0 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -35,7 +35,7 @@ typetag = { version = "0.2", optional = true } [features] default = [] -serde = ["dep:serde"] +serde = ["dep:serde", "dep:typetag"] ndarray-bindings = ["dep:ndarray"] datasets = ["dep:rand_distr", "std_rand", "serde"] std_rand = ["rand/std_rng", "rand/std"]