Coverage for / home / jenkins / .local / lib / python3.10 / site-packages / hyper_parallel / core / dtensor / _ragged_utils.py: 75%
97 statements
« prev ^ index » next coverage.py v7.13.1, created at 2026-08-21 04:29 +0800
« prev ^ index » next coverage.py v7.13.1, created at 2026-08-21 04:29 +0800
1# Copyright 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"""Flat-storage geometry and distribution helpers for RaggedShard."""
16from math import prod
17from typing import NamedTuple, Optional, Sequence
19import numpy as np
21from hyper_parallel.core.dtensor._collective_utils import mesh_scatter_ragged
22from hyper_parallel.core.dtensor.layout import Layout, RaggedShardInfo
23from hyper_parallel.platform import get_platform
25platform = get_platform()
26Tensor = platform.Tensor
29def _layout_has_ragged_shard(layout: object) -> bool:
30 """Return whether a concrete Layout carries RaggedShard metadata."""
31 return isinstance(getattr(layout, "ragged_shard", None), RaggedShardInfo)
34class _RaggedSlice(NamedTuple):
35 """Flat interval owned by one rank in a RaggedShard layout."""
37 flat_start: int
38 flat_end: int
40 @property
41 def local_numel(self) -> int:
42 """Return the number of flat elements in the interval."""
43 return self.flat_end - self.flat_start
46def _normalize_global_shape(shape: Sequence[int]) -> tuple[int, ...]:
47 """Normalize a concrete global tensor shape."""
48 normalized = tuple(shape)
49 if any(
50 not isinstance(size, (int, np.integer)) or isinstance(size, bool) or size < 0
51 for size in normalized
52 ):
53 raise ValueError(
54 f"DTensor global shape must contain non-negative integers, got {normalized!r}"
55 )
56 return tuple(int(size) for size in normalized)
59def _compute_ragged_slice(
60 global_shape: Sequence[int],
61 layout: Layout,
62 local_rank: Optional[int] = None,
63) -> _RaggedSlice:
64 """Compute one rank's flat RaggedShard interval.
66 Phase one supports exactly one RaggedShard and Replicate placements on all
67 other mesh dimensions.
68 """
69 info = layout.ragged_shard
70 if info is None:
71 raise ValueError("RaggedShard slice computation requires a ragged layout")
73 for mesh_dim, placement in enumerate(layout.placements):
74 if mesh_dim == info.mesh_dim:
75 continue
76 if not placement.is_replicate():
77 raise NotImplementedError(
78 "RaggedShard phase one only supports Replicate on other mesh dimensions, "
79 f"got mesh_dim={mesh_dim}, placement={placement!r}"
80 )
82 shape = _normalize_global_shape(global_shape)
83 ragged = info.placement
84 prefix_ndim = len(ragged.dims)
85 if prefix_ndim > len(shape):
86 raise ValueError(
87 f"RaggedShard dims {ragged.dims!r} exceed global shape rank {len(shape)}"
88 )
90 mesh_dim_size = layout.mesh.size(info.mesh_dim)
91 if len(ragged.local_units) != mesh_dim_size:
92 raise ValueError(
93 "RaggedShard len(local_units) must equal mesh.size(mesh_dim), "
94 f"got len(local_units)={len(ragged.local_units)}, mesh_dim_size={mesh_dim_size}"
95 )
97 prefix_cells = prod(shape[:prefix_ndim])
98 total_units = sum(ragged.local_units)
99 if prefix_cells % total_units != 0:
100 raise ValueError(
101 "RaggedShard prefix cell count must be divisible by sum(local_units), "
102 f"got prefix_cells={prefix_cells}, local_units={ragged.local_units!r}"
103 )
105 if local_rank is None:
106 local_rank = layout.mesh.get_local_rank(info.mesh_dim)
107 if local_rank < 0 or local_rank >= mesh_dim_size:
108 raise ValueError(
109 f"RaggedShard local rank must be in [0, {mesh_dim_size}), got {local_rank}"
110 )
111 cells_per_unit = prefix_cells // total_units
112 prefix_start = sum(ragged.local_units[:local_rank]) * cells_per_unit
113 local_prefix_cells = ragged.local_units[local_rank] * cells_per_unit
114 suffix_numel = prod(shape[prefix_ndim:])
115 flat_start = prefix_start * suffix_numel
116 flat_end = flat_start + local_prefix_cells * suffix_numel
117 return _RaggedSlice(flat_start, flat_end)
120def _compute_ragged_splits(
121 global_shape: Sequence[int],
122 layout: Layout,
123) -> tuple[int, ...]:
124 """Return flat element counts contributed by all ranks on the ragged axis."""
125 info = layout.ragged_shard
126 if info is None:
127 raise ValueError("RaggedShard split computation requires a ragged layout")
128 return tuple(
129 _compute_ragged_slice(global_shape, layout, local_rank=rank).local_numel
130 for rank in range(layout.mesh.size(info.mesh_dim))
131 )
134def _interval_overlap_size(first: _RaggedSlice, second: _RaggedSlice) -> int:
135 """Return the size of two half-open flat intervals' intersection."""
136 return max(0, min(first.flat_end, second.flat_end) - max(first.flat_start, second.flat_start))
139def _compute_ragged_all_to_all_splits(
140 global_shape: Sequence[int],
141 from_layout: Layout,
142 to_layout: Layout,
143) -> tuple[tuple[int, ...], tuple[int, ...]]:
144 """Compute variable all-to-all splits for a local-units-only change."""
145 source_info = from_layout.ragged_shard
146 target_info = to_layout.ragged_shard
147 if source_info is None or target_info is None:
148 raise ValueError("RaggedShard all-to-all split computation requires ragged source and target layouts")
149 if source_info.mesh_dim != target_info.mesh_dim:
150 raise ValueError(
151 "RaggedShard all-to-all requires the same ragged mesh dimension, "
152 f"got source={source_info.mesh_dim}, target={target_info.mesh_dim}"
153 )
154 if source_info.placement.dims != target_info.placement.dims:
155 raise ValueError(
156 "RaggedShard all-to-all only supports local_units changes; dims must stay unchanged, "
157 f"got source={source_info.placement.dims!r}, target={target_info.placement.dims!r}"
158 )
160 mesh_dim = source_info.mesh_dim
161 source_rank = from_layout.mesh.get_local_rank(mesh_dim)
162 target_rank = to_layout.mesh.get_local_rank(mesh_dim)
163 source_interval = _compute_ragged_slice(global_shape, from_layout, source_rank)
164 target_interval = _compute_ragged_slice(global_shape, to_layout, target_rank)
165 input_splits = tuple(
166 _interval_overlap_size(
167 source_interval,
168 _compute_ragged_slice(global_shape, to_layout, rank),
169 )
170 for rank in range(to_layout.mesh.size(mesh_dim))
171 )
172 output_splits = tuple(
173 _interval_overlap_size(
174 _compute_ragged_slice(global_shape, from_layout, rank),
175 target_interval,
176 )
177 for rank in range(from_layout.mesh.size(mesh_dim))
178 )
179 return input_splits, output_splits
182def _slice_ragged_tensor(tensor: Tensor, layout: Layout) -> Tensor:
183 """Return an independent flat local shard from a complete global tensor."""
184 if hasattr(tensor, "is_contiguous") and not tensor.is_contiguous():
185 raise ValueError("distribute_tensor with RaggedShard requires a contiguous tensor")
186 ragged_slice = _compute_ragged_slice(tuple(tensor.shape), layout)
187 flat_tensor = tensor.reshape((-1,))
188 return flat_tensor[ragged_slice.flat_start:ragged_slice.flat_end].clone()
191def _scatter_ragged_tensor(
192 tensor: Tensor,
193 layout: Layout,
194 src_data_rank: int,
195) -> Tensor:
196 """Distribute variable flat shards from one group-relative source rank."""
197 if hasattr(tensor, "is_contiguous") and not tensor.is_contiguous():
198 raise ValueError("distribute_tensor with RaggedShard requires a contiguous tensor")
199 info = layout.ragged_shard
200 ragged_slice = _compute_ragged_slice(tuple(tensor.shape), layout)
201 flat_tensor = tensor.reshape((-1,))
202 output = platform.empty(
203 (ragged_slice.local_numel,),
204 dtype=tensor.dtype,
205 device=getattr(tensor, "device", None),
206 )
208 scatter_list = None
209 if layout.mesh.get_local_rank(info.mesh_dim) == src_data_rank:
210 scatter_list = []
211 for destination_rank in range(len(info.placement.local_units)):
212 destination_slice = _compute_ragged_slice(
213 tuple(tensor.shape),
214 layout,
215 local_rank=destination_rank,
216 )
217 scatter_list.append(
218 flat_tensor[destination_slice.flat_start:destination_slice.flat_end]
219 )
221 return mesh_scatter_ragged(
222 output,
223 scatter_list,
224 layout.mesh,
225 info.mesh_dim,
226 group_src=src_data_rank,
227 )