Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions docs/user-guide/cli-reference.md
Original file line number Diff line number Diff line change
Expand Up @@ -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. |
Expand Down
2 changes: 2 additions & 0 deletions miles/backends/fsdp_utils/actor.py
Original file line number Diff line number Diff line change
Expand Up @@ -656,13 +656,15 @@ 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)"
)

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.
Expand Down
13 changes: 13 additions & 0 deletions miles/utils/arguments.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -97,6 +97,7 @@ def main():
args=Namespace(
diffusion_forward_dtype="bf16",
fsdp_reduce_dtype="fp32",
fsdp_reshard_after_forward=True,
gradient_checkpointing=False,
),
)
Expand Down
Loading