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

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 

18 

19import numpy as np 

20 

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 

24 

25platform = get_platform() 

26Tensor = platform.Tensor 

27 

28 

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) 

32 

33 

34class _RaggedSlice(NamedTuple): 

35 """Flat interval owned by one rank in a RaggedShard layout.""" 

36 

37 flat_start: int 

38 flat_end: int 

39 

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 

44 

45 

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) 

57 

58 

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. 

65 

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

72 

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 ) 

81 

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 ) 

89 

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 ) 

96 

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 ) 

104 

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) 

118 

119 

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 ) 

132 

133 

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

137 

138 

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 ) 

159 

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 

180 

181 

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

189 

190 

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 ) 

207 

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 ) 

220 

221 return mesh_scatter_ragged( 

222 output, 

223 scatter_list, 

224 layout.mesh, 

225 info.mesh_dim, 

226 group_src=src_data_rank, 

227 )