Coverage for /opt/cloud/slavespace/usr1/096471637100f3de0fcfc01072822a80/dttest/api/python/ge/ge/es/tensor_holder.py: 92%
119 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-07-27 10:03 +0800
« prev ^ index » next coverage.py v7.15.2, created at 2026-07-27 10:03 +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"""TensorHolder module for tensor operations in eager-style graph construction."""
16import ctypes
17from typing import TYPE_CHECKING, List, Union
19from ge._capi.pyes_graph_builder_wrapper import (
20 DEFAULT_GENERATED_LIB_NAME,
21 MATH_LIB_NAME,
22 EsCTensorHolderPtr,
23 c_int64,
24 esb_lib,
25 get_generated_lib,
26)
27from ge.es.tensor_like import convert_to_tensor_holder
28from ge.graph.types import DataType, Format
30if TYPE_CHECKING:
31 from ge.es.graph_builder import GraphBuilder
32 from ge.es.tensor_like import TensorLike
33 from ge.graph.node import Node
36class TensorHolder:
37 """TensorHolder for tensor operations in eager-style graph construction.
39 This class provides a Pythonic interface for tensor operations using
40 the eager-style graph builder C API.
42 The TensorHolder automatically resolves and maintains a strong reference to its
43 GraphBuilder to ensure that the underlying C++ resources remain valid as long
44 as the TensorHolder object exists. This prevents dangling references when the
45 GraphBuilder is garbage collected while TensorHolder objects are still in use.
48 Note:
49 # TensorHolder cannot call setter methods after GraphBuilder.build_and_reset() is called.
51 Example:
52 >>> builder = GraphBuilder("my_graph")
53 >>> tensor1 = builder.create_const_float(1.0)
54 >>> tensor2 = builder.create_const_float(2.0)
55 >>> result = tensor1 + tensor2 # Uses operator overloading
56 >>> result = Add(tensor1, tensor2) # Or explicit method call
57 >>> # Even if 'builder' goes out of scope, 'tensor1' and 'tensor2' remain valid
58 """
60 def __init__(self):
61 """Prevent direct instantiation of TensorHolder objects."""
62 raise RuntimeError("TensorHolder objects should not be created directly")
64 @classmethod
65 def _create_from(cls, handle: EsCTensorHolderPtr, owner_builder: "GraphBuilder") -> "TensorHolder":
66 """Create TensorHolder object from C++ pointer. (internal use only by e.g GraphBuilder.create_input(),
67 do not use this method directly)
69 Args:
70 handle: C++ EsCTensorHolder object pointer.
71 owner_builder: The GraphBuilder that owns this tensor.
73 Returns:
74 TensorHolder object.
76 Raises:
77 ValueError: If handle or owner_builder is None.
78 """
79 if not handle:
80 raise ValueError("Tensor handle cannot be None")
81 if not owner_builder:
82 raise ValueError("Owner builder cannot be None")
84 instance = cls.__new__(cls)
85 instance._handle = handle
86 instance._builder = owner_builder
88 return instance
90 def _check_usable(self, operation: str) -> None:
91 """Check if tensor holder is usable.
93 Args:
94 operation: Operation name.
95 """
96 self._builder._check_usable(operation)
98 def _validate_operation(self, other: "TensorHolder", op_name: str) -> None:
99 """Validate operation between two tensor holders.
101 Args:
102 other: Another TensorHolder object.
103 op_name: Operation name.
104 """
105 if self._builder is not other._builder:
106 raise ValueError(
107 f"Cannot perform {op_name}: tensors from different GraphBuilders "
108 f"('{self._builder.name}' vs '{other._builder.name}')"
109 )
110 other._check_usable(op_name)
111 self._check_usable(op_name)
113 @property
114 def name(self) -> str:
115 """Get node name.
117 Returns:
118 Producer node name.
119 """
120 return self._get_node_snapshot().name
122 def set_data_type(self, data_type: DataType) -> "TensorHolder":
123 """Set tensor data type.
125 Args:
126 data_type: Data type using DataType enum.
128 Returns:
129 TensorHolder object.
131 Raises:
132 TypeError: If data_type is not a DataType enum.
133 RuntimeError: If operation fails.
134 """
135 self._check_usable("set data type")
137 if not isinstance(data_type, DataType):
138 raise TypeError("Data type must be a DataType enum")
140 if esb_lib.EsSetDataType(self._handle, ctypes.c_int(data_type.value)) != 0:
141 raise RuntimeError("Failed to set data type")
143 return self
145 def set_format(self, format: Format) -> "TensorHolder":
146 """Set tensor data format.
148 Args:
149 format: Data format using Format enum.
151 Returns:
152 TensorHolder object.
154 Raises:
155 TypeError: If format is not a Format enum.
156 RuntimeError: If operation fails.
157 """
158 self._check_usable("set format")
160 if not isinstance(format, Format):
161 raise TypeError("Format must be a Format enum")
163 if esb_lib.EsSetFormat(self._handle, ctypes.c_int(format.value)) != 0:
164 raise RuntimeError("Failed to set format")
166 return self
168 def set_shape(self, shape: List[int]) -> "TensorHolder":
169 """Set tensor shape.
171 Args:
172 shape: List of shape dimensions.
174 Returns:
175 TensorHolder object.
177 Raises:
178 TypeError: If shape is not a list of integers.
179 RuntimeError: If operation fails.
180 """
181 self._check_usable("set shape")
183 if not isinstance(shape, list):
184 raise TypeError("Shape must be a list of integers")
186 if not all(isinstance(dim, int) for dim in shape):
187 raise TypeError("All shape dimensions must be integers")
189 dim_num = len(shape)
190 shape_array = (c_int64 * dim_num)(*shape)
191 if esb_lib.EsSetShape(self._handle, shape_array, c_int64(dim_num)) != 0:
192 raise RuntimeError("Failed to set shape")
194 return self
196 def _get_math_operator_lib(self):
197 """Get library containing math operators (Add/Sub/Mul/Div).
199 Tries libes_math.so first, falls back to default library if unavailable.
201 Returns:
202 ctypes library object containing EsAdd/EsSub/EsMul/EsDiv.
204 Raises:
205 RuntimeError: If neither library is available.
206 """
207 try:
208 return get_generated_lib()
209 except RuntimeError:
210 try:
211 return get_generated_lib(MATH_LIB_NAME)
212 except RuntimeError as exc:
213 raise RuntimeError(
214 f"Math operators (Add/Sub/Mul/Div) not available: neither {MATH_LIB_NAME} "
215 f"nor {DEFAULT_GENERATED_LIB_NAME} could be loaded. Please ensure at least one is accessible."
216 ) from exc
218 # Operator overloading support
219 def __add__(self, other: Union["TensorHolder", "TensorLike"]) -> "TensorHolder":
220 """Support + operator."""
221 return self.add(other)
223 def __sub__(self, other: Union["TensorHolder", "TensorLike"]) -> "TensorHolder":
224 """Support - operator."""
225 return self.sub(other)
227 def __mul__(self, other: Union["TensorHolder", "TensorLike"]) -> "TensorHolder":
228 """Support * operator."""
229 return self.mul(other)
231 def __truediv__(self, other: Union["TensorHolder", "TensorLike"]) -> "TensorHolder":
232 """Support / operator."""
233 return self.div(other)
235 def __radd__(self, other: Union["TensorHolder", "TensorLike"]) -> "TensorHolder":
236 """Support right addition operator."""
237 return self.add(other)
239 def __rsub__(self, other: Union["TensorHolder", "TensorLike"]) -> "TensorHolder":
240 """Support right subtraction operator."""
241 other = convert_to_tensor_holder(other, self._builder)
242 return other.sub(self)
244 def __rmul__(self, other: Union["TensorHolder", "TensorLike"]) -> "TensorHolder":
245 """Support right multiplication operator."""
246 return self.mul(other)
248 def __rtruediv__(self, other: Union["TensorHolder", "TensorLike"]) -> "TensorHolder":
249 """Support right division operator."""
250 other = convert_to_tensor_holder(other, self._builder)
251 return other.div(self)
253 def add(self, other: Union["TensorHolder", "TensorLike"]) -> "TensorHolder":
254 """Add two tensors.
256 Args:
257 other: Another TensorHolder object.
259 Returns:
260 New TensorHolder representing the result.
262 Raises:
263 TypeError: If other is not a TensorHolder.
264 RuntimeError: If operation fails or library is not available.
265 """
266 other = convert_to_tensor_holder(other, self._builder)
268 self._validate_operation(other, "add")
269 generated_lib = self._get_math_operator_lib()
271 result_handle = generated_lib.EsAdd(self._handle, other._handle)
272 if not result_handle:
273 raise RuntimeError("Failed to create add operation")
275 return self._builder._apply_scope_infos_to_node(TensorHolder._create_from(result_handle, self._builder))
277 def sub(self, other: Union["TensorHolder", "TensorLike"]) -> "TensorHolder":
278 """Subtract two tensors.
280 Args:
281 other: Another TensorHolder object.
283 Returns:
284 New TensorHolder representing the result.
286 Raises:
287 TypeError: If other is not a TensorHolder.
288 RuntimeError: If operation fails or library is not available.
289 """
290 other = convert_to_tensor_holder(other, self._builder)
292 self._validate_operation(other, "sub")
293 generated_lib = self._get_math_operator_lib()
295 result_handle = generated_lib.EsSub(self._handle, other._handle)
296 if not result_handle:
297 raise RuntimeError("Failed to create sub operation")
299 return self._builder._apply_scope_infos_to_node(TensorHolder._create_from(result_handle, self._builder))
301 def mul(self, other: Union["TensorHolder", "TensorLike"]) -> "TensorHolder":
302 """Multiply two tensors.
304 Args:
305 other: Another TensorHolder object.
307 Returns:
308 New TensorHolder representing the result.
310 Raises:
311 TypeError: If other is not a TensorHolder.
312 RuntimeError: If operation fails or library is not available.
313 """
314 other = convert_to_tensor_holder(other, self._builder)
316 self._validate_operation(other, "mul")
317 generated_lib = self._get_math_operator_lib()
319 result_handle = generated_lib.EsMul(self._handle, other._handle)
320 if not result_handle:
321 raise RuntimeError("Failed to create mul operation")
323 return self._builder._apply_scope_infos_to_node(TensorHolder._create_from(result_handle, self._builder))
325 def div(self, other: Union["TensorHolder", "TensorLike"]) -> "TensorHolder":
326 """Divide two tensors.
328 Args:
329 other: Another TensorHolder object.
331 Returns:
332 New TensorHolder representing the result.
334 Raises:
335 TypeError: If other is not a TensorHolder.
336 RuntimeError: If operation fails or library is not available.
337 """
338 other = convert_to_tensor_holder(other, self._builder)
340 self._validate_operation(other, "div")
341 generated_lib = self._get_math_operator_lib()
343 result_handle = generated_lib.EsDiv(self._handle, other._handle)
344 if not result_handle:
345 raise RuntimeError("Failed to create div operation")
347 return self._builder._apply_scope_infos_to_node(TensorHolder._create_from(result_handle, self._builder))
349 def get_owner_builder(self):
350 return self._builder
352 def _get_node_snapshot(self) -> "Node":
353 """Get a snapshot of the producer node of this tensor. (internal use by internal api)
355 Returns:
356 Node object representing the producer node snapshot.
357 The returned Node object is a snapshot and does not own the underlying pointer.
359 Raises:
360 RuntimeError: If getting node snapshot fails.
361 """
362 from ge.graph.node import Node
364 node_ptr = esb_lib.EsGetProducer(self._handle)
365 if not node_ptr:
366 raise RuntimeError("Failed to get producer node")
368 return Node._create_from(ctypes.c_void_p(node_ptr), owns_handle=False)