Coverage for  / home / jenkins / .local / lib / python3.10 / site-packages / hyper_parallel / auto_parallel / sapp_nd / nd / common / hardware.py: 98%

173 statements  

« prev     ^ index     » next       coverage.py v7.13.1, created at 2026-08-04 05:18 +0800

1# Copyright 2025-2026 Huawei Technologies Co., Ltd 

2# 

3# Licensed under the Apache License, Version 2.0 (the "License"); 

4# you may not use this file except in compliance with the License. 

5# You may obtain a copy of the License at 

6# 

7# http://www.apache.org/licenses/LICENSE-2.0 

8# 

9# Unless required by applicable law or agreed to in writing, software 

10# distributed under the License is distributed on an "AS IS" BASIS, 

11# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. 

12# See the License for the specific language governing permissions and 

13# limitations under the License. 

14# ============================================================================ 

15"""hardware abstraction""" 

16 

17from __future__ import annotations 

18 

19import hyper_parallel.auto_parallel.sapp_nd.nd.dimensions as Dim 

20from hyper_parallel.auto_parallel.sapp_nd.nd.logger import logger 

21 

22 

23class Type: 

24 """Machine type""" 

25 

26 name: str 

27 levels: int # levels in hierarchy 

28 level_bound_number: list[int] # devices per level 

29 level_bandwidth: list[int] # bandwidth (GB/s) per level 

30 

31 def __init__(self, name, bounds, bandwidths): 

32 self.name = name 

33 self.level_bound_number = bounds 

34 self.level_bandwidth = bandwidths 

35 if len(bounds) != len(bandwidths): 

36 raise ValueError("bounds and bandwidths must have the same length") 

37 self.levels = len(bounds) 

38 

39 def __str__(self): 

40 return self.name 

41 

42 def __repr__(self): 

43 return str(self) 

44 

45 def devices_below_level(self, level): 

46 """Number of devices below the given hierarchy level""" 

47 devices = 1 

48 for lvl in range(min(level, self.levels)): 

49 devices *= self.level_bound_number[lvl] 

50 return devices 

51 

52 def intra_node_num(self): 

53 """Number of devices in a node""" 

54 return self.devices_below_level(1) 

55 

56 def levels_used(self, device_number): 

57 """Number of hierarchy level used""" 

58 devices = 1 

59 for lvl in range(self.levels): 

60 if self.level_bound_number[lvl]: 

61 devices *= self.level_bound_number[lvl] 

62 if device_number <= devices: 

63 return lvl 

64 else: 

65 return lvl 

66 return self.levels 

67 

68 def level_assign(self, dp=1, tp=1, cp=1, pp=1, ep=1): 

69 """device assignment of the different parallel dimensions""" 

70 # EP borrows from DP, not counted in total devices; kept for 

71 # topology tracking only — callers should NOT pass ep > 1 

72 # unless they also account for EP-in-DP convention. 

73 device_number = dp * tp * cp * pp * ep 

74 logger.debug("DP = %d, TP = %d, EP = %d, CP = %d, PP = %d", dp, tp, ep, cp, pp) 

75 assignment = {} 

76 assignment[Dim.TP] = [] 

77 assignment[Dim.EP] = [] 

78 assignment[Dim.CP] = [] 

79 assignment[Dim.DP] = [] 

80 assignment[Dim.PP] = [] 

81 for level in range(self.levels): 

82 bound = self.level_bound_number[level] 

83 if bound: 

84 level_device_number = min(device_number, bound) 

85 device_number = device_number // bound 

86 else: 

87 level_device_number = device_number 

88 remaining_devices = max(level_device_number, 1) 

89 

90 tp_level = min(tp, remaining_devices) 

91 assignment[Dim.TP].append(tp_level) 

92 tp = tp // tp_level 

93 remaining_devices = remaining_devices // tp_level 

94 

95 ep_level = min(ep, remaining_devices) 

96 assignment[Dim.EP].append(ep_level) 

97 ep = ep // ep_level 

98 remaining_devices = remaining_devices // ep_level 

99 

100 cp_level = min(cp, remaining_devices) 

101 assignment[Dim.CP].append(cp_level) 

102 cp = cp // cp_level 

103 remaining_devices = remaining_devices // cp_level 

104 

105 dp_level = min(dp, remaining_devices) 

106 assignment[Dim.DP].append(dp_level) 

107 dp = dp // dp_level 

108 remaining_devices = remaining_devices // dp_level 

109 

110 pp_level = min(pp, remaining_devices) 

111 assignment[Dim.PP].append(pp_level) 

112 pp = pp // pp_level 

113 remaining_devices = remaining_devices // pp_level 

114 

115 return assignment 

116 

117 

118# Device_A2 = Machine(devices_per_node=8, inter_node_bw=10, intra_node_bw=50) 

119Device_A2 = Type(name="A2", bounds=[8, None], bandwidths=[50, 10]) 

120Device_A3 = Type( 

121 name="A3", bounds=[16, 24, None], bandwidths=[200, 25, 10] 

122) 

123device_map = { 

124 "A2": Device_A2, 

125 "A3": Device_A3, 

126 "V100": Type(name="V100", bounds=[8, None], bandwidths=[50, 10]), 

127} 

128 

129 

130class Machine: 

131 """Hardware description""" 

132 

133 number: int 

134 device: Type 

135 

136 def __init__(self, number, device): 

137 self.number = number 

138 if isinstance(device, int): 

139 if device == 2: 

140 self.device = Device_A2 

141 elif device == 3: 

142 self.device = Device_A3 

143 else: 

