Returns:
tuple[int, ...]: Tuple of offsets for each dimension.
"""
if dtensor_layout is None:
# If layout is None, return all zeros (no sharding)
return tuple(0 for _ in global_shape)
# Validate layout attributes
if not hasattr(dtensor_layout, 'mesh_shape') or dtensor_layout.mesh_shape is None:
raise ValueError("Layout must have mesh_shape attribute")
if not hasattr(dtensor_layout, 'tensor_map') or dtensor_layout.tensor_map is None:
raise ValueError("Layout must have tensor_map attribute")
if not hasattr(dtensor_layout, 'rank_list') or dtensor_layout.rank_list is None:
raise ValueError("Layout must have rank_list attribute")
if current_rank not in dtensor_layout.rank_list:
raise ValueError(
f"Current rank {current_rank} not found in layout's rank_list {dtensor_layout.rank_list}")
inner_rank_id = dtensor_layout.rank_list.index(current_rank)
# Calculate slice area using infer_slice_area_by_rank
slice_area = infer_slice_area_by_layout(
dtensor_layout,
inner_rank_id,
global_shape,
)