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

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"""Handling tensor-like operands in eager-style graph construction.""" 

15 

16from __future__ import annotations 

17 

18from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union 

19 

20if TYPE_CHECKING: 

21 from ge.es import GraphBuilder, TensorHolder 

22 

23Number = Union[int, float] 

24TensorLike = Union[Number, List["TensorLike"]] 

25 

26 

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 

41 

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) 

46 

47 if isinstance(value, list): 

48 return _convert_list_to_tensor_holder(value, owner_builder) 

49 

50 raise TypeError("Value must be TensorHolder or int|float or nested lists of int|float or None") 

51 

52 

53def resolve_builder( 

54 *values: Union["TensorLike", "TensorHolder", "GraphBuilder"], 

55) -> "GraphBuilder": 

56 """ 

57 Resolve the owning GraphBuilder from inputs. 

58 

59 Args: 

60 *values: TensorHolder or TensorLike or GraphBuilder arguments. 

61 

62 Returns: 

63 The owner GraphBuilder. 

64 

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

70 

71 from ge.es import GraphBuilder, TensorHolder 

72 

73 for value in values: 

74 if isinstance(value, GraphBuilder): 

75 return value 

76 if isinstance(value, TensorHolder): 

77 return value.get_owner_builder() 

78 

79 raise ValueError("Please ensure at least one input tensor or an explicit owner_builder is provided when supported") 

80 

81 

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. 

87 

88 Args: 

89 value: List whose leaves are ints/floats. 

90 

91 Returns: 

92 A pair of (flat_values, shape) (e.g., [[1, 2], [3, 4]] -> ([1, 2, 3, 4], [2, 2])). 

93 

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] 

100 

101 def _recurse(current: "TensorLike") -> Tuple[List["Number"], List[int]]: 

102 if isinstance(current, (int, float)): 

103 return [current], [] 

104 

105 if not isinstance(current, list) or not current: 

106 raise TypeError("Value must be int|float or nested lists of int|float") 

107 

108 first_flat, inner_shape = _recurse(current[0]) 

109 flat_acc: List[Number] = list(first_flat) 

110 

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) 

116 

117 return flat_acc, [len(current)] + inner_shape 

118 

119 return _recurse(value) 

120 

121 

122def _unflatten(value: List["Number"], shape: List[int]) -> List["TensorLike"]: 

123 """ 

124 Unflatten data 

125 

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 

131 

132 Raises: 

133 ValueError: If shape is empty 

134 ValueError: If shape and value does not match 

135 """ 

136 

137 # shape check 

138 if not shape: 

139 raise ValueError("Shape cannot be empty") 

140 

141 # length check 

142 total_size = 1 

143 for dim in shape: 

144 total_size *= dim 

145 

146 if len(value) != total_size: 

147 raise ValueError(f"Length of list {len(value)} does not match shape {total_size}") 

148 

149 # recurse 

150 def _recurse(flat_values, shape): 

151 # last dimension 

152 if len(shape) == 1: 

153 return flat_values[: shape[0]] 

154 

155 size = shape[0] 

156 step = int(len(flat_values) / size) 

157 result = [] 

158 index = 0 

159 

160 for _ in range(size): 

161 block = flat_values[index : index + step] 

162 result.append(_recurse(block, shape[1:])) 

163 index += step 

164 

165 return result 

166 

167 return _recurse(value, shape) 

168 

169 

170def _convert_scalar_to_tensor_holder(value: "Number", owner_builder: "GraphBuilder") -> "TensorHolder": 

171 """ 

172 Convert scalar to TensorHolder object. 

173 

174 Args: 

175 value: int or float scaler 

176 owner_builder: The GraphBuilder that owns this value. 

177 

178 Returns: 

179 TensorHolder representing the scalar. 

180 

181 Raises: 

182 TypeError: If value is not a number. 

183 """ 

184 if isinstance(value, int): 

185 return owner_builder.create_scalar_int64(value) 

186 

187 if isinstance(value, float): 

188 return owner_builder.create_scalar_float(value) 

189 

190 raise TypeError("Value must be a number") 

191 

192 

193def _convert_list_to_tensor_holder(value: List[Any], owner_builder: "GraphBuilder") -> "TensorHolder": 

194 """ 

195 Convert list to TensorHolder object. 

196 

197 Args: 

198 value: nested lists of int/float 

199 owner_builder: The GraphBuilder that owns this value. 

200 

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)