144 raise ValueError(f"Ascend A{device} unknown") 

145 elif isinstance(device, str): 

146 if device not in device_map: 

147 raise ValueError( 

148 f"Device {device} is not supported. " 

149 f"Supported devices: {list(device_map.keys())}" 

150 ) 

151 self.device = device_map[device] 

152 else: 

153 self.device = device 

154 

155 def update_num_if_none(self, num): 

156 """Assign number of device if not already precised""" 

157 if self.number is None: 

158 self.number = num 

159 

160 def pipeline_bound(self): 

161 """Return pipeline bound from hardware topology because as pipeline may currently not cross hierarchy levels""" 

162 max_bound = 1 

163 devices = self.number 

164 while devices > 1: 

165 max_bound = max( 

166 max_bound, 

167 devices 

168 // self.device.devices_below_level( 

169 self.device.levels_used(devices) 

170 ), 

171 ) 

172 devices = devices // 2 

173 # devices = self.devices_below_level(self.levels_used(device_number)) 

174 # return device_number // devices 

175 return max_bound 

176 

177 

178def prime_factors(n): 

179 """Decompose n into a product of prime factors""" 

180 divisor = 2 

181 factors = [] 

182 while n > 1: 

183 while n % divisor != 0: 

184 divisor += 1 

185 factors.append(divisor) 

186 n = n // divisor 

187 return factors 

188 

189 

190def all_factors_combinations(factors): 

191 """Computes all divisors from a prime factor list""" 

192 def rec_factors(n, factors): 

193 combinations = {n} 

194 for u in set(factors): 

195 remaining = factors.copy() 

196 remaining.remove(u) 

197 combinations = combinations.union(rec_factors(n * u, remaining)) 

198 return combinations 

199 return rec_factors(1, factors) 

200 

201 

202def all_divisors(n, reverse=False, min_bound=1, max_bound=float("inf")): 

203 """Computes all divisors of an integer n""" 

204 divisors = sorted( 

205 all_factors_combinations(prime_factors(n)), reverse=reverse 

206 ) 

207 div_in_bound = [] 

208 for d in divisors: 

209 if min_bound <= d <= max_bound: 

210 div_in_bound.append(d) 

211 

212 return div_in_bound 

213 

214 

215def from_prime_factors(factors): 

216 """Compute a number from its prime factor decomposition""" 

217 number = 1 

218 for f in factors: 

219 number *= f 

220 return number 

221 

222 

223def split_node(n, device): 

224 """Split decompositions into intra & inter devices""" 

225 devices_per_node = device.intra_node_num() 

226 nodes = prime_factors(max(1, n // devices_per_node)) 

227 intra = prime_factors(min(n, devices_per_node)) 

228 return [intra, nodes] 

229 

230 

231def unique_factors(factors): 

232 """Remove duplicates. Factors are sorted""" 

233 offset = 0 

234 for i, f in enumerate(factors[:-1]): 

235 j = i - offset 

236 if factors[j + 1] == f: 

237 factors.pop(j) 

238 offset += 1 

239 return factors 

240 

241 

242def highest_power_of_2_divisor(divisor_of): 

243 """Compute the highest number that is both a divisor of 'divisor_of' and a power of 2""" 

244 divisor = 1 

245 factors = prime_factors(divisor_of) 

246 for f in factors: 

247 if f == 2: 

248 divisor *= f 

249 return divisor 

250 

251 

252def get_cp_topology(tp_degree: int, cp_degree: int, device_per_node: int) -> tuple: 

253 """Determine CP topology and effective bandwidth. 

254 

255 Args: 

256 tp_degree: Tensor parallelism degree. 

257 cp_degree: Context parallelism degree. 

258 device_per_node: Number of devices per node. 

259 

260 Returns: 

261 Tuple of (topology_type, effective_bandwidth, is_intra_node). 

262 - topology_type: "intra-node" or "cross-node" 

263 - effective_bandwidth: Bandwidth in GB/s 

264 - is_intra_node: True if CP stays within node 

265 """ 

266 total_devices_needed = tp_degree * cp_degree 

267 

268 if total_devices_needed <= device_per_node: 

269 topology_type = "intra-node" 

270 is_intra_node = True 

271 effective_bandwidth = 300.0 

272 else: 

273 topology_type = "cross-node" 

274 is_intra_node = False 

275 effective_bandwidth = 25.0 

276 

277 return topology_type, effective_bandwidth, is_intra_node 

278 

279 

280def get_cp_bandwidth(topology_type: str, device_type: str = "A2") -> float: 

281 """Get effective bandwidth for CP communication based on topology. 

282 

283 Args: 

284 topology_type: "intra-node" or "cross-node" 

285 device_type: Device type string (e.g., "A2", "A3") 

286 

287 Returns: 

288 Bandwidth in GB/s 

289 """ 

290 device = device_map.get(device_type, Device_A2) 

291 

292 if topology_type == "intra-node": 

293 return device.level_bandwidth[0] if device.level_bandwidth else 300.0 

294 return device.level_bandwidth[1] if len(device.level_bandwidth) > 1 else 25.0 

295 

296 

297def recommend_cp_max_by_attention(attention_type: str) -> int: 

298 """Recommend maximum CP degree based on attention type. 

299 

300 Args: 

301 attention_type: "mla", "gqa", or "mha" 

302 

303 Returns: 

304 Recommended maximum CP degree 

305 """ 

306 attention_type_upper = attention_type.upper() 

307 if attention_type_upper == "MLA": 

308 return 16 

309 if attention_type_upper == "GQA": 

310 return 8 

311 return 4