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: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# ----------------------------------------------------------------------------------------------------------- 

13 

14"""TensorHolder module for tensor operations in eager-style graph construction.""" 

15 

16import ctypes 

17from typing import TYPE_CHECKING, List, Union 

18 

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 

29 

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 

34 

35 

36class TensorHolder: 

37 """TensorHolder for tensor operations in eager-style graph construction. 

38 

39 This class provides a Pythonic interface for tensor operations using 

40 the eager-style graph builder C API. 

41 

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. 

46 

47 

48 Note: 

49 # TensorHolder cannot call setter methods after GraphBuilder.build_and_reset() is called. 

50 

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

59 

60 def __init__(self): 

61 """Prevent direct instantiation of TensorHolder objects.""" 

62 raise RuntimeError("TensorHolder objects should not be created directly") 

63 

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) 

68 

69 Args: 

70 handle: C++ EsCTensorHolder object pointer. 

71 owner_builder: The GraphBuilder that owns this tensor. 

72 

73 Returns: 

74 TensorHolder object. 

75 

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

83 

84 instance = cls.__new__(cls) 

85 instance._handle = handle 

86 instance._builder = owner_builder 

87 

88 return instance 

89 

90 def _check_usable(self, operation: str) -> None: 

91 """Check if tensor holder is usable. 

92 

93 Args: 

94 operation: Operation name. 

95 """ 

96 self._builder._check_usable(operation) 

97 

98 def _validate_operation(self, other: "TensorHolder", op_name: str) -> None: 

99 """Validate operation between two tensor holders. 

100 

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) 

112 

113 @property 

114 def name(self) -> str: 

115 """Get node name. 

116 

117 Returns: 

118 Producer node name. 

119 """ 

120 return self._get_node_snapshot().name 

121 

122 def set_data_type(self, data_type: DataType) -> "TensorHolder": 

123 """Set tensor data type. 

124 

125 Args: 

126 data_type: Data type using DataType enum. 

127 

128 Returns: 

129 TensorHolder object. 

130 

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

136 

137 if not isinstance(data_type, DataType): 

138 raise TypeError("Data type must be a DataType enum") 

139 

140 if esb_lib.EsSetDataType(self._handle, ctypes.c_int(data_type.value)) != 0: 

141 raise RuntimeError("Failed to set data type") 

142 

143 return self 

144 

145 def set_format(self, format: Format) -> "TensorHolder": 

146 """Set tensor data format. 

147 

148 Args: 

149 format: Data format using Format enum. 

150 

151 Returns: 

152 TensorHolder object. 

153 

154 Raises: 

155 TypeError: If format is not a Format enum. 

156 RuntimeError: If operation fails. 

157 """ 

158 self._check_usable("set format") 

159 

160 if not isinstance(format, Format): 

161 raise TypeError("Format must be a Format enum") 

162 

163 if esb_lib.EsSetFormat(self._handle, ctypes.c_int(format.value)) != 0: 

164 raise RuntimeError("Failed to set format") 

165 

166 return self 

167 

168 def set_shape(self, shape: List[int]) -> "TensorHolder": 

169 """Set tensor shape. 

170 

171 Args: 

172 shape: List of shape dimensions. 

173 

174 Returns: 

175 TensorHolder object. 

176 

177 Raises: 

178 TypeError: If shape is not a list of integers. 

179 RuntimeError: If operation fails. 

180 """ 

181 self._check_usable("set shape") 

182 

183 if not isinstance(shape, list): 

184 raise TypeError("Shape must be a list of integers") 

185 

186 if not all(isinstance(dim, int) for dim in shape): 

187 raise TypeError("All shape dimensions must be integers") 

188 

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

193 

194 return self 

195 

196 def _get_math_operator_lib(self): 

197 """Get library containing math operators (Add/Sub/Mul/Div). 

198 

199 Tries libes_math.so first, falls back to default library if unavailable. 

200 

201 Returns: 

202 ctypes library object containing EsAdd/EsSub/EsMul/EsDiv. 

203 

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 

217 

218 # Operator overloading support 

219 def __add__(self, other: Union["TensorHolder", "TensorLike"]) -> "TensorHolder": 

220 """Support + operator.""" 

221 return self.add(other) 

222 

