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
« 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"""
17from __future__ import annotations
19import hyper_parallel.auto_parallel.sapp_nd.nd.dimensions as Dim
20from hyper_parallel.auto_parallel.sapp_nd.nd.logger import logger
23class Type:
24 """Machine type"""
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
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)
39 def __str__(self):
40 return self.name
42 def __repr__(self):
43 return str(self)
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
52 def intra_node_num(self):
53 """Number of devices in a node"""
54 return self.devices_below_level(1)
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
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)
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
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
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
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
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
115 return assignment
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}
130class Machine:
131 """Hardware description"""
133 number: int
134 device: Type
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
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
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
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
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)
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)
212 return div_in_bound
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
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]
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
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
252def get_cp_topology(tp_degree: int, cp_degree: int, device_per_node: int) -> tuple:
253 """Determine CP topology and effective bandwidth.
255 Args:
256 tp_degree: Tensor parallelism degree.
257 cp_degree: Context parallelism degree.
258 device_per_node: Number of devices per node.
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
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
277 return topology_type, effective_bandwidth, is_intra_node
280def get_cp_bandwidth(topology_type: str, device_type: str = "A2") -> float:
281 """Get effective bandwidth for CP communication based on topology.
283 Args:
284 topology_type: "intra-node" or "cross-node"
285 device_type: Device type string (e.g., "A2", "A3")
287 Returns:
288 Bandwidth in GB/s
289 """
290 device = device_map.get(device_type, Device_A2)
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
297def recommend_cp_max_by_attention(attention_type: str) -> int:
298 """Recommend maximum CP degree based on attention type.
300 Args:
301 attention_type: "mla", "gqa", or "mha"
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