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_preserving_vllm_linear ¶
init_parameters_preserving_vllm_linear(
module: Module,
dtype: dtype | None,
device: device | None = None,
) -> None
make_attention_mask ¶
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 |
| 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.