diff --git a/docs/user-guide/cli-reference.md b/docs/user-guide/cli-reference.md index 237df50f..4c9286e8 100644 --- a/docs/user-guide/cli-reference.md +++ b/docs/user-guide/cli-reference.md @@ -130,6 +130,7 @@ See [Dtype Control](../advanced/dtype-control.md). | `--train-backend` | enum | `fsdp` | Only value. | | `--fsdp-master-dtype` | enum | `fp32` | `fp32` / `bf16` / `fp16`. Load, shard, and optimizer-state precision. | | `--fsdp-reduce-dtype` | enum | `fp32` | `bf16` matches flow_grpo's all-bf16 policy but adds cross-rank add-noise. | +| `--fsdp-reshard-after-forward` / `--no-fsdp-reshard-after-forward` | flag | on | On = ZeRO-3. Off keeps gathered params between forward and backward (ZeRO-2): no backward all-gather, higher memory. | | `--diffusion-forward-dtype` | enum | `bf16` | `bf16` / `fp16` / `fp32`. | | `--fsdp-cpu-offload` | flag | off | Offloads params, grads, optimizer state; the optimizer then runs on CPU. | | `--fsdp-cpu-backend` | str | `gloo` | CPU process group for the above. | diff --git a/miles/backends/fsdp_utils/actor.py b/miles/backends/fsdp_utils/actor.py index 332cfb09..2dc8017c 100644 --- a/miles/backends/fsdp_utils/actor.py +++ b/miles/backends/fsdp_utils/actor.py @@ -656,6 +656,7 @@ def apply_fsdp2( logger.info( f"FSDP: wrapping {len(modules)} modules of type {layer_cls_to_wrap}, " f"param_dtype={param_dtype}, reduce_dtype={reduce_dtype}, " + f"reshard_after_forward={args.fsdp_reshard_after_forward}, " f"param_dtype_overrides={param_dtype_maps.override_count} " f"({param_dtype_maps.override_numel:,} parameters)" ) @@ -663,6 +664,7 @@ def apply_fsdp2( fsdp_kwargs = { "offload_policy": offload_policy, "mesh": mesh, + "reshard_after_forward": args.fsdp_reshard_after_forward, } # input_dtype_policy owns boundary casts; autocast owns compute and keeps grad-ckpt recompute consistent. diff --git a/miles/utils/arguments.py b/miles/utils/arguments.py index deeb75f9..98ab8fe8 100644 --- a/miles/utils/arguments.py +++ b/miles/utils/arguments.py @@ -146,6 +146,19 @@ def add_train_arguments(parser): "of bf16 add-non-associativity noise across ranks." ), ) + parser.add_argument( + "--fsdp-reshard-after-forward", + action=argparse.BooleanOptionalAction, + default=True, + help=( + "Whether FSDP2 reshards parameters after forward " + "(fully_shard's reshard_after_forward). True (default) is " + "ZeRO-3; --no-fsdp-reshard-after-forward keeps gathered " + "parameters between forward and backward (ZeRO-2), saving " + "the backward all-gather at the cost of holding unsharded " + "parameters through the step." + ), + ) parser.add_argument( "--fsdp-flow-shift", type=float, diff --git a/tests/fast-gpu/backends/fsdp_utils/_param_dtype_map_integration_worker.py b/tests/fast-gpu/backends/fsdp_utils/_param_dtype_map_integration_worker.py index 47a00200..344080b7 100644 --- a/tests/fast-gpu/backends/fsdp_utils/_param_dtype_map_integration_worker.py +++ b/tests/fast-gpu/backends/fsdp_utils/_param_dtype_map_integration_worker.py @@ -97,6 +97,7 @@ def main(): args=Namespace( diffusion_forward_dtype="bf16", fsdp_reduce_dtype="fp32", + fsdp_reshard_after_forward=True, gradient_checkpointing=False, ), )