Skip to content

vllm_omni.diffusion.distributed.hsdp

logger module-attribute

logger = init_logger(__name__)

HSDPInferenceConfig dataclass

Configuration for HSDP inference.

This is a runtime config created from DiffusionParallelConfig's HSDP settings.

enabled class-attribute instance-attribute

enabled: bool = False

hsdp_replicate_size class-attribute instance-attribute

hsdp_replicate_size: int = 1

hsdp_shard_size class-attribute instance-attribute

hsdp_shard_size: int = -1

output_dtype class-attribute instance-attribute

output_dtype: dtype | None = None

param_dtype class-attribute instance-attribute

param_dtype: dtype = torch.bfloat16

reduce_dtype class-attribute instance-attribute

reduce_dtype: dtype = torch.float32

reshard_after_forward class-attribute instance-attribute

reshard_after_forward: bool = True

HSDPShardContext dataclass

Reusable state for incrementally sharding one HSDP model.

hsdp_kwargs instance-attribute

hsdp_kwargs: dict[str, Any]

ignored_params class-attribute instance-attribute

ignored_params: set[Parameter] | None = None

apply_hsdp_to_model

apply_hsdp_to_model(
    model: Module,
    hsdp_config: HSDPInferenceConfig,
    target_device: device | None = None,
) -> Module

Apply HSDP sharding to a model that already has weights loaded.

This function redistributes the model's parameters across GPUs using HSDP. The model should already have its weights loaded via the standard load_weights method.

Parameters:

Name Type Description Default
model Module

Model instance with weights already loaded

required
hsdp_config HSDPInferenceConfig

HSDP configuration with HSDP mesh dimensions

required
target_device device | None

Worker's execution device. When the model declares _hsdp_ignored_modules, those modules are excluded from FSDP's mesh-driven device placement, so the caller must specify where to put them. Optional only when there are no ignored modules.

None

Returns:

Type Description
Module

HSDP-wrapped model ready for inference

finalize_hsdp_root

finalize_hsdp_root(
    model: Module, context: HSDPShardContext
) -> None

Shard the model root after all selected child modules are sharded.

prepare_hsdp_shard_context

prepare_hsdp_shard_context(
    model: Module,
    hsdp_config: HSDPInferenceConfig,
    target_device: device | None = None,
) -> HSDPShardContext

Prepare mesh, precision, and ignored-parameter state for HSDP.

The returned context can be reused to shard completed child modules while checkpoint weights are still streaming, then to shard the root once all children have been loaded.

shard_hsdp_module

shard_hsdp_module(
    module: Module, context: HSDPShardContext
) -> None

Shard one completed child module with a prepared HSDP context.

shard_model

shard_model(
    model: Module,
    *,
    reshard_after_forward: bool = True,
    mp_policy: MixedPrecisionPolicy | None = None,
    mesh: DeviceMesh | None = None,
    hsdp_shard_conditions: list[
        Callable[[str, Module], bool]
    ],
    ignored_params: set[Parameter] | None = None,
    context: HSDPShardContext | None = None,
) -> None

Apply HSDP sharding to model modules based on shard conditions.

ignored_params (if provided) are excluded from the root fully_shard wrap, so they are not collected into the root flat-parameter, are not subject to MixedPrecisionPolicy, and retain their original dtype. Each per-submodule wrap receives the subset of ignored_params that it owns. This is required for packed integer parameters inside sharded transformer blocks; the root wrap receives the full set for all remaining parameters.