diff --git a/README.md b/README.md index 1f1525b..29a92dd 100644 --- a/README.md +++ b/README.md @@ -74,7 +74,7 @@ See [examples/README.md](examples/README.md) for the raw-weight workflow. ## ONNX conversion -The converter accepts `ai.onnx` opsets 11 through 18. Static dimension overrides and optional constant +The converter accepts `ai.onnx` opsets 11 through 20. Static dimension overrides and optional constant folding can resolve shape-critical ONNX expressions. Experimental bounded dynamic input metadata is available, but operations whose arguments must be static still require concrete values. diff --git a/docs/onnx-lowering.md b/docs/onnx-lowering.md index 4af15cc..801fd68 100644 --- a/docs/onnx-lowering.md +++ b/docs/onnx-lowering.md @@ -1,7 +1,7 @@ # ONNX to WebNN lowering The optional `onnx` feature converts ONNX models into the `GraphJson` AST and serializes them as `.webnn` or -JSON. The converter accepts `ai.onnx` opsets 11 through 18. Other domains are retained for operator-specific +JSON. The converter accepts `ai.onnx` opsets 11 through 20. Other domains are retained for operator-specific handling rather than being checked by the `ai.onnx` opset guard. ## Conversion flow @@ -33,6 +33,24 @@ operation that requires a static reshape target, axis, permutation, slice bound, The exact supported behavior is operator- and opset-specific. Source tests are authoritative; this page does not claim that every variant of a named ONNX operator is supported. +Standard-domain `Gelu` is supported from opset 20. Exact GELU (an absent `approximate` attribute or +`"none"`) maps to WebNN `gelu`; `approximate="tanh"` lowers to the ONNX polynomial with +`mul`/`add`/`tanh`, not to exact GELU. FP16 inputs are promoted to FP32 for the intermediate +polynomial and rounded back to FP16 at the output. This avoids half-precision intermediate overflow +and excessive negative-tail cancellation. Scalar inputs remain rank zero. Invalid approximation +attributes fail conversion; the older `com.microsoft::Gelu` accepts no approximation attribute. + +The opset-19/20 audit preserves `AveragePool` dilation, and reductions resolve constant axes inputs, +`keepdims=0`, and `noop_with_empty_axes`. Dynamic reduction axes still fail conversion. Added float8, +bfloat16, string, sequence, and optional types remain unsupported rather than being reinterpreted as +FP32. `Cast`'s float8-only `saturate` option has no effect on supported destination types. New operators +without handlers (including `CastLike`, `Resize`, `QuantizeLinear`, `DequantizeLinear`, `GridSample`, +`AffineGrid`, and `DFT`) remain explicit errors; accepting an opset does not claim every operator in it. + +`tests/onnx_gelu.rs` covers optimized/unoptimized full import, reference numerical results, typed +constants, scalar inputs, generated-name collisions, and emitted JavaScript signatures. WebNN WPT +does not import ONNX, so passing exact-GELU WPT is independent of support for ONNX's tanh variant. + ## Constants and output artifacts By default, `convert-onnx` extracts initializers and large inline constants into a headerless `.weights` blob and @@ -75,7 +93,7 @@ webnn-graph --debug convert-onnx --input model.onnx --optimize When conversion fails, first check: -- whether the model uses `ai.onnx` opset 11–18; +- whether the model uses `ai.onnx` opset 11–20; - whether every required symbolic dimension has an override or a usable bounded representation; - whether `--optimize` can fold the shape-producing expression; - whether the specific operator form and attributes have a registered lowering. diff --git a/src/emit_js.rs b/src/emit_js.rs index 758db6e..6ffa6ca 100644 --- a/src/emit_js.rs +++ b/src/emit_js.rs @@ -250,8 +250,18 @@ pub fn emit_builder_js(g: &GraphJson) -> String { .join(", "); let mut opts_val = serde_json::Value::Object(n.options.clone()); normalize_options_for_js(&mut opts_val); + let cast_type = if n.op == "cast" { + opts_val + .as_object_mut() + .and_then(|options| options.remove("to")) + } else { + None + }; let opts = opts_val.to_string(); - let call = if ins.is_empty() { + let call = if let Some(dtype) = cast_type { + // WebNN takes the destination dtype positionally, not in options. + format!("builder[\"cast\"]({ins}, {dtype}, {opts})") + } else if ins.is_empty() { format!("builder[{op:?}]({opts})", op = n.op, opts = opts) } else { format!( @@ -477,7 +487,7 @@ mod tests { }); g.outputs.insert("y".to_string(), "y".to_string()); let js = emit_builder_js(&g); - assert!(js.contains("\"to\":\"int32\"")); + assert!(js.contains("builder[\"cast\"](env.get(\"x\"), \"int32\", {})")); } #[test] diff --git a/src/onnx/convert.rs b/src/onnx/convert.rs index add1777..4c9e64e 100644 --- a/src/onnx/convert.rs +++ b/src/onnx/convert.rs @@ -14,7 +14,7 @@ use thiserror::Error; use webnn_onnx_utils::{data_types as utils_data_types, identifiers}; const MIN_SUPPORTED_OPSET: i64 = 11; -const MAX_SUPPORTED_OPSET: i64 = 18; +const MAX_SUPPORTED_OPSET: i64 = 20; #[derive(Debug, Error)] pub enum OnnxError { @@ -33,6 +33,14 @@ pub enum OnnxError { #[error("missing required attribute: {attr} in {op}")] MissingAttribute { attr: String, op: String }, + #[error("invalid attribute '{attr}' in {op} (node: {node}): {reason}")] + InvalidAttribute { + attr: String, + op: String, + node: String, + reason: String, + }, + #[error("invalid tensor shape: {0}")] InvalidShape(String), @@ -74,6 +82,61 @@ pub(crate) fn map_onnx_data_type(onnx_type: i32) -> Result }) } +/// Check opset-19/20 type extensions before either import path folds constants. +/// A foldable intermediate is not permission to reinterpret an unsupported type. +fn validate_extended_opset_types(model: &ModelProto) -> Result<(), OnnxError> { + let graph = model + .graph + .as_ref() + .ok_or_else(|| OnnxError::ProtobufError("Missing graph in model".to_string()))?; + if !model.opset_import.iter().any(|import| { + (import.domain.is_empty() || import.domain == "ai.onnx") && import.version >= 19 + }) { + return Ok(()); + } + for value in graph + .input + .iter() + .chain(&graph.output) + .chain(&graph.value_info) + { + if let Some(type_proto) = &value.r#type { + match &type_proto.value { + Some(TypeProtoValue::TensorType(tensor)) => { + map_onnx_data_type(tensor.elem_type)?; + } + _ => { + return Err(OnnxError::UnsupportedOp { + op: "non-tensor value type".to_string(), + node: value.name.clone(), + }) + } + } + } + } + for tensor in &graph.initializer { + map_onnx_data_type(tensor.data_type)?; + } + for node in &graph.node { + for attribute in &node.attribute { + if let Some(tensor) = &attribute.t { + map_onnx_data_type(tensor.data_type)?; + } + if node.op_type == "Cast" && attribute.name == "to" { + map_onnx_data_type(i32::try_from(attribute.i).map_err(|_| { + OnnxError::InvalidAttribute { + attr: "to".to_string(), + op: "Cast".to_string(), + node: node.name.clone(), + reason: "dtype code is out of range".to_string(), + } + })?)?; + } + } + } + Ok(()) +} + /// Infer output shape for an ONNX node based on its operation type and inputs fn infer_shape( node: &crate::protos::onnx::NodeProto, @@ -191,11 +254,11 @@ fn infer_shape( .as_slice() .iter() .find(|a| a.name.as_str() == "keepdims") - .and_then(|a| if a.i != 0 { Some(a.i != 0) } else { None }) + .map(|a| a.i != 0) .unwrap_or(true); // Get axes attribute - let axes: Vec = node + let mut axes: Vec = node .attribute .as_slice() .iter() @@ -203,6 +266,18 @@ fn infer_shape( .map(|a| a.ints.clone()) .unwrap_or_default(); + if let Some(input) = ins.get(1).filter(|name| !name.is_empty()) { + axes = const_values.get(input)?.clone(); + } + if axes.is_empty() + && node + .attribute + .iter() + .any(|a| a.name == "noop_with_empty_axes" && a.i != 0) + { + return Some(input_shape.clone()); + } + if axes.is_empty() { // Reduce all dimensions if keepdims { @@ -1393,6 +1468,32 @@ impl OnnxConverter { } let onnx_graph = self.model.graph.as_ref().unwrap(); + let standard_opset = self + .model + .opset_import + .iter() + .find(|import| import.domain.is_empty() || import.domain == "ai.onnx") + .map(|import| import.version); + validate_extended_opset_types(&self.model)?; + for node in &onnx_graph.node { + if node.op_type == "Gelu" + && (node.domain.is_empty() || node.domain == "ai.onnx") + && standard_opset.is_none_or(|version| version < 20) + { + return Err(OnnxError::UnsupportedOp { + op: "Gelu requires ai.onnx opset 20".to_string(), + node: node.name.clone(), + }); + } + } + let mut reserved_ids: HashSet = onnx_graph + .node + .iter() + .flat_map(|node| node.input.iter().chain(&node.output)) + .chain(onnx_graph.input.iter().map(|value| &value.name)) + .chain(onnx_graph.initializer.iter().map(|value| &value.name)) + .map(|name| sanitize_identifier(name)) + .collect(); let mut value_name_map: HashMap = HashMap::new(); let mut effective_overrides = options.free_dim_overrides.clone(); let mut inference_overrides = effective_overrides.clone(); @@ -1569,10 +1670,6 @@ Provide --override-dim {}= or enable --experimental-dynamic-inputs.", ))); }; - if shape.is_empty() { - continue; - } - self.graph.inputs.insert( name.clone(), crate::ast::OperandDesc { @@ -2786,7 +2883,8 @@ Provide --override-dim {}= or enable --experimental-dynamic-inputs.", value_types: &value_types, }; - let converted = registry.convert_node(onnx_node, &context)?; + let mut converted = registry.convert_node(onnx_node, &context)?; + converted.reserve_private_values(&mut reserved_ids); for (name, mut decl) in converted.consts { if let crate::ast::ConstInit::InlineBytes { bytes } = &decl.init { @@ -2895,6 +2993,8 @@ pub fn convert_onnx>( let mut model: ModelProto = ModelProto::decode(&onnx_bytes[..]).map_err(|e| OnnxError::ProtobufError(e.to_string()))?; + validate_extended_opset_types(&model)?; + // Apply constant folding if optimize flag is set if options.optimize { crate::debug_println!("Running constant folding..."); diff --git a/src/onnx/ops/activation.rs b/src/onnx/ops/activation.rs index d661a14..4b31fce 100644 --- a/src/onnx/ops/activation.rs +++ b/src/onnx/ops/activation.rs @@ -1,10 +1,11 @@ // Activation and unary math operators: Relu, Gelu, Tanh, Sigmoid, Sqrt, Exp, Log, Abs, Neg, Erf -use crate::ast::Node; +use crate::ast::{ConstDecl, ConstInit, DataType, Node}; use crate::onnx::convert::{sanitize_identifier, OnnxError}; use crate::onnx::ops::{ConversionContext, ConversionResult, OpHandler}; use crate::protos::onnx::NodeProto; -use serde_json::Map; +use serde_json::{json, Map}; +use std::collections::HashSet; pub struct ActivationHandler; @@ -40,6 +41,38 @@ impl OpHandler for ActivationHandler { "unnamed".to_string() }; + if op_type == "Gelu" { + if !matches!(node.domain.as_str(), "" | "ai.onnx" | "com.microsoft") { + return Err(OnnxError::UnsupportedOp { + op: format!("{}::Gelu", node.domain), + node: node_name, + }); + } + if node.input.len() != 1 + || node.output.len() != 1 + || node.input[0].is_empty() + || node.output[0].is_empty() + { + return Err(OnnxError::InvalidShape(format!( + "Gelu '{}' expects one input and one output", + node_name + ))); + } + if context + .value_types + .get(&node.input[0]) + .is_some_and(|dtype| !matches!(dtype, DataType::Float16 | DataType::Float32)) + { + return Err(OnnxError::UnsupportedOp { + op: "Gelu requires float16 or float32 input".to_string(), + node: node_name, + }); + } + if Self::gelu_uses_tanh(node, &node_name)? { + return self.convert_tanh_gelu(node, &node_name, context); + } + } + // Map ONNX operator to WebNN operation name let webnn_op = match op_type { "Relu" => "relu", @@ -68,6 +101,158 @@ impl OpHandler for ActivationHandler { } impl ActivationHandler { + fn gelu_uses_tanh(node: &NodeProto, node_name: &str) -> Result { + let invalid = |reason: &str| OnnxError::InvalidAttribute { + attr: "approximate".to_string(), + op: "Gelu".to_string(), + node: node_name.to_string(), + reason: reason.to_string(), + }; + let mut attributes = node.attribute.iter().filter(|a| a.name == "approximate"); + let Some(attribute) = attributes.next() else { + return Ok(false); + }; + if node.domain == "com.microsoft" { + return Err(invalid("com.microsoft Gelu does not define this attribute")); + } + if attributes.next().is_some() { + return Err(invalid("attribute must not be repeated")); + } + if attribute.r#type != crate::protos::onnx::attribute_proto::AttributeType::String as i32 { + return Err(invalid("expected a string")); + } + match attribute.s.as_slice() { + b"none" => Ok(false), + b"tanh" => Ok(true), + _ => Err(invalid("expected 'none' or 'tanh'")), + } + } + + fn convert_tanh_gelu( + &self, + node: &NodeProto, + node_name: &str, + context: &ConversionContext, + ) -> Result { + let input = context.resolve_input(&node.input[0]); + let output = sanitize_identifier(&node.output[0]); + let dtype = context + .value_types + .get(&node.input[0]) + .or_else(|| context.value_types.get(&input)) + .ok_or_else(|| { + OnnxError::InvalidShape(format!("Gelu '{}' requires a known input type", node_name)) + })?; + if !matches!(dtype, DataType::Float16 | DataType::Float32) { + return Err(OnnxError::UnsupportedOp { + op: format!("Gelu with {:?} input", dtype), + node: node_name.to_string(), + }); + } + + let mut used: HashSet = context.value_ids.values().cloned().collect(); + used.insert(input.clone()); + used.insert(output.clone()); + let mut private_values = Vec::new(); + let mut fresh = |suffix: &str| { + let base = format!("{}__gelu_{}", output, suffix); + let mut id = base.clone(); + let mut index = 1; + while !used.insert(id.clone()) { + id = format!("{}_{}", base, index); + index += 1; + } + private_values.push(id.clone()); + id + }; + let mut result = ConversionResult::default(); + let x = if *dtype == DataType::Float16 { + let id = fresh("float32"); + result.nodes.push(Node { + id: id.clone(), + op: "cast".to_string(), + inputs: vec![input], + options: Map::from_iter([("to".to_string(), json!("float32"))]), + outputs: None, + }); + id + } else { + input + }; + + // WebNN gelu is erf-based. Preserve ONNX's separate tanh formula using + // primitive operations. Promote half inputs for the polynomial and + // cancellation near the negative tail, then round only the result. + let mut scalar = |suffix: &str, value: f32| { + let id = fresh(suffix); + result.consts.push(( + id.clone(), + ConstDecl { + data_type: DataType::Float32, + shape: vec![], + init: ConstInit::InlineBytes { + bytes: value.to_le_bytes().to_vec(), + }, + }, + )); + id + }; + let half = scalar("half", 0.5); + let one = scalar("one", 1.0); + let coefficient = scalar("coefficient", 0.044715); + // Round the mathematical coefficient once, not pi and the division + // separately; the latter gives the preceding float32 value. + let scale = scalar("scale", (2.0_f64 / std::f64::consts::PI).sqrt() as f32); + let mut operation = |suffix: &str, op: &str, inputs: Vec| { + let id = fresh(suffix); + result.nodes.push(Node { + id: id.clone(), + op: op.to_string(), + inputs, + options: Map::new(), + outputs: None, + }); + id + }; + let square = operation("square", "mul", vec![x.clone(), x.clone()]); + let cube = operation("cube", "mul", vec![square, x.clone()]); + let cubic = operation("cubic", "mul", vec![coefficient, cube]); + let polynomial = operation("polynomial", "add", vec![x.clone(), cubic]); + let scaled = operation("scaled", "mul", vec![scale, polynomial]); + let tanh = operation("tanh", "tanh", vec![scaled]); + let gate = operation("gate", "add", vec![one, tanh]); + let half_x = operation("half_x", "mul", vec![half, x]); + let result_id = if *dtype == DataType::Float16 { + fresh("result") + } else { + output.clone() + }; + result.nodes.push(Node { + id: result_id.clone(), + op: "mul".to_string(), + inputs: vec![half_x, gate], + options: Map::new(), + outputs: None, + }); + if *dtype == DataType::Float16 { + result.nodes.push(Node { + id: output.clone(), + op: "cast".to_string(), + inputs: vec![result_id], + options: Map::from_iter([("to".to_string(), json!("float16"))]), + outputs: None, + }); + } + result + .output_mappings + .insert(node.output[0].clone(), output); + result + .output_types + .insert(node.output[0].clone(), dtype.clone()); + result.private_values = private_values; + Ok(result) + } + /// Convert ONNX unary/activation operation to WebNN fn convert_unary( &self, @@ -217,6 +402,87 @@ mod tests { assert_eq!(result.nodes[0].op, "gelu"); } + fn convert_gelu_attributes( + attributes: Vec, + ) -> Result { + let mut node = create_test_node("Gelu", vec!["x"], vec!["y"]); + node.attribute = attributes; + let initializers = std::collections::HashMap::new(); + let value_shapes = std::collections::HashMap::new(); + let const_values = std::collections::HashMap::new(); + let value_ids = std::collections::HashMap::new(); + let value_types = std::collections::HashMap::from([("x".to_string(), DataType::Float32)]); + ActivationHandler.convert( + &node, + &ConversionContext { + initializers: &initializers, + value_shapes: &value_shapes, + value_shape_dims: crate::onnx::ops::empty_value_shape_dims(), + const_values: &const_values, + value_ids: &value_ids, + value_types: &value_types, + }, + ) + } + + fn approximate(value: &[u8]) -> crate::protos::onnx::AttributeProto { + crate::protos::onnx::AttributeProto { + name: "approximate".to_string(), + r#type: 3, // AttributeProto::STRING + s: value.to_vec(), + ..Default::default() + } + } + + #[test] + fn test_gelu_default_and_none_preserve_exact_operation() { + for attributes in [vec![], vec![approximate(b"none")]] { + let result = convert_gelu_attributes(attributes).expect("exact GELU"); + assert_eq!(result.nodes.len(), 1); + assert_eq!(result.nodes[0].op, "gelu"); + assert_eq!(result.nodes[0].inputs, ["x"]); + assert!(result.nodes[0].options.is_empty()); + assert_eq!(result.output_mappings.get("y"), Some(&"y".to_string())); + } + } + + #[test] + fn test_gelu_tanh_is_not_silently_replaced_with_exact_gelu() { + let result = convert_gelu_attributes(vec![approximate(b"tanh")]).unwrap(); + assert!(result.nodes.iter().any(|node| node.op == "tanh")); + assert!(result.nodes.iter().all(|node| node.op != "gelu")); + assert_eq!(result.consts.len(), 4); + } + + #[test] + fn test_gelu_rejects_invalid_or_malformed_approximation() { + let mut wrong_type = approximate(b"none"); + wrong_type.r#type = 2; // INT, even if the string field is populated. + let mut missing_type = approximate(b"none"); + missing_type.r#type = 0; + for attributes in [ + vec![approximate(b"invalid")], + vec![approximate(b"")], + vec![approximate(b"TANH")], + vec![approximate(&[0xff])], + vec![wrong_type], + vec![missing_type], + vec![approximate(b"none"), approximate(b"tanh")], + ] { + let error = convert_gelu_attributes(attributes) + .expect_err("invalid approximation must not become exact GELU"); + assert!(matches!( + &error, + OnnxError::InvalidAttribute { attr, op, node, .. } + if attr == "approximate" && op == "Gelu" && node == "test_gelu" + )); + let message = error.to_string(); + assert!(message.contains("approximate"), "{message}"); + assert!(message.contains("Gelu"), "{message}"); + assert!(message.contains("test_gelu"), "{message}"); + } + } + #[test] fn test_convert_cos() { let handler = ActivationHandler; diff --git a/src/onnx/ops/matmul.rs b/src/onnx/ops/matmul.rs index 4e18bce..9a48150 100644 --- a/src/onnx/ops/matmul.rs +++ b/src/onnx/ops/matmul.rs @@ -287,6 +287,7 @@ impl MatMulHandler { consts, output_mappings: std::collections::HashMap::new(), output_types: std::collections::HashMap::new(), + private_values: Vec::new(), }; if let Some(output) = node.output.as_slice().first() { diff --git a/src/onnx/ops/mod.rs b/src/onnx/ops/mod.rs index 3d0166d..436ad84 100644 --- a/src/onnx/ops/mod.rs +++ b/src/onnx/ops/mod.rs @@ -3,7 +3,7 @@ use crate::ast::{ConstDecl, Node}; use crate::onnx::convert::OnnxError; use crate::protos::onnx::{NodeProto, TensorProto}; -use std::collections::HashMap; +use std::collections::{HashMap, HashSet}; use std::sync::OnceLock; pub mod activation; @@ -118,6 +118,9 @@ pub struct ConversionResult { pub output_mappings: HashMap, /// ONNX output name -> data type pub output_types: HashMap, + /// Private values introduced by a decomposition, renamed against the whole + /// model before insertion so future ONNX outputs cannot collide with them. + pub private_values: Vec, } impl ConversionResult { @@ -127,8 +130,45 @@ impl ConversionResult { consts: Vec::new(), output_mappings: HashMap::new(), output_types: HashMap::new(), + private_values: Vec::new(), } } + + pub(crate) fn reserve_private_values(&mut self, reserved: &mut HashSet) { + let mut renamed = HashMap::new(); + for original in &self.private_values { + let mut candidate = original.clone(); + let mut suffix = 1; + while !reserved.insert(candidate.clone()) { + candidate = format!("{}_{}", original, suffix); + suffix += 1; + } + renamed.insert(original.clone(), candidate); + } + let rename = |id: &mut String| { + if let Some(replacement) = renamed.get(id) { + *id = replacement.clone(); + } + }; + for (id, _) in &mut self.consts { + rename(id); + } + for node in &mut self.nodes { + rename(&mut node.id); + for input in &mut node.inputs { + rename(input); + } + if let Some(outputs) = &mut node.outputs { + for output in outputs { + rename(output); + } + } + } + // Keep all inserted IDs reserved, including values produced by handlers + // that do not yet register their private intermediates explicitly. + reserved.extend(self.nodes.iter().map(|node| node.id.clone())); + reserved.extend(self.consts.iter().map(|(id, _)| id.clone())); + } } /// Trait for handling ONNX operator conversion diff --git a/src/onnx/ops/pool.rs b/src/onnx/ops/pool.rs index dd2ac72..f1663cc 100644 --- a/src/onnx/ops/pool.rs +++ b/src/onnx/ops/pool.rs @@ -5,7 +5,7 @@ // ONNX MaxPool / AveragePool attributes (spatial-rank-aware): // * kernel_shape: required, length = spatial_rank // * strides: default = [1; spatial_rank] -// * dilations: default = [1; spatial_rank] (MaxPool only) +// * dilations: default = [1; spatial_rank] (AveragePool since opset 19) // * pads: default = [0; 2*spatial_rank], layout [b1, b2, ..., e1, e2, ...] // * auto_pad: NOTSET | SAME_UPPER | SAME_LOWER | VALID // * ceil_mode: 0 (floor) | 1 (ceil) @@ -255,7 +255,7 @@ impl PoolHandler { options.insert("windowDimensions".to_string(), json!(kernel)); options.insert("strides".to_string(), json!(strides)); - // AveragePool in ONNX has no dilations; only emit dilations when non-default + // AveragePool gained dilation in opset 19; emit non-default values. // to keep generated calls minimal for the average case. if matches!(kind, PoolKind::Max) || dilations.iter().any(|&d| d != 1) { options.insert("dilations".to_string(), json!(dilations)); diff --git a/src/onnx/ops/reduction.rs b/src/onnx/ops/reduction.rs index 8100a78..8be2533 100644 --- a/src/onnx/ops/reduction.rs +++ b/src/onnx/ops/reduction.rs @@ -2,9 +2,7 @@ use crate::ast::Node; use crate::onnx::convert::{sanitize_identifier, OnnxError}; -use crate::onnx::ops::{ - normalize_axes_best_effort, ConversionContext, ConversionResult, OpHandler, -}; +use crate::onnx::ops::{normalize_axes, ConversionContext, ConversionResult, OpHandler}; use crate::protos::onnx::NodeProto; use serde_json::Map; @@ -53,9 +51,9 @@ impl ReductionHandler { context: &ConversionContext, ) -> Result { let inputs = node.input.as_slice(); - if inputs.is_empty() { + if inputs.is_empty() || inputs.len() > 2 { return Err(OnnxError::InvalidShape(format!( - "{} expects at least 1 input", + "{} expects 1 or 2 inputs", webnn_op ))); } @@ -63,18 +61,49 @@ impl ReductionHandler { // Extract attributes let mut axes: Option> = None; let mut keepdims = 1i64; // ONNX default is 1 (keep dimensions) + let mut noop_with_empty_axes = false; for attr in node.attribute.as_slice() { match attr.name.as_str() { "axes" => { axes = Some(attr.ints.clone()); } - "keepdims" if attr.i != 0 => { + "keepdims" => { keepdims = attr.i; } + "noop_with_empty_axes" => noop_with_empty_axes = attr.i != 0, _ => {} } } + if let Some(axes_input) = inputs.get(1).filter(|name| !name.is_empty()) { + if axes.is_some() { + return Err(OnnxError::InvalidShape(format!( + "{} '{}' specifies both an axes input and attribute", + webnn_op, node_name + ))); + } + axes = Some( + context + .const_values + .get(axes_input) + .cloned() + .ok_or_else(|| OnnxError::UnsupportedOp { + op: format!("{} with nonconstant axes", node.op_type), + node: node_name.to_string(), + })?, + ); + } + let empty_axes = axes.as_ref().is_none_or(Vec::is_empty); + // Identity is correct for the four supported reductions only. Composite + // reductions must retain their non-reduction steps (for example, + // ReduceLogSum still takes log and ReduceSumSquare still squares). + // They are rejected by supports()/convert(), not routed through here. + let no_op = empty_axes && noop_with_empty_axes; + // ONNX reduces every dimension for absent/empty axes unless noop=1; + // WebNN's explicitly empty axes instead mean no reduction. + if empty_axes { + axes = None; + } let output_name = if node.output.as_slice().is_empty() { format!("{}_output", node_name) @@ -89,7 +118,7 @@ impl ReductionHandler { // Add axes if specified if let Some(axes_values) = axes { let axes_values = if let Some(rank) = context.input_rank(inputs[0].as_str()) { - normalize_axes_best_effort(&axes_values, rank) + normalize_axes(&axes_values, rank)? } else { axes_values }; @@ -104,9 +133,9 @@ impl ReductionHandler { let mut result = ConversionResult::new(vec![Node { id: output_name.clone(), - op: webnn_op.to_string(), + op: if no_op { "identity" } else { webnn_op }.to_string(), inputs: vec![input0], - options, + options: if no_op { Map::new() } else { options }, outputs: None, }]); diff --git a/src/onnx/shape_inference.rs b/src/onnx/shape_inference.rs index 1fec275..34affcd 100644 --- a/src/onnx/shape_inference.rs +++ b/src/onnx/shape_inference.rs @@ -156,7 +156,7 @@ fn seed_initializers( DataType::Int32 | DataType::Int64 | DataType::Uint32 | DataType::Uint64 ) { let values = read_int_tensor(init); - if !values.is_empty() { + if !values.is_empty() || init.dims.contains(&0) { result.const_values.insert(name, values); } } @@ -871,21 +871,31 @@ fn infer_node_shape(node: &NodeProto, ctx: &InferenceResult) -> Option> "ReduceMean" | "ReduceSum" | "ReduceMax" | "ReduceMin" => { let input = node.input.as_slice().first()?; let input_shape = ctx.value_shapes.get(input)?; - let axes: Vec = node + let mut axes: Vec = node .attribute .as_slice() .iter() .find(|a| a.name.as_str() == "axes") .map(|a| a.ints.clone()) .unwrap_or_default(); + if let Some(input) = node.input.get(1).filter(|name| !name.is_empty()) { + axes = ctx.const_values.get(input)?.clone(); + } let keepdims = node .attribute .as_slice() .iter() - .find(|a| a.name.as_str() == "keepdims" && a.i != 0) + .find(|a| a.name.as_str() == "keepdims") .map(|a| a.i != 0) .unwrap_or(true); if axes.is_empty() { + if node + .attribute + .iter() + .any(|a| a.name == "noop_with_empty_axes" && a.i != 0) + { + return Some(input_shape.clone()); + } if keepdims { Some(vec![1; input_shape.len()]) } else { diff --git a/tests/onnx_gelu.rs b/tests/onnx_gelu.rs new file mode 100644 index 0000000..1937f27 --- /dev/null +++ b/tests/onnx_gelu.rs @@ -0,0 +1,424 @@ +#![cfg(feature = "onnx")] + +use prost::Message; +use std::collections::HashMap; +use webnn_graph::ast::{ConstInit, DataType, GraphJson}; +use webnn_graph::emit_js::emit_builder_js; +use webnn_graph::onnx::convert::{convert_onnx, ConvertOptions, OnnxConverter, OnnxError}; +use webnn_graph::protos::onnx::{ + tensor_shape_proto, type_proto, AttributeProto, GraphProto, ModelProto, NodeProto, + OperatorSetIdProto, TensorShapeProto, TypeProto, ValueInfoProto, +}; +use webnn_graph::validate::validate_graph; + +fn model(domain: &str, opset: i64, dtype: i32, approximate: Option<&str>) -> ModelProto { + let tensor_type = TypeProto { + value: Some(type_proto::Value::TensorType(type_proto::Tensor { + elem_type: dtype, + shape: Some(TensorShapeProto { + dim: vec![tensor_shape_proto::Dimension { + value: Some(tensor_shape_proto::dimension::Value::DimValue(3)), + ..Default::default() + }], + }), + })), + ..Default::default() + }; + let value = |name: &str| ValueInfoProto { + name: name.to_string(), + r#type: Some(tensor_type.clone()), + ..Default::default() + }; + let attribute = approximate + .map(|value| AttributeProto { + name: "approximate".to_string(), + r#type: 3, + s: value.as_bytes().to_vec(), + ..Default::default() + }) + .into_iter() + .collect(); + ModelProto { + ir_version: 9, + graph: Some(GraphProto { + name: "gelu_regression".to_string(), + input: vec![value("x")], + output: vec![value("y")], + node: vec![NodeProto { + name: "activation".to_string(), + op_type: "Gelu".to_string(), + domain: domain.to_string(), + input: vec!["x".to_string()], + output: vec!["y".to_string()], + attribute, + ..Default::default() + }], + ..Default::default() + }), + opset_import: vec![OperatorSetIdProto { + domain: domain.to_string(), + version: opset, + }], + ..Default::default() + } +} + +#[test] +fn standard_gelu_opset20_converts_both_approximations() { + for approximate in [None, Some("none"), Some("tanh")] { + for optimize in [false, true] { + let graph = OnnxConverter::new(model("", 20, 1, approximate)) + .unwrap() + .convert(&ConvertOptions { + optimize, + ..Default::default() + }) + .unwrap(); + validate_graph(&graph).unwrap(); + if approximate == Some("tanh") { + assert!(graph.nodes.iter().any(|node| node.op == "tanh")); + assert!(graph.nodes.iter().all(|node| node.op != "gelu")); + } else { + assert_eq!(graph.nodes[0].op, "gelu"); + } + } + } +} + +#[test] +fn legacy_exact_gelu_still_converts_validates_and_emits_js() { + // The older com.microsoft Gelu is exact and has no approximate attribute. + for (dtype, expected_type) in [(1, DataType::Float32), (10, DataType::Float16)] { + for optimize in [false, true] { + let graph = OnnxConverter::new(model("com.microsoft", 1, dtype, None)) + .unwrap() + .convert(&ConvertOptions { + optimize, + ..Default::default() + }) + .unwrap(); + validate_graph(&graph).unwrap(); + assert_eq!(graph.inputs["x"].data_type, expected_type); + assert_eq!(graph.nodes.len(), 1); + assert_eq!(graph.nodes[0].op, "gelu"); + assert_eq!(graph.nodes[0].inputs, ["x"]); + assert!(graph.nodes[0].options.is_empty()); + assert!(emit_builder_js(&graph).contains("builder[\"gelu\"](env.get(\"x\"), {})")); + } + } +} + +#[test] +fn legacy_gelu_rejects_an_unsupported_approximation_attribute() { + // Deliberately malformed legacy nodes used to reach the handler and have + // their approximation discarded. Optimizing must not hide this error. + for approximate in ["tanh", "invalid"] { + for optimize in [false, true] { + let error = OnnxConverter::new(model("com.microsoft", 1, 1, Some(approximate))) + .unwrap() + .convert(&ConvertOptions { + optimize, + ..Default::default() + }) + .unwrap_err(); + assert!(matches!(error, OnnxError::InvalidAttribute { .. })); + assert!(error.to_string().contains("activation")); + } + } +} + +fn convert(model: ModelProto, optimize: bool) -> GraphJson { + let directory = tempfile::tempdir().unwrap(); + let path = directory.path().join("gelu.onnx"); + std::fs::write(&path, model.encode_to_vec()).unwrap(); + let graph = convert_onnx( + &path, + ConvertOptions { + extract_weights: false, + optimize, + ..Default::default() + }, + ) + .unwrap(); + validate_graph(&graph).unwrap(); + graph +} + +fn evaluate(graph: &GraphJson, input: f32) -> f32 { + let mut values = HashMap::from([("x".to_string(), input)]); + for (name, constant) in &graph.consts { + assert_eq!(constant.data_type, DataType::Float32); + assert!(constant.shape.is_empty()); + let ConstInit::InlineBytes { bytes } = &constant.init else { + panic!("inline scalar") + }; + values.insert( + name.clone(), + f32::from_le_bytes(bytes.as_slice().try_into().unwrap()), + ); + } + for node in &graph.nodes { + let x = values[&node.inputs[0]]; + let result = match node.op.as_str() { + "mul" => x * values[&node.inputs[1]], + "add" => x + values[&node.inputs[1]], + "tanh" => x.tanh(), + "identity" => x, + "cast" => match node.options["to"].as_str().unwrap() { + "float16" => half::f16::from_f32(x).to_f32(), + "float32" => x, + other => panic!("unexpected dtype {other}"), + }, + other => panic!("unexpected operation {other}"), + }; + values.insert(node.id.clone(), result); + } + values[&graph.outputs["y"]] +} + +#[test] +fn tanh_gelu_matches_onnx_runtime_including_half_precision_tails() { + // ONNX checker-valid Gelu-20, ORT 1.30 CPUExecutionProvider. Values at + // +/-2.707 separate exact/tanh GELU; the far tails expose half overflow. + let inputs = [ + -65504.0, -256.0, -20.0, -10.0, -5.0, -4.0, -3.0, -2.707, -2.0, -1.0, -0.5, -0.1, 0.0, 0.1, + 0.5, 1.0, 2.0, 2.707, 3.0, 4.0, 5.0, 10.0, 20.0, 256.0, 65504.0, + ]; + let expected: [f32; 25] = [ + -0.0, + -0.0, + -0.0, + -0.0, + -2.9802322e-7, + -7.009506e-5, + -0.0036375225, + -0.008716196, + -0.045402348, + -0.15880796, + -0.154286, + -0.046017252, + 0.0, + 0.053982753, + 0.345714, + 0.841192, + 1.9545977, + 2.698284, + 2.9963627, + 3.99993, + 4.9999995, + 10.0, + 20.0, + 256.0, + 65504.0, + ]; + let expected_half: [f32; 25] = [ + -0.0, + -0.0, + -0.0, + -0.0, + -2.9802322e-7, + -7.021427e-5, + -0.0036373138, + -0.008712769, + -0.045410156, + -0.15881348, + -0.15429688, + -0.046020508, + 0.0, + 0.053955078, + 0.34570313, + 0.8413086, + 1.9550781, + 2.6992188, + 2.9960938, + 4.0, + 5.0, + 10.0, + 20.0, + 256.0, + 65504.0, + ]; + for optimize in [false, true] { + for (dtype, reference) in [(1, &expected), (10, &expected_half)] { + let graph = convert(model("", 20, dtype, Some("tanh")), optimize); + let scale = graph + .consts + .iter() + .find(|(id, _)| id.contains("__gelu_scale")) + .unwrap() + .1; + let ConstInit::InlineBytes { bytes } = &scale.init else { + panic!("inline scale") + }; + assert_eq!( + u32::from_le_bytes(bytes.as_slice().try_into().unwrap()), + 0x3f4c422a + ); + for (&input, &expected) in inputs.iter().zip(reference) { + let input = if dtype == 10 { + half::f16::from_f32(input).to_f32() + } else { + input + }; + let actual = evaluate(&graph, input); + let tolerance = if dtype == 10 { + 2.0e-7 + expected.abs() * 1.0e-3 + } else { + 2.0e-7 + expected.abs() * 1.0e-6 + }; + assert!( + (actual - expected).abs() <= tolerance, + "dtype={dtype}, input={input}, actual={actual}, expected={expected}" + ); + } + assert!(evaluate(&graph, f32::NAN).is_nan()); + assert!(evaluate(&graph, f32::NEG_INFINITY).is_nan()); + assert_eq!(evaluate(&graph, f32::INFINITY), f32::INFINITY); + assert_eq!(evaluate(&graph, -0.0).to_bits(), (-0.0_f32).to_bits()); + } + } +} + +#[test] +fn tanh_gelu_keeps_scalar_inputs_and_rejects_unknown_rank() { + for optimize in [false, true] { + for dtype in [1, 10] { + let mut scalar = model("", 20, dtype, Some("tanh")); + let graph = scalar.graph.as_mut().unwrap(); + for value in graph.input.iter_mut().chain(&mut graph.output) { + let Some(type_proto::Value::TensorType(tensor)) = + &mut value.r#type.as_mut().unwrap().value + else { + unreachable!() + }; + tensor.shape.as_mut().unwrap().dim.clear(); + } + let result = convert(scalar.clone(), optimize); + assert!(result.inputs["x"].shape.is_empty()); + assert!((evaluate(&result, 1.0) - 0.8412).abs() < 0.0002); + let Some(type_proto::Value::TensorType(tensor)) = + &mut scalar.graph.as_mut().unwrap().input[0] + .r#type + .as_mut() + .unwrap() + .value + else { + unreachable!() + }; + tensor.shape = None; + assert!(OnnxConverter::new(scalar) + .unwrap() + .convert(&ConvertOptions::default()) + .is_err()); + } + } +} + +#[test] +fn tanh_gelu_private_values_cannot_shadow_future_outputs_or_inputs() { + for optimize in [false, true] { + let mut source = model("", 20, 1, Some("tanh")); + let graph = source.graph.as_mut().unwrap(); + graph.node.push(NodeProto { + op_type: "Identity".to_string(), + input: vec!["y".to_string()], + output: vec!["y__gelu_half".to_string()], + ..Default::default() + }); + let mut extra_input = graph.input[0].clone(); + extra_input.name = "y__gelu_square".to_string(); + graph.input.push(extra_input); + let converted = convert(source, optimize); + assert!(!converted.consts.contains_key("y__gelu_half")); + assert_eq!( + converted + .nodes + .iter() + .filter(|node| node.id == "y__gelu_half") + .count(), + 1 + ); + assert!(converted + .nodes + .iter() + .all(|node| node.id != "y__gelu_square")); + assert!((evaluate(&converted, 1.0) - 0.841192).abs() < 2.0e-7); + } +} + +#[test] +fn unsupported_opsets_types_and_invalid_gelu_modes_fail_closed() { + for version in [10, 21] { + assert!(matches!( + OnnxConverter::new(model("", version, 1, None)) + .unwrap() + .convert(&ConvertOptions::default()), + Err(OnnxError::UnsupportedOpset { .. }) + )); + } + for dtype in [11, 16, 17, 18, 19, 20] { + assert!( + OnnxConverter::new(model("", 20, dtype, Some("tanh"))) + .unwrap() + .convert(&ConvertOptions::default()) + .is_err(), + "dtype {dtype}" + ); + } + for mode in ["invalid", "", "TANH"] { + assert!(matches!( + OnnxConverter::new(model("", 20, 1, Some(mode))) + .unwrap() + .convert(&ConvertOptions::default()), + Err(OnnxError::InvalidAttribute { .. }) + )); + } + assert!(OnnxConverter::new(model("", 18, 1, None)) + .unwrap() + .convert(&ConvertOptions::default()) + .is_err()); +} + +#[test] +fn emitted_tanh_gelu_executes_with_webnn_shaped_builder_signatures() { + use std::io::Write; + use std::process::{Command, Stdio}; + if Command::new("node").arg("--version").output().is_err() { + eprintln!("Node.js unavailable: emitted JavaScript execution not tested"); + return; + } + for dtype in [1, 10] { + let graph = convert(model("", 20, dtype, Some("tanh")), true); + let mut child = Command::new("node") + .args(["--input-type=module", "-"]) + .stdin(Stdio::piped()) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .spawn() + .unwrap(); + let mut script = include_str!("support/gelu_builder.mjs").to_string(); + script.push_str(&emit_builder_js(&graph)); + script.push_str("\nconst result = await buildGraph({x: [-2.707, 0, 2.707]}); console.log(JSON.stringify(result.y.data));\n"); + child + .stdin + .take() + .unwrap() + .write_all(script.as_bytes()) + .unwrap(); + let output = child.wait_with_output().unwrap(); + assert!( + output.status.success(), + "{}", + String::from_utf8_lossy(&output.stderr) + ); + let actual: Vec = serde_json::from_slice(&output.stdout).unwrap(); + for (&actual, input) in actual.iter().zip([-2.707_f32, 0.0, 2.707]) { + let input = if dtype == 10 { + half::f16::from_f32(input).to_f32() + } else { + input + }; + assert!((actual - evaluate(&graph, input)).abs() < 1.0e-6); + } + } +} diff --git a/tests/onnx_opset20.rs b/tests/onnx_opset20.rs new file mode 100644 index 0000000..9f3a90a --- /dev/null +++ b/tests/onnx_opset20.rs @@ -0,0 +1,322 @@ +#![cfg(feature = "onnx")] + +use prost::Message; +use std::collections::HashMap; +use webnn_graph::onnx::convert::{convert_onnx, ConvertOptions, OnnxConverter, OnnxError}; +use webnn_graph::onnx::shape_inference::infer_static_shapes; +use webnn_graph::protos::onnx::{ + tensor_shape_proto, type_proto, AttributeProto, GraphProto, ModelProto, NodeProto, + OperatorSetIdProto, TensorProto, TensorShapeProto, TypeProto, ValueInfoProto, +}; +use webnn_graph::validate::validate_graph; + +fn value(name: &str, dtype: i32, shape: &[i64]) -> ValueInfoProto { + ValueInfoProto { + name: name.to_string(), + r#type: Some(TypeProto { + value: Some(type_proto::Value::TensorType(type_proto::Tensor { + elem_type: dtype, + shape: Some(TensorShapeProto { + dim: shape + .iter() + .map(|&dim| tensor_shape_proto::Dimension { + value: Some(tensor_shape_proto::dimension::Value::DimValue(dim)), + ..Default::default() + }) + .collect(), + }), + })), + ..Default::default() + }), + ..Default::default() + } +} + +fn model(op: &str, version: i64, input_shape: &[i64], output_shape: &[i64]) -> ModelProto { + ModelProto { + ir_version: 9, + opset_import: vec![OperatorSetIdProto { + domain: String::new(), + version, + }], + graph: Some(GraphProto { + name: "opset_audit".to_string(), + input: vec![value("x", 1, input_shape)], + output: vec![value("y", 1, output_shape)], + node: vec![NodeProto { + op_type: op.to_string(), + name: "audited_operator".to_string(), + input: vec!["x".to_string()], + output: vec!["y".to_string()], + ..Default::default() + }], + ..Default::default() + }), + ..Default::default() + } +} + +fn integer(name: &str, value: i64) -> AttributeProto { + AttributeProto { + name: name.to_string(), + r#type: 2, + i: value, + ..Default::default() + } +} + +#[test] +fn opset19_and_20_preserve_average_pool_dilation() { + for version in [19, 20] { + let mut model = model("AveragePool", version, &[1, 1, 5, 5], &[1, 1, 3, 3]); + model.graph.as_mut().unwrap().node[0].attribute = ["kernel_shape", "dilations"] + .into_iter() + .map(|name| AttributeProto { + name: name.to_string(), + r#type: 7, + ints: vec![2, 2], + ..Default::default() + }) + .collect(); + let graph = OnnxConverter::new(model) + .unwrap() + .convert(&ConvertOptions::default()) + .unwrap(); + validate_graph(&graph).unwrap(); + assert_eq!(graph.nodes[0].op, "averagePool2d"); + assert_eq!( + graph.nodes[0].options["dilations"], + serde_json::json!([2, 2]) + ); + } +} + +#[test] +fn modern_reductions_preserve_constant_axes_and_keepdims_zero() { + for op in ["ReduceMean", "ReduceSum", "ReduceMin", "ReduceMax"] { + for optimize in [false, true] { + let mut model = model(op, 20, &[2, 3], &[2]); + let graph = model.graph.as_mut().unwrap(); + graph.node[0].input.push("axes".to_string()); + graph.node[0].attribute = vec![integer("keepdims", 0)]; + graph.initializer.push(TensorProto { + name: "axes".to_string(), + data_type: 7, + dims: vec![1], + raw_data: (-1_i64).to_le_bytes().to_vec(), + ..Default::default() + }); + let inferred = infer_static_shapes(&model, &HashMap::new()).unwrap(); + assert_eq!(inferred.value_shapes["y"], [2]); + let graph = OnnxConverter::new(model) + .unwrap() + .convert(&ConvertOptions { + optimize, + ..Default::default() + }) + .unwrap(); + validate_graph(&graph).unwrap(); + let reduction = graph.nodes.iter().find(|node| node.id == "y").unwrap(); + assert_eq!(reduction.options["axes"], serde_json::json!([1])); + assert_eq!(reduction.options["keepDimensions"], false); + } + } +} + +#[test] +fn modern_reductions_preserve_noop_and_reject_dynamic_axes() { + for noop in [0, 1] { + let output_shape = if noop == 0 { vec![1, 1] } else { vec![2, 3] }; + let mut model = model("ReduceMax", 20, &[2, 3], &output_shape); + model.graph.as_mut().unwrap().node[0].attribute = + vec![integer("noop_with_empty_axes", noop)]; + let inferred = infer_static_shapes(&model, &HashMap::new()).unwrap(); + assert_eq!(inferred.value_shapes["y"], output_shape); + let graph = OnnxConverter::new(model) + .unwrap() + .convert(&ConvertOptions::default()) + .unwrap(); + assert_eq!( + graph.nodes[0].op, + if noop == 0 { "reduceMax" } else { "identity" } + ); + assert!(!graph.nodes[0].options.contains_key("axes")); + } + let mut source = model("ReduceMin", 20, &[2, 3], &[2]); + let graph = source.graph.as_mut().unwrap(); + graph.node[0].input.push("axes".to_string()); + graph.input.push(value("axes", 7, &[1])); + assert!(matches!( + OnnxConverter::new(source) + .unwrap() + .convert(&ConvertOptions::default()), + Err(OnnxError::UnsupportedOp { .. }) + )); +} + +#[test] +fn modern_reductions_accept_empty_axes_initializer() { + for optimize in [false, true] { + for noop in [0, 1] { + let output_shape = if noop == 0 { vec![1, 1] } else { vec![2, 3] }; + let mut source = model("ReduceMean", 20, &[2, 3], &output_shape); + let graph = source.graph.as_mut().unwrap(); + graph.node[0].input.push("axes".to_string()); + graph.node[0].attribute = vec![integer("noop_with_empty_axes", noop)]; + graph.initializer.push(TensorProto { + name: "axes".to_string(), + data_type: 7, + dims: vec![0], + ..Default::default() + }); + let result = OnnxConverter::new(source) + .unwrap() + .convert(&ConvertOptions { + optimize, + ..Default::default() + }) + .unwrap(); + validate_graph(&result).unwrap(); + assert_eq!( + result.nodes[0].op, + if noop == 0 { "reduceMean" } else { "identity" } + ); + assert!(!result.nodes[0].options.contains_key("axes")); + } + } +} + +#[test] +fn cast19_saturation_does_not_admit_unsupported_float8_types() { + for target in [10, 17, 18, 19, 20] { + let mut source = model("Cast", 19, &[3], &[3]); + source.graph.as_mut().unwrap().node[0].attribute = + vec![integer("to", target), integer("saturate", 0)]; + let result = OnnxConverter::new(source) + .unwrap() + .convert(&ConvertOptions::default()); + if target == 10 { + assert_eq!(result.unwrap().nodes[0].options["to"], "float16"); + } else { + assert!(matches!(result, Err(OnnxError::TypeConversion(_)))); + } + } +} + +#[test] +fn opset20_does_not_silently_accept_unimplemented_operators() { + for op in [ + "CastLike", + "QuantizeLinear", + "DequantizeLinear", + "Resize", + "GridSample", + "AffineGrid", + "DFT", + ] { + let error = OnnxConverter::new(model(op, 20, &[3], &[3])) + .unwrap() + .convert(&ConvertOptions::default()) + .unwrap_err(); + assert!( + matches!(error, OnnxError::UnsupportedOp { .. }), + "{op}: {error}" + ); + } +} + +#[test] +fn empty_axes_does_not_turn_unsupported_composite_reductions_into_identity() { + // With noop=1, ONNX still applies log/abs/square for these operators. + // Until their complete semantics are implemented, neither empty axes nor + // optimization may turn them into a successful identity conversion. + for op in [ + "ReduceLogSum", + "ReduceSumSquare", + "ReduceL1", + "ReduceL2", + "ReduceLogSumExp", + ] { + for optimize in [false, true] { + for explicit_axes in [false, true] { + let mut source = model(op, 20, &[6], &[6]); + let graph = source.graph.as_mut().unwrap(); + graph.node[0].attribute = vec![integer("noop_with_empty_axes", 1)]; + if explicit_axes { + graph.node[0].input.push("axes".to_string()); + graph.initializer.push(TensorProto { + name: "axes".to_string(), + data_type: 7, + dims: vec![0], + ..Default::default() + }); + } + assert!( + matches!( + OnnxConverter::new(source) + .unwrap() + .convert(&ConvertOptions { + optimize, + ..Default::default() + }), + Err(OnnxError::UnsupportedOp { .. }) + ), + "{op}" + ); + } + } + } +} + +#[test] +fn file_import_rejects_unsupported_types_before_outer_constant_folding() { + use webnn_graph::onnx::constant_folding::{ + evaluators::get_evaluators, fold_constants_in_model, + }; + let mut source = model("Cast", 20, &[], &[3]); + let graph = source.graph.as_mut().unwrap(); + graph.input.clear(); + graph.node[0].input = vec!["double_value".to_string()]; + graph.node[0].attribute = vec![integer("to", 1)]; + graph.node.insert( + 0, + NodeProto { + op_type: "Constant".to_string(), + output: vec!["double_value".to_string()], + attribute: vec![AttributeProto { + name: "value".to_string(), + r#type: 4, + t: Some(TensorProto { + data_type: 11, + dims: vec![3], + double_data: vec![-4.0, 0.5, 4.0], + ..Default::default() + }), + ..Default::default() + }], + ..Default::default() + }, + ); + let mut folded = source.clone(); + assert_eq!( + fold_constants_in_model(&mut folded, &get_evaluators()).unwrap(), + 2 + ); + assert!(folded.graph.unwrap().node.is_empty()); + let directory = tempfile::tempdir().unwrap(); + let path = directory.path().join("unsupported_foldable.onnx"); + std::fs::write(&path, source.encode_to_vec()).unwrap(); + for optimize in [false, true] { + assert!(matches!( + convert_onnx( + &path, + ConvertOptions { + extract_weights: false, + optimize, + ..Default::default() + } + ), + Err(OnnxError::TypeConversion(_)) + )); + } +} diff --git a/tests/support/gelu_builder.mjs b/tests/support/gelu_builder.mjs new file mode 100644 index 0000000..6cf03e5 --- /dev/null +++ b/tests/support/gelu_builder.mjs @@ -0,0 +1,45 @@ +// Small numerical builder double: exercises emitted signatures and dtype +// agreement, not a substitute for execution on a WebNN implementation. +import assert from 'node:assert/strict'; + +function roundHalf(value) { + if (!Number.isFinite(value) || value === 0) return value; + const magnitude = Math.abs(value); + const step = 2 ** Math.max(-24, Math.floor(Math.log2(magnitude)) - 10); + const scaled = magnitude / step; + const lower = Math.floor(scaled); + const rounded = scaled - lower === 0.5 ? lower + (lower % 2) : Math.round(scaled); + const result = rounded * step; + return Math.sign(value) * (result > 65504 ? Infinity : result); +} + +class MLGraphBuilder { + constructor(data) { this.data = data; } + input(name, descriptor) { + const round = descriptor.dataType === 'float16' ? roundHalf : Math.fround; + return {type: descriptor.dataType, data: this.data[name].map(round)}; + } + constant(descriptor, bytes) { + assert.equal(descriptor.dataType, 'float32'); + assert.deepEqual(descriptor.shape, []); + return {type: 'float32', data: Array.from(new Float32Array(bytes))}; + } + cast(input, type, options = {}) { + assert.ok(type === 'float16' || type === 'float32'); + assert.deepEqual(options, {}); + return {type, data: input.data.map(type === 'float16' ? roundHalf : Math.fround)}; + } + binary(a, b, operation) { + assert.equal(a.type, b.type); + assert.equal(a.type, 'float32'); + const length = Math.max(a.data.length, b.data.length); + return {type: a.type, data: Array.from({length}, (_, i) => + Math.fround(operation(a.data[a.data.length === 1 ? 0 : i], b.data[b.data.length === 1 ? 0 : i])))}; + } + mul(a, b) { return this.binary(a, b, (x, y) => x * y); } + add(a, b) { return this.binary(a, b, (x, y) => x + y); } + tanh(input) { + return {type: input.type, data: input.data.map(x => Math.fround(Math.tanh(x)))}; + } + async build(outputs) { return outputs; } +}