fix: drop context parallelism from the FSDP backend (#2386)

Co-authored-by: Kangrui Du <89372739+Rockdu@users.noreply.github.com>
This commit is contained in:
Zhichen Zeng
2026-08-11 15:05:07 -07:00
committed by GitHub
co-authored by Kangrui Du
parent 1b3ee448d8
commit 0ad5cb1835
4 changed files with 23 additions and 67 deletions
+2 -18
View File
@@ -659,25 +659,9 @@ class FSDPTrainRayActor(TrainRayActor):
raise NotImplementedError(f"Loading from checkpoint file {ref_load_path} not yet implemented")
def _get_model_inputs_args(self, batch: dict) -> dict:
input_ids = batch["tokens"]
position_ids = batch["position_ids"]
if get_parallel_state().cp.size > 1:
# TODO: Pin ring_flash_attn for torch 2.11+ compatibility; keep this local import to unblock non-FSDP+CP paths.
from ring_flash_attn import update_ring_flash_attn_params
if "cu_seqlens" in batch:
cu_seqlens = batch["cu_seqlens"]
if not cu_seqlens.is_cuda:
cu_seqlens = cu_seqlens.cuda()
update_ring_flash_attn_params(cu_seqlens, self.cp_group)
input_ids = torch.chunk(input_ids, get_parallel_state().cp.size, dim=1)[get_parallel_state().cp.rank]
position_ids = torch.chunk(position_ids, get_parallel_state().cp.size, dim=1)[get_parallel_state().cp.rank]
model_args = {
"input_ids": input_ids,
"position_ids": position_ids,
"input_ids": batch["tokens"],
"position_ids": batch["position_ids"],
"attention_mask": None,
}
+5 -14
View File
@@ -50,8 +50,9 @@ class FSDPArgs:
deterministic_mode: bool = False # This name must be the same as Megatron's
# Context Parallelism
context_parallel_size: int = 1 # Context Parallelism size
# The FSDP backend is pure data parallel. This knob only exists so shared argument
# validation can reject a context-parallel run with a clear message.
context_parallel_size: int = 1
# Profile
record_memory_history: bool = False
memory_snapshot_path: str = "snapshot.pickle"
@@ -117,15 +118,5 @@ def validate_hybrid_shard_args(args) -> None:
raise ValueError(f"dp_replicate_size must be at least 1, got {replicate_size}")
world_size = args.actor_num_nodes * args.actor_num_gpus_per_node
if args.context_parallel_size < 1:
raise ValueError(f"context_parallel_size must be at least 1, got {args.context_parallel_size}")
if world_size % args.context_parallel_size:
raise ValueError(
f"world_size({world_size}) must be divisible by " f"context_parallel_size({args.context_parallel_size})"
)
data_parallel_size = world_size // args.context_parallel_size
if data_parallel_size % replicate_size:
raise ValueError(
f"data_parallel_size({data_parallel_size}) must be divisible by " f"dp_replicate_size({replicate_size})"
)
if world_size % replicate_size:
raise ValueError(f"world_size({world_size}) must be divisible by dp_replicate_size({replicate_size})")
+16 -34
View File
@@ -14,30 +14,24 @@ logger = logging.getLogger(__name__)
def build_fsdp_meshes(
device_type: str,
world_size: int,
context_parallel_size: int,
dp_replicate_size: int,
) -> dict[str, DeviceMesh]:
"""Build the data/context-parallel views and the FSDP2 shard mesh."""
data_parallel_size = world_size // context_parallel_size
dp_cp_mesh = init_device_mesh(
"""Build the data-parallel view and the FSDP2 shard mesh."""
dp_mesh = init_device_mesh(
device_type,
mesh_shape=(data_parallel_size, context_parallel_size),
mesh_dim_names=("dp", "cp"),
mesh_shape=(world_size,),
mesh_dim_names=("dp",),
)
dp_mesh = dp_cp_mesh["dp"]
fsdp_mesh = dp_mesh
if dp_replicate_size > 1:
fsdp_mesh = dp_mesh._unflatten(
0,
(dp_replicate_size, data_parallel_size // dp_replicate_size),
(dp_replicate_size, world_size // dp_replicate_size),
("dp_replicate", "dp_shard"),
)
return {
"dp_cp": dp_cp_mesh,
"dp": dp_mesh,
"cp": dp_cp_mesh["cp"],
"fsdp": fsdp_mesh,
}
@@ -47,41 +41,29 @@ def create_fsdp_parallel_state(args: Namespace) -> ParallelState:
world_size = dist.get_world_size()
rank = dist.get_rank()
cp_size = args.context_parallel_size
dp_rank = rank // cp_size
cp_rank = rank % cp_size
meshes = build_fsdp_meshes(
device_type="cuda",
world_size=world_size,
context_parallel_size=cp_size,
dp_replicate_size=args.dp_replicate_size,
)
dp_mesh = meshes["dp"]
cp_mesh = meshes["cp"]
fsdp_mesh = meshes["fsdp"]
logger.info(
f"[Rank {rank}] FSDP mesh shape={fsdp_mesh.shape}, "
f"dp_replicate_size={args.dp_replicate_size}, "
f"dp_shard_size={(world_size // cp_size) // args.dp_replicate_size}, "
f"dp_rank={dp_rank}, cp_rank={cp_rank}"
f"dp_shard_size={world_size // args.dp_replicate_size}, "
f"dp_rank={rank}"
)
# Setup Ring Flash Attention with CP group from mesh (only when cp_size > 1)
if cp_size > 1:
# TODO: Pin ring_flash_attn for torch 2.11+ compatibility; keep this local import to unblock non-FSDP+CP paths.
from ring_flash_attn import substitute_hf_flash_attn
substitute_hf_flash_attn(cp_mesh.get_group(), heads_k_stride=1)
logger.info(f"[Rank {rank}] CP initialized via device mesh")
else:
logger.info(f"[Rank {rank}] Pure DP mode (cp_size=1)")
# The FSDP backend is pure data parallel: every parallelism axis other than dp is
# a single-rank group, so collectives issued on them are no-ops.
self_group = dist.new_group([rank])
parallel_state = ParallelState(
intra_dp=GroupInfo(
rank=dp_rank,
size=world_size // cp_size,
rank=rank,
size=world_size,
group=dp_mesh.get_group(),
),
intra_dp_cp=GroupInfo(
@@ -91,14 +73,14 @@ def create_fsdp_parallel_state(args: Namespace) -> ParallelState:
gloo_group=get_gloo_group(),
),
cp=GroupInfo(
rank=cp_rank,
size=cp_size,
group=cp_mesh.get_group(),
rank=0,
size=1,
group=self_group,
),
tp=GroupInfo(
rank=0,
size=1,
group=dist.new_group([rank]),
group=self_group,
),
pp=GroupInfo(rank=0, size=1, group=None),
ep=GroupInfo(rank=0, size=1, group=None),
@@ -81,7 +81,6 @@ def main() -> None:
meshes = build_fsdp_meshes(
device_type="cuda",
world_size=world_size,
context_parallel_size=1,
dp_replicate_size=args.replicate_size,
)
fsdp_mesh = meshes["fsdp"]