5cbd3f29e3
Fuzz / Run fuzz harnesses (${{ github.event_name == 'schedule' && 'nightly' || 'smoke' }}) (push) Has been cancelled
Create Releases / call-mac (push) Has been cancelled
Create Releases / call-linux (push) Has been cancelled
Create Releases / call-sdist (push) Has been cancelled
Create Releases / call-win (push) Has been cancelled
Create Releases / call-pyodide (push) Has been cancelled
Windows_No_Exception_CI / build (x64, 3.10) (push) Has been cancelled
Check URLs / build (push) Has been cancelled
Create Releases / Attest CI build artifacts (push) Has been cancelled
Create Releases / Check for Publish release build to pypi (push) Has been cancelled
Create Releases / Check for Publish preview build to test.pypi-weekly (push) Has been cancelled
Create Releases / Publish preview build to test.pypi-weekly (push) Has been cancelled
Create Releases / Check for Publish release build to test.pypi (rc-candidates) (push) Has been cancelled
Create Releases / Publish release build to test.pypi (push) Has been cancelled
Create Releases / Check for Publish preview build to pypi-weekly (push) Has been cancelled
Create Releases / Publish preview build to pypi-weekly (push) Has been cancelled
Create Releases / Publish release build to pypi (push) Has been cancelled
Create Releases / test source distribution (push) Has been cancelled
clang-tidy / clang-tidy (push) Has been cancelled
Lint / Validate SBOM (push) Has been cancelled
Lint / Enforce style (push) Has been cancelled
CI / Test windows-2022, 3.14, External, debug=0, unity_build=0, onnx_ml=1, autogen=0 (push) Has been cancelled
CI / Test windows-latest, 3.10, Internal, debug=0, unity_build=0, onnx_ml=1, autogen=0 (push) Has been cancelled
CI / Test windows-latest, 3.14, Internal, debug=0, unity_build=0, onnx_ml=1, autogen=0 (push) Has been cancelled
CI / Test windows-latest, 3.14t, Internal, debug=0, unity_build=0, onnx_ml=1, autogen=0 (push) Has been cancelled
CI / Test ubuntu-24.04, 3.14, Internal, debug=1, unity_build=0, onnx_ml=1, autogen=0 (push) Has been cancelled
CI / Test ubuntu-24.04, 3.14, External, debug=0, unity_build=1, onnx_ml=1, autogen=1 (push) Has been cancelled
CI / Test ubuntu-24.04, 3.14, External, debug=0, unity_build=0, onnx_ml=0, autogen=0 (push) Has been cancelled
CI / Test macos-latest, 3.10, Internal, debug=0, unity_build=0, onnx_ml=1, autogen=0 (push) Has been cancelled
CI / Test macos-latest, 3.14, Internal, debug=0, unity_build=0, onnx_ml=1, autogen=0 (push) Has been cancelled
CI / Test macos-latest, 3.14t, Internal, debug=0, unity_build=0, onnx_ml=1, autogen=0 (push) Has been cancelled
CI / Test ubuntu-24.04, 3.14, External, debug=0, unity_build=0, onnx_ml=1, autogen=0 (push) Has been cancelled
CI / Test ubuntu-24.04, 3.10, Internal, debug=0, unity_build=0, onnx_ml=1, autogen=0 (push) Has been cancelled
CI / Test ubuntu-24.04, 3.14, Internal, debug=0, unity_build=0, onnx_ml=1, autogen=0 (push) Has been cancelled
CI / Test ubuntu-24.04, 3.14t, Internal, debug=0, unity_build=0, onnx_ml=1, autogen=0 (push) Has been cancelled
Pixi CI / Install and lint (ubuntu-24.04-arm) (push) Has been cancelled
Pixi CI / Install and lint (windows-2022) (push) Has been cancelled
Pixi CI / Xcode generator build (push) Has been cancelled
Pixi CI / Install and test (macos-latest, default) (push) Has been cancelled
Pixi CI / Install and test (ubuntu-24.04-arm, default) (push) Has been cancelled
Pixi CI / Install and test (ubuntu-latest, default) (push) Has been cancelled
Pixi CI / Install and test (windows-2022, default) (push) Has been cancelled
Pixi CI / Install and test (macos-latest, oldies) (push) Has been cancelled
Pixi CI / Install and test (ubuntu-24.04-arm, oldies) (push) Has been cancelled
Pixi CI / Install and test (ubuntu-latest, oldies) (push) Has been cancelled
Pixi CI / Install and test (windows-2022, oldies) (push) Has been cancelled
CodeQL / Analyze (actions) (push) Has been cancelled
CodeQL / Analyze (cpp) (push) Has been cancelled
CodeQL / Analyze (python) (push) Has been cancelled
Copilot Setup Steps / copilot-setup-steps (push) Has been cancelled
Generate and publish ONNX docs / build (push) Has been cancelled
Generate and publish ONNX docs / deploy (push) Has been cancelled
Scorecard supply-chain security / Scorecard analysis (push) Has been cancelled
111 lines
3.9 KiB
Python
111 lines
3.9 KiB
Python
# Copyright (c) ONNX Project Contributors
|
|
|
|
# SPDX-License-Identifier: Apache-2.0
|
|
from __future__ import annotations
|
|
|
|
import numpy as np
|
|
|
|
import onnx
|
|
from onnx.backend.test.case.base import Base
|
|
from onnx.backend.test.case.model import expect
|
|
from onnx.defs import AI_ONNX_PREVIEW_TRAINING_DOMAIN, ONNX_DOMAIN
|
|
|
|
|
|
class Gradient(Base):
|
|
@staticmethod
|
|
def export_gradient_scalar_add() -> None:
|
|
add_node = onnx.helper.make_node("Add", ["a", "b"], ["c"], name="my_add")
|
|
gradient_node = onnx.helper.make_node(
|
|
"Gradient",
|
|
["a", "b"],
|
|
["dc_da", "dc_db"],
|
|
name="my_gradient",
|
|
domain=AI_ONNX_PREVIEW_TRAINING_DOMAIN,
|
|
xs=["a", "b"],
|
|
y="c",
|
|
)
|
|
|
|
a = np.array(1.0).astype(np.float32)
|
|
b = np.array(2.0).astype(np.float32)
|
|
c = a + b
|
|
# dc / da = d(a+b) / da = 1
|
|
dc_da = np.array(1).astype(np.float32)
|
|
# db / db = d(a+b) / db = 1
|
|
dc_db = np.array(1).astype(np.float32)
|
|
|
|
graph = onnx.helper.make_graph(
|
|
nodes=[add_node, gradient_node],
|
|
name="GradientOfAdd",
|
|
inputs=[
|
|
onnx.helper.make_tensor_value_info("a", onnx.TensorProto.FLOAT, []),
|
|
onnx.helper.make_tensor_value_info("b", onnx.TensorProto.FLOAT, []),
|
|
],
|
|
outputs=[
|
|
onnx.helper.make_tensor_value_info("c", onnx.TensorProto.FLOAT, []),
|
|
onnx.helper.make_tensor_value_info("dc_da", onnx.TensorProto.FLOAT, []),
|
|
onnx.helper.make_tensor_value_info("dc_db", onnx.TensorProto.FLOAT, []),
|
|
],
|
|
)
|
|
opsets = [
|
|
onnx.helper.make_operatorsetid(ONNX_DOMAIN, 12),
|
|
onnx.helper.make_operatorsetid(AI_ONNX_PREVIEW_TRAINING_DOMAIN, 1),
|
|
]
|
|
model = onnx.helper.make_model_gen_version(
|
|
graph, producer_name="backend-test", opset_imports=opsets
|
|
)
|
|
expect(
|
|
model, inputs=[a, b], outputs=[c, dc_da, dc_db], name="test_gradient_of_add"
|
|
)
|
|
|
|
@staticmethod
|
|
def export_gradient_scalar_add_and_mul() -> None:
|
|
add_node = onnx.helper.make_node("Add", ["a", "b"], ["c"], name="my_add")
|
|
mul_node = onnx.helper.make_node("Mul", ["c", "a"], ["d"], name="my_mul")
|
|
gradient_node = onnx.helper.make_node(
|
|
"Gradient",
|
|
["a", "b"],
|
|
["dd_da", "dd_db"],
|
|
name="my_gradient",
|
|
domain=AI_ONNX_PREVIEW_TRAINING_DOMAIN,
|
|
xs=["a", "b"],
|
|
y="d",
|
|
)
|
|
|
|
a = np.array(1.0).astype(np.float32)
|
|
b = np.array(2.0).astype(np.float32)
|
|
c = a + b
|
|
# d = a * c = a * (a + b)
|
|
d = a * c
|
|
# dd / da = d(a*a+a*b) / da = 2 * a + b
|
|
dd_da = (2 * a + b).astype(np.float32)
|
|
# dd / db = d(a*a+a*b) / db = a
|
|
dd_db = a
|
|
|
|
graph = onnx.helper.make_graph(
|
|
nodes=[add_node, mul_node, gradient_node],
|
|
name="GradientOfTwoOperators",
|
|
inputs=[
|
|
onnx.helper.make_tensor_value_info("a", onnx.TensorProto.FLOAT, []),
|
|
onnx.helper.make_tensor_value_info("b", onnx.TensorProto.FLOAT, []),
|
|
],
|
|
outputs=[
|
|
onnx.helper.make_tensor_value_info("d", onnx.TensorProto.FLOAT, []),
|
|
onnx.helper.make_tensor_value_info("dd_da", onnx.TensorProto.FLOAT, []),
|
|
onnx.helper.make_tensor_value_info("dd_db", onnx.TensorProto.FLOAT, []),
|
|
],
|
|
)
|
|
|
|
opsets = [
|
|
onnx.helper.make_operatorsetid(ONNX_DOMAIN, 12),
|
|
onnx.helper.make_operatorsetid(AI_ONNX_PREVIEW_TRAINING_DOMAIN, 1),
|
|
]
|
|
model = onnx.helper.make_model_gen_version(
|
|
graph, producer_name="backend-test", opset_imports=opsets
|
|
)
|
|
expect(
|
|
model,
|
|
inputs=[a, b],
|
|
outputs=[d, dd_da, dd_db],
|
|
name="test_gradient_of_add_and_mul",
|
|
)
|