[diffusion] refactor: merge redundant default_dtype and param_dtype parameters in FSDP loader (#18789)
This commit is contained in:
@@ -88,8 +88,7 @@ class BridgeLoader(ComponentLoader):
|
|||||||
cpu_offload=server_args.dit_cpu_offload,
|
cpu_offload=server_args.dit_cpu_offload,
|
||||||
pin_cpu_memory=server_args.pin_cpu_memory,
|
pin_cpu_memory=server_args.pin_cpu_memory,
|
||||||
fsdp_inference=server_args.use_fsdp_inference,
|
fsdp_inference=server_args.use_fsdp_inference,
|
||||||
default_dtype=default_dtype,
|
param_dtype=default_dtype,
|
||||||
param_dtype=torch.bfloat16,
|
|
||||||
reduce_dtype=torch.float32,
|
reduce_dtype=torch.float32,
|
||||||
output_dtype=None,
|
output_dtype=None,
|
||||||
strict=False,
|
strict=False,
|
||||||
|
|||||||
@@ -105,7 +105,6 @@ class TransformerLoader(ComponentLoader):
|
|||||||
pin_cpu_memory=server_args.pin_cpu_memory,
|
pin_cpu_memory=server_args.pin_cpu_memory,
|
||||||
fsdp_inference=server_args.use_fsdp_inference,
|
fsdp_inference=server_args.use_fsdp_inference,
|
||||||
# TODO(will): make these configurable
|
# TODO(will): make these configurable
|
||||||
default_dtype=default_dtype,
|
|
||||||
param_dtype=torch.bfloat16,
|
param_dtype=torch.bfloat16,
|
||||||
reduce_dtype=torch.float32,
|
reduce_dtype=torch.float32,
|
||||||
output_dtype=None,
|
output_dtype=None,
|
||||||
|
|||||||
@@ -54,7 +54,6 @@ def maybe_load_fsdp_model(
|
|||||||
device: torch.device,
|
device: torch.device,
|
||||||
hsdp_replicate_dim: int,
|
hsdp_replicate_dim: int,
|
||||||
hsdp_shard_dim: int,
|
hsdp_shard_dim: int,
|
||||||
default_dtype: torch.dtype,
|
|
||||||
param_dtype: torch.dtype,
|
param_dtype: torch.dtype,
|
||||||
reduce_dtype: torch.dtype,
|
reduce_dtype: torch.dtype,
|
||||||
cpu_offload: bool = False,
|
cpu_offload: bool = False,
|
||||||
@@ -63,8 +62,15 @@ def maybe_load_fsdp_model(
|
|||||||
pin_cpu_memory: bool = True,
|
pin_cpu_memory: bool = True,
|
||||||
strict: bool = True,
|
strict: bool = True,
|
||||||
) -> torch.nn.Module:
|
) -> torch.nn.Module:
|
||||||
"""
|
"""Load a model with optional FSDP (Fully Sharded Data Parallel) support.
|
||||||
Load the model with FSDP if is training, else load the model without FSDP.
|
|
||||||
|
Args:
|
||||||
|
param_dtype: Data type for model parameters, also used for:
|
||||||
|
- Model initialization context (set_default_torch_dtype)
|
||||||
|
- FSDP mixed precision policy
|
||||||
|
- Weight loading and casting
|
||||||
|
reduce_dtype: Data type for gradient reduction in FSDP mixed precision.
|
||||||
|
strict: If True, enforce strict state dict loading (all keys must match).
|
||||||
"""
|
"""
|
||||||
# NOTE(will): cast_forward_inputs=True shouldn't be needed as we are
|
# NOTE(will): cast_forward_inputs=True shouldn't be needed as we are
|
||||||
# manually casting the inputs to the model
|
# manually casting the inputs to the model
|
||||||
@@ -79,7 +85,7 @@ def maybe_load_fsdp_model(
|
|||||||
mp_policy=mp_policy,
|
mp_policy=mp_policy,
|
||||||
)
|
)
|
||||||
|
|
||||||
with set_default_torch_dtype(default_dtype), torch.device("meta"):
|
with set_default_torch_dtype(param_dtype), torch.device("meta"):
|
||||||
model = model_cls(**init_params)
|
model = model_cls(**init_params)
|
||||||
|
|
||||||
# Check if we should use FSDP
|
# Check if we should use FSDP
|
||||||
@@ -120,7 +126,7 @@ def maybe_load_fsdp_model(
|
|||||||
model,
|
model,
|
||||||
weight_iterator,
|
weight_iterator,
|
||||||
device,
|
device,
|
||||||
default_dtype,
|
param_dtype,
|
||||||
strict=strict,
|
strict=strict,
|
||||||
cpu_offload=cpu_offload,
|
cpu_offload=cpu_offload,
|
||||||
param_names_mapping=param_names_mapping_fn,
|
param_names_mapping=param_names_mapping_fn,
|
||||||
|
|||||||
Reference in New Issue
Block a user