diff --git a/src/python/graph_builder.rs b/src/python/graph_builder.rs index 23a1d81..4086056 100644 --- a/src/python/graph_builder.rs +++ b/src/python/graph_builder.rs @@ -499,38 +499,32 @@ impl PyMLGraphBuilder { let pads = pads.unwrap_or_else(|| vec![0, 0, 0, 0]); let groups = groups.unwrap_or(1); - let input_layout_s = match input_layout.unwrap_or("nchw") { - "nchw" => "nchw", - "nhwc" => "nhwc", - other => { - return Err(pyo3::exceptions::PyValueError::new_err(format!( - "Invalid input_layout '{}', must be 'nchw' or 'nhwc'", - other - ))); - } - }; - - let filter_layout_s = match filter_layout.unwrap_or("oihw") { - "oihw" => "oihw", - "hwio" => "hwio", - "ohwi" => "ohwi", - "ihwo" => "ihwo", - other => { - return Err(pyo3::exceptions::PyValueError::new_err(format!( - "Invalid filter_layout '{}', must be 'oihw', 'hwio', 'ohwi', or 'ihwo'", - other - ))); - } - }; - let conv2d_options = MLConv2dOptions { label: String::new(), padding: pads, strides, dilations, groups, - input_layout: input_layout_s.to_string(), - filter_layout: filter_layout_s.to_string(), + input_layout: input_layout + .map(|layout| { + serde_json::from_str(&format!("\"{layout}\"")).map_err(|_err| { + pyo3::exceptions::PyValueError::new_err(format!( + "Invalid input_layout '{layout}', must be 'nchw' or 'nhwc'", + )) + }) + }) + .transpose()? + .unwrap_or_default(), + filter_layout: filter_layout + .map(|layout| { + serde_json::from_str(&format!("\"{layout}\"")).map_err(|_err| { + pyo3::exceptions::PyValueError::new_err(format!( + "Invalid filter_layout '{layout}', must be 'oihw', 'hwio', 'ohwi', or 'ihwo'", + )) + }) + }) + .transpose()? + .unwrap_or_default(), bias: bias.map(|b| b.id), }; @@ -596,30 +590,6 @@ impl PyMLGraphBuilder { let output_padding = output_padding.unwrap_or_else(|| vec![0, 0]); let groups = groups.unwrap_or(1); - let input_layout_s = match input_layout.unwrap_or("nchw") { - "nchw" => "nchw", - "nhwc" => "nhwc", - other => { - return Err(pyo3::exceptions::PyValueError::new_err(format!( - "Invalid input_layout '{}', must be 'nchw' or 'nhwc'", - other - ))); - } - }; - - let filter_layout_s = match filter_layout.unwrap_or("iohw") { - "iohw" => "iohw", - "hwoi" => "hwoi", - "ohwi" => "ohwi", - "oihw" => "oihw", - other => { - return Err(pyo3::exceptions::PyValueError::new_err(format!( - "Invalid filter_layout '{}', must be 'iohw', 'hwoi', 'ohwi', or 'oihw'", - other - ))); - } - }; - let conv_t_options = MLConvTranspose2dOptions { label: String::new(), padding: pads, @@ -628,8 +598,26 @@ impl PyMLGraphBuilder { output_padding, output_sizes: output_sizes.clone(), groups, - input_layout: input_layout_s.to_string(), - filter_layout: filter_layout_s.to_string(), + input_layout: input_layout + .map(|layout| { + serde_json::from_str(&format!("\"{layout}\"")).map_err(|_err| { + pyo3::exceptions::PyValueError::new_err(format!( + "Invalid input_layout '{layout}', must be 'nchw' or 'nhwc'", + )) + }) + }) + .transpose()? + .unwrap_or_default(), + filter_layout: filter_layout + .map(|layout| { + serde_json::from_str(&format!("\"{layout}\"")).map_err(|_err| { + pyo3::exceptions::PyValueError::new_err(format!( + "Invalid filter_layout '{layout}', must be 'iohw' | 'hwoi' | 'ohwi'", + )) + }) + }) + .transpose()? + .unwrap_or_default(), bias: bias.map(|b| b.id), };