Skip to content

vllm_omni.diffusion.models.utils

Style module-attribute

Style = Literal[
    "colwise",
    "colwise_rep",
    "rowwise",
    "rowwise_rep",
    "replicate",
]

create_transformers_model

create_transformers_model(
    auto_cls: _BaseAutoModelClass,
    od_config: OmniDiffusionConfig,
    hf_config: PretrainedConfig,
    dtype: dtype | None = None,
    device: device | None = None,
) -> PreTrainedModel

Create a HuggingFace model using the given auto class and model name.

create_transformers_model_with_vllm_linears

create_transformers_model_with_vllm_linears(
    auto_cls: _BaseAutoModelClass,
    hf_config: PretrainedConfig,
    quant_config: QuantizationConfig | None,
    dtype: dtype,
    device: device,
    prefix: str = "",
    skip_modules: Sequence[str] = (),
) -> PreTrainedModel

Build a Transformers model with vLLM linears.

init_parameters

init_parameters(
    module: Module,
    dtype: dtype | None,
    device: device | None = None,
)

init_parameters_preserving_vllm_linear

init_parameters_preserving_vllm_linear(
    module: Module,
    dtype: dtype | None,
    device: device | None = None,
) -> None

make_attention_mask

make_attention_mask(
    hidden_states: Tensor, seq_lengths: list[int]
) -> Tensor | None

Build a padding mask only when the batch contains padding.

seq_lengths must be non-empty and contain valid sequence lengths. Returns a boolean mask of shape (len(seq_lengths), max(seq_lengths)) on the device of hidden_states, with True marking valid tokens, or None when all sequence lengths are equal.

Passing an all-true CUDA mask makes attention backends inspect a device scalar to select the dense path, which introduces a host-device sync.

recursive_replace_linear

recursive_replace_linear(
    model: Module, od_config: OmniDiffusionConfig
)

Recursively replace modules in the model as needed. Currently, this replaces: - nn.Linear with vLLM's tensor parallel linear classes

recursive_replace_linear_with_quantization_config

recursive_replace_linear_with_quantization_config(
    model: Module,
    quant_config: QuantizationConfig | None,
    prefix: str = "",
    skip_modules: Sequence[str] = (),
) -> None

replace_linear_class

replace_linear_class(
    linear: Linear,
    style: Style = "replicate",
    quant_config: QuantizationConfig | None = None,
    *,
    prefix: str = "",
) -> (
    ColumnParallelLinear
    | RowParallelLinear
    | ReplicatedLinear
)

Replace nn.Linear with one of vLLM's tensor parallel linear classes.

Parameters:

Name Type Description Default
linear Linear

nn.Linear to be replaced.

required
style Style

Tensor parallel style of the new linear, e.g. "colwise".

'replicate'
quant_config QuantizationConfig | None

Quantization config for the new linear.

None

Returns: The new linear.