Files
onnx--onnx/onnx/reference/ops/aionnxml/op_tree_ensemble_regressor.py
wehub-resource-sync 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
chore: import upstream snapshot with attribution
2026-07-13 12:41:19 +08:00

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.")