mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
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:
co-authored by
Kangrui Du
parent
1b3ee448d8
commit
0ad5cb1835
@@ -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,
|
||||
}
|
||||
|
||||
|
||||
@@ -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})")
|
||||
|
||||
@@ -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"]
|
||||
|
||||
Reference in New Issue
Block a user