223 def __sub__(self, other: Union["TensorHolder", "TensorLike"]) -> "TensorHolder": 

224 """Support - operator.""" 

225 return self.sub(other) 

226 

227 def __mul__(self, other: Union["TensorHolder", "TensorLike"]) -> "TensorHolder": 

228 """Support * operator.""" 

229 return self.mul(other) 

230 

231 def __truediv__(self, other: Union["TensorHolder", "TensorLike"]) -> "TensorHolder": 

232 """Support / operator.""" 

233 return self.div(other) 

234 

235 def __radd__(self, other: Union["TensorHolder", "TensorLike"]) -> "TensorHolder": 

236 """Support right addition operator.""" 

237 return self.add(other) 

238 

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) 

243 

244 def __rmul__(self, other: Union["TensorHolder", "TensorLike"]) -> "TensorHolder": 

245 """Support right multiplication operator.""" 

246 return self.mul(other) 

247 

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) 

252 

253 def add(self, other: Union["TensorHolder", "TensorLike"]) -> "TensorHolder": 

254 """Add two tensors. 

255 

256 Args: 

257 other: Another TensorHolder object. 

258 

259 Returns: 

260 New TensorHolder representing the result. 

261 

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) 

267 

268 self._validate_operation(other, "add") 

269 generated_lib = self._get_math_operator_lib() 

270 

271 result_handle = generated_lib.EsAdd(self._handle, other._handle) 

272 if not result_handle: 

273 raise RuntimeError("Failed to create add operation") 

274 

275 return self._builder._apply_scope_infos_to_node(TensorHolder._create_from(result_handle, self._builder)) 

276 

277 def sub(self, other: Union["TensorHolder", "TensorLike"]) -> "TensorHolder": 

278 """Subtract two tensors. 

279 

280 Args: 

281 other: Another TensorHolder object. 

282 

283 Returns: 

284 New TensorHolder representing the result. 

285 

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) 

291 

292 self._validate_operation(other, "sub") 

293 generated_lib = self._get_math_operator_lib() 

294 

295 result_handle = generated_lib.EsSub(self._handle, other._handle) 

296 if not result_handle: 

297 raise RuntimeError("Failed to create sub operation") 

298 

299 return self._builder._apply_scope_infos_to_node(TensorHolder._create_from(result_handle, self._builder)) 

300 

301 def mul(self, other: Union["TensorHolder", "TensorLike"]) -> "TensorHolder": 

302 """Multiply two tensors. 

303 

304 Args: 

305 other: Another TensorHolder object. 

306 

307 Returns: 

308 New TensorHolder representing the result. 

309 

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) 

315 

316 self._validate_operation(other, "mul") 

317 generated_lib = self._get_math_operator_lib() 

318 

319 result_handle = generated_lib.EsMul(self._handle, other._handle) 

320 if not result_handle: 

321 raise RuntimeError("Failed to create mul operation") 

322 

323 return self._builder._apply_scope_infos_to_node(TensorHolder._create_from(result_handle, self._builder)) 

324 

325 def div(self, other: Union["TensorHolder", "TensorLike"]) -> "TensorHolder": 

326 """Divide two tensors. 

327 

328 Args: 

329 other: Another TensorHolder object. 

330 

331 Returns: 

332 New TensorHolder representing the result. 

333 

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) 

339 

340 self._validate_operation(other, "div") 

341 generated_lib = self._get_math_operator_lib() 

342 

343 result_handle = generated_lib.EsDiv(self._handle, other._handle) 

344 if not result_handle: 

345 raise RuntimeError("Failed to create div operation") 

346 

347 return self._builder._apply_scope_infos_to_node(TensorHolder._create_from(result_handle, self._builder)) 

348 

349 def get_owner_builder(self): 

350 return self._builder 

351 

352 def _get_node_snapshot(self) -> "Node": 

353 """Get a snapshot of the producer node of this tensor. (internal use by internal api) 

354 

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. 

358 

359 Raises: 

360 RuntimeError: If getting node snapshot fails. 

361 """ 

362 from ge.graph.node import Node 

363 

364 node_ptr = esb_lib.EsGetProducer(self._handle) 

365 if not node_ptr: 

366 raise RuntimeError("Failed to get producer node") 

367 

368 return Node._create_from(ctypes.c_void_p(node_ptr), owns_handle=False)