Coverage for /opt/cloud/slavespace/usr1/096471637100f3de0fcfc01072822a80/dttest/api/python/ge/ge/es/tensor_like.py: 97%
73 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-07-27 10:02 +0800
« prev ^ index » next coverage.py v7.15.2, created at 2026-07-27 10:02 +0800
1#!/usr/bin/env python3
2# -*- coding: utf-8 -*-
3# -------------------------------------------------------------------
4# -----------------------------------------------------------------------------------------------------------
5# Copyright (c) 2025 Huawei Technologies Co., Ltd.
6# This program is free software, you can redistribute it and/or modify it under the terms and conditions of
7# CANN Open Software License Agreement Version 2.0 (the "License").
8# Please refer to the License for details. You may not use this file except in compliance with the License.
9# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
10# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
11# See LICENSE in the root of the software repository for the full text of the License.
12# -----------------------------------------------------------------------------------------------------------
14"""Handling tensor-like operands in eager-style graph construction."""
16from __future__ import annotations
18from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union
20if TYPE_CHECKING:
21 from ge.es import GraphBuilder, TensorHolder
23Number = Union[int, float]
24TensorLike = Union[Number, List["TensorLike"]]
27def convert_to_tensor_holder(
28 value: Optional[Union["TensorHolder", "TensorLike"]], owner_builder: "GraphBuilder"
29) -> Optional["TensorHolder"]:
30 """
31 Convert value to TensorHolder object.
32 Args:
33 value: TensorHolder or int|float or nested lists of int|float or None
34 owner_builder: The GraphBuilder that owns this value.
35 Returns:
36 TensorHolder object or None if value is None
37 Raises:
38 TypeError: If value is not TensorHolder or TensorLike or None.
39 """
40 from ge.es.tensor_holder import TensorHolder
42 if value is None or isinstance(value, TensorHolder):
43 return value
44 if isinstance(value, (int, float)):
45 return _convert_scalar_to_tensor_holder(value, owner_builder)
47 if isinstance(value, list):
48 return _convert_list_to_tensor_holder(value, owner_builder)
50 raise TypeError("Value must be TensorHolder or int|float or nested lists of int|float or None")
53def resolve_builder(
54 *values: Union["TensorLike", "TensorHolder", "GraphBuilder"],
55) -> "GraphBuilder":
56 """
57 Resolve the owning GraphBuilder from inputs.
59 Args:
60 *values: TensorHolder or TensorLike or GraphBuilder arguments.
62 Returns:
63 The owner GraphBuilder.
65 Raises:
66 ValueError: If no TensorHolder is provided to infer the GraphBuilder.
67 """
68 if not values:
69 raise ValueError("At least one argument is required to resolve GraphBuilder")
71 from ge.es import GraphBuilder, TensorHolder
73 for value in values:
74 if isinstance(value, GraphBuilder):
75 return value
76 if isinstance(value, TensorHolder):
77 return value.get_owner_builder()
79 raise ValueError("Please ensure at least one input tensor or an explicit owner_builder is provided when supported")
82def _flatten_and_infer_shape(
83 value: List["TensorLike"],
84) -> Tuple[List["Number"], List[int]]:
85 """
86 Flatten nested numeric lists and infer their tensor shape.
88 Args:
89 value: List whose leaves are ints/floats.
91 Returns:
92 A pair of (flat_values, shape) (e.g., [[1, 2], [3, 4]] -> ([1, 2, 3, 4], [2, 2])).
94 Raises:
95 TypeError: value or any nested element is neither list nor numeric.
96 ValueError: Nested shapes differ.
97 """
98 if not value:
99 return value, [0]
101 def _recurse(current: "TensorLike") -> Tuple[List["Number"], List[int]]:
102 if isinstance(current, (int, float)):
103 return [current], []
105 if not isinstance(current, list) or not current:
106 raise TypeError("Value must be int|float or nested lists of int|float")
108 first_flat, inner_shape = _recurse(current[0])
109 flat_acc: List[Number] = list(first_flat)
111 for elem in current[1:]:
112 sub_flat, sub_shape = _recurse(elem)
113 if sub_shape != inner_shape:
114 raise ValueError("Irregular nested list: all sub-lists must have the same shape")
115 flat_acc.extend(sub_flat)
117 return flat_acc, [len(current)] + inner_shape
119 return _recurse(value)
122def _unflatten(value: List["Number"], shape: List[int]) -> List["TensorLike"]:
123 """
124 Unflatten data
126 Args:
127 value: flattened data like [1,2,3,4,5,6]
128 shape: list of shape like [2,3]
129 Returns:
130 List of TensorLike
132 Raises:
133 ValueError: If shape is empty
134 ValueError: If shape and value does not match
135 """
137 # shape check
138 if not shape:
139 raise ValueError("Shape cannot be empty")
141 # length check
142 total_size = 1
143 for dim in shape:
144 total_size *= dim
146 if len(value) != total_size:
147 raise ValueError(f"Length of list {len(value)} does not match shape {total_size}")
149 # recurse
150 def _recurse(flat_values, shape):
151 # last dimension
152 if len(shape) == 1:
153 return flat_values[: shape[0]]
155 size = shape[0]
156 step = int(len(flat_values) / size)
157 result = []
158 index = 0
160 for _ in range(size):
161 block = flat_values[index : index + step]
162 result.append(_recurse(block, shape[1:]))
163 index += step
165 return result
167 return _recurse(value, shape)
170def _convert_scalar_to_tensor_holder(value: "Number", owner_builder: "GraphBuilder") -> "TensorHolder":
171 """
172 Convert scalar to TensorHolder object.
174 Args:
175 value: int or float scaler
176 owner_builder: The GraphBuilder that owns this value.
178 Returns:
179 TensorHolder representing the scalar.
181 Raises:
182 TypeError: If value is not a number.
183 """
184 if isinstance(value, int):
185 return owner_builder.create_scalar_int64(value)
187 if isinstance(value, float):
188 return owner_builder.create_scalar_float(value)
190 raise TypeError("Value must be a number")
193def _convert_list_to_tensor_holder(value: List[Any], owner_builder: "GraphBuilder") -> "TensorHolder":
194 """
195 Convert list to TensorHolder object.
197 Args:
198 value: nested lists of int/float
199 owner_builder: The GraphBuilder that owns this value.
201 Returns:
202 TensorHolder representing the list.
203 """
204 flat_values, shape = _flatten_and_infer_shape(value)
205 has_float = any(isinstance(v, float) for v in flat_values)
206 if has_float:
207 return owner_builder.create_const_float(flat_values, shape=shape)
208 return owner_builder.create_const_int64(flat_values, shape=shape)