core.context_parallel.utils#

Context-parallel batch partitioning helpers.

Module Contents#

Functions#

_get_batch_on_this_cp_rank_contiguous

Use contiguous CP shards while keeping dense attention-mask queries zigzag.

API#

core.context_parallel.utils._get_batch_on_this_cp_rank_contiguous(
batch: Dict[str, Any],
cp_group: torch.distributed.ProcessGroup,
) Dict[str, Any]#

Use contiguous CP shards while keeping dense attention-mask queries zigzag.