Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand Down
22 changes: 20 additions & 2 deletions docs/onnx-lowering.md
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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.
Expand Down
14 changes: 12 additions & 2 deletions src/emit_js.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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!(
Expand Down Expand Up @@ -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]
Expand Down
116 changes: 108 additions & 8 deletions src/onnx/convert.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand All @@ -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),

Expand Down Expand Up @@ -74,6 +82,61 @@ pub(crate) fn map_onnx_data_type(onnx_type: i32) -> Result<DataType, OnnxError>
})
}

/// 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,
Expand Down Expand Up @@ -191,18 +254,30 @@ 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<i64> = node
let mut axes: Vec<i64> = node
.attribute
.as_slice()
.iter()
.find(|a| a.name.as_str() == "axes")
.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 {
Expand Down Expand Up @@ -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<String> = 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<String, String> = HashMap::new();
let mut effective_overrides = options.free_dim_overrides.clone();
let mut inference_overrides = effective_overrides.clone();
Expand Down Expand Up @@ -1569,10 +1670,6 @@ Provide --override-dim {}=<value> or enable --experimental-dynamic-inputs.",
)));
};

if shape.is_empty() {
continue;
}

self.graph.inputs.insert(
name.clone(),
crate::ast::OperandDesc {
Expand Down Expand Up @@ -2786,7 +2883,8 @@ Provide --override-dim {}=<value> 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 {
Expand Down Expand Up @@ -2895,6 +2993,8 @@ pub fn convert_onnx<P: AsRef<Path>>(
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...");
Expand Down
Loading