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
110 lines
4.0 KiB
Python
110 lines
4.0 KiB
Python
# Copyright (c) ONNX Project Contributors
|
|
|
|
# SPDX-License-Identifier: Apache-2.0
|
|
from __future__ import annotations
|
|
|
|
import numpy as np
|
|
|
|
from onnx.reference.ops.aionnxml._op_run_aionnxml import OpRunAiOnnxMl
|
|
from onnx.reference.ops.aionnxml.op_tree_ensemble_helper import TreeEnsemble
|
|
|
|
|
|
class TreeEnsembleRegressor(OpRunAiOnnxMl):
|
|
"""`nodes_hitrates` and `nodes_hitrates_as_tensor` are not used."""
|
|
|
|
def _run(
|
|
self,
|
|
X,
|
|
aggregate_function=None,
|
|
base_values=None,
|
|
base_values_as_tensor=None,
|
|
n_targets=None,
|
|
nodes_falsenodeids=None,
|
|
nodes_featureids=None,
|
|
nodes_hitrates=None,
|
|
nodes_hitrates_as_tensor=None,
|
|
nodes_missing_value_tracks_true=None,
|
|
nodes_modes=None,
|
|
nodes_nodeids=None,
|
|
nodes_treeids=None,
|
|
nodes_truenodeids=None,
|
|
nodes_values=None,
|
|
nodes_values_as_tensor=None,
|
|
post_transform=None,
|
|
target_ids=None,
|
|
target_nodeids=None,
|
|
target_treeids=None,
|
|
target_weights=None,
|
|
target_weights_as_tensor=None,
|
|
):
|
|
nmv = nodes_missing_value_tracks_true
|
|
tr = TreeEnsemble(
|
|
base_values=base_values,
|
|
base_values_as_tensor=base_values_as_tensor,
|
|
nodes_falsenodeids=nodes_falsenodeids,
|
|
nodes_featureids=nodes_featureids,
|
|
nodes_hitrates=nodes_hitrates,
|
|
nodes_hitrates_as_tensor=nodes_hitrates_as_tensor,
|
|
nodes_missing_value_tracks_true=nmv,
|
|
nodes_modes=nodes_modes,
|
|
nodes_nodeids=nodes_nodeids,
|
|
nodes_treeids=nodes_treeids,
|
|
nodes_truenodeids=nodes_truenodeids,
|
|
nodes_values=nodes_values,
|
|
nodes_values_as_tensor=nodes_values_as_tensor,
|
|
target_weights=target_weights,
|
|
target_weights_as_tensor=target_weights_as_tensor,
|
|
)
|
|
# unused unless for debugging purposes
|
|
self._tree = tr
|
|
leaves_index = tr.leave_index_tree(X)
|
|
res = np.zeros((leaves_index.shape[0], n_targets), dtype=X.dtype)
|
|
n_trees = len(set(tr.atts.nodes_treeids))
|
|
|
|
target_index = {}
|
|
for i, (tid, nid) in enumerate(
|
|
zip(target_treeids, target_nodeids, strict=False)
|
|
):
|
|
if (tid, nid) not in target_index:
|
|
target_index[tid, nid] = []
|
|
target_index[tid, nid].append(i)
|
|
for i in range(res.shape[0]):
|
|
indices = leaves_index[i]
|
|
t_index = [
|
|
target_index[nodes_treeids[i], nodes_nodeids[i]] for i in indices
|
|
]
|
|
if aggregate_function in ("SUM", "AVERAGE"):
|
|
for its in t_index:
|
|
for it in its:
|
|
res[i, target_ids[it]] += tr.atts.target_weights[it]
|
|
elif aggregate_function == "MIN":
|
|
res[i, :] = np.finfo(res.dtype).max
|
|
for its in t_index:
|
|
for it in its:
|
|
res[i, target_ids[it]] = min(
|
|
res[i, target_ids[it]],
|
|
tr.atts.target_weights[it],
|
|
)
|
|
elif aggregate_function == "MAX":
|
|
res[i, :] = np.finfo(res.dtype).min
|
|
for its in t_index:
|
|
for it in its:
|
|
res[i, target_ids[it]] = max(
|
|
res[i, target_ids[it]],
|
|
tr.atts.target_weights[it],
|
|
)
|
|
else:
|
|
raise NotImplementedError(
|
|
f"aggregate_transform={aggregate_function!r} not supported yet."
|
|
)
|
|
if aggregate_function == "AVERAGE":
|
|
res /= n_trees
|
|
|
|
# Convention is to add base_values after aggregate function
|
|
if base_values is not None:
|
|
res[:, :] += np.array(base_values).reshape((1, -1))
|
|
|
|
if post_transform in (None, "NONE"):
|
|
return (res,)
|
|
raise NotImplementedError(f"post_transform={post_transform!r} not implemented.")
|