[diffusion] chore: clean excessive document (#16986)

This commit is contained in:
Mick
2026-01-13 21:33:50 +08:00
committed by GitHub
parent 888d7e54d1
commit 7a869045b6
27 changed files with 36 additions and 236 deletions

View File

@@ -124,27 +124,29 @@ def _discover_and_register_pipelines():
)
for cls in entry_cls_list:
if hasattr(cls, "pipeline_name"):
if cls.pipeline_name in _PIPELINE_REGISTRY:
logger.warning(
f"Duplicate pipeline name '{cls.pipeline_name}' found. Overwriting."
)
_PIPELINE_REGISTRY[cls.pipeline_name] = cls
if not issubclass(cls, ComposedPipelineBase):
continue
if cls.pipeline_name in _PIPELINE_REGISTRY:
logger.warning(
f"Duplicate pipeline name '{cls.pipeline_name}' found. Overwriting."
)
_PIPELINE_REGISTRY[cls.pipeline_name] = cls
# Auto-register config classes if Pipeline class has them defined
# because comfyui get model from a single weight file, so we need to register the config classes here
if hasattr(cls, "pipeline_config_cls") and hasattr(
cls, "sampling_params_cls"
):
_PIPELINE_CONFIG_REGISTRY[cls.pipeline_name] = (
cls.pipeline_config_cls,
cls.sampling_params_cls,
)
logger.debug(
f"Auto-registered config classes for pipeline '{cls.pipeline_name}': "
f"PipelineConfig={cls.pipeline_config_cls.__name__}, "
f"SamplingParams={cls.sampling_params_cls.__name__}"
)
# Special handling for ComfyUI Pipelines:
# Auto-register config classes if Pipeline class has them defined
# since comfyui get model from a single weight file, so we need to register the config classes here
if hasattr(cls, "pipeline_config_cls") and hasattr(
cls, "sampling_params_cls"
):
_PIPELINE_CONFIG_REGISTRY[cls.pipeline_name] = (
cls.pipeline_config_cls,
cls.sampling_params_cls,
)
logger.debug(
f"Auto-registered config classes for pipeline '{cls.pipeline_name}': "
f"PipelineConfig={cls.pipeline_config_cls.__name__}, "
f"SamplingParams={cls.sampling_params_cls.__name__}"
)
logger.debug(
f"Registering pipelines complete, {len(_PIPELINE_REGISTRY)} pipelines registered"
)
@@ -155,12 +157,6 @@ def get_pipeline_config_classes(
) -> Tuple[Type[PipelineConfig], Type[Any]] | None:
"""
Get the configuration classes for a pipeline.
Args:
pipeline_class_name: The name of the pipeline class
Returns:
A tuple of (PipelineConfig class, SamplingParams class) if found, None otherwise
"""
# Ensure pipelines are discovered first
_discover_and_register_pipelines()
@@ -325,11 +321,8 @@ def get_model_info(
manually registered mapping based on the model path.
Args:
model_path: Path to the model or HuggingFace model ID
backend: Backend to use ('auto', 'sglang', 'diffusers'). If None, uses 'auto'.
Returns:
ModelInfo with the resolved pipeline class and config classes, or None if not found.
"""
# import Backend enum here to avoid circular imports
from sglang.multimodal_gen.runtime.server_args import Backend

View File

@@ -633,8 +633,6 @@ def patch_tensor_parallel_group(tp_group: GroupCoordinator):
This method is for draft workers of speculative decoding to run draft model
with different tp degree from that of target model workers.
Args:
tp_group (GroupCoordinator): the tp group coordinator
"""
global _TP_STATE_PATCHED
assert not _TP_STATE_PATCHED, "Should not call when it's already patched"

View File

@@ -87,12 +87,6 @@ class DiffGenerator:
"""
Create a DiffGenerator from a pretrained model.
Args:
**kwargs: Additional arguments to customize model loading, set any ServerArgs or PipelineConfig attributes here.
Returns:
The created DiffGenerator
Priority level: Default pipeline config < User's pipeline config < User's kwargs
"""
# If users also provide some kwargs, it will override the ServerArgs and PipelineConfig.

View File

@@ -139,12 +139,6 @@ class AttentionImpl(ABC, Generic[T]):
Called AFTER all_to_all for distributed attention
Args:
qkv: The query-key-value tensor
attn_metadata: Metadata for the attention operation
Returns:
Processed QKV tensor
"""
return qkv
@@ -161,12 +155,6 @@ class AttentionImpl(ABC, Generic[T]):
Called BEFORE all_to_all for distributed attention
Args:
output: The output tensor from the attention operation
attn_metadata: Metadata for the attention operation
Returns:
Postprocessed output tensor
"""
return output

View File

@@ -170,7 +170,6 @@ def global_force_attn_backend_context_manager(
manager.
Arguments:
* attn_backend: attention backend to force
Returns:

View File

@@ -254,9 +254,6 @@ class ComponentLoader(ABC):
Args:
module_type: Type of module (e.g., "vae", "text_encoder", "transformer", "scheduler")
transformers_or_diffusers: Whether the module is from transformers or diffusers
Returns:
A component loader for the specified module type
"""
# Map of module types to their loader classes and expected library
module_type = _normalize_module_type(module_type)
@@ -803,8 +800,6 @@ class PipelineComponentLoader:
component_model_path: Path to the component model
transformers_or_diffusers: Whether the module is from transformers or diffusers
Returns:
The loaded module
"""
# Get the appropriate loader for this module type

View File

@@ -181,8 +181,6 @@ def shard_model(
which modules to shard with FSDP.
pin_cpu_memory (bool): If set to True, FSDP will pin the CPU memory of the offloaded parameters.
Raises:
ValueError: If no layer modules were sharded, indicating that no shard_condition was triggered.
"""
if fsdp_shard_conditions is None or len(fsdp_shard_conditions) == 0:
logger.warning(
@@ -244,8 +242,6 @@ def load_model_from_full_model_state_dict(
* **missing_keys** is a list of str containing the missing keys
* **unexpected_keys** is a list of str containing the unexpected keys
Raises:
NotImplementedError: If got FSDP with more than 1D.
"""
meta_sd = model.state_dict()
param_dict = dict(model.named_parameters())

View File

@@ -447,10 +447,6 @@ class CLIPTextTransformer(nn.Module):
inputs_embeds: torch.Tensor | None = None,
output_hidden_states: bool | None = None,
) -> BaseEncoderOutput:
r"""
Returns:
"""
output_hidden_states = (
output_hidden_states
if output_hidden_states is not None

View File

@@ -196,16 +196,6 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler
) -> torch.FloatTensor:
"""
Forward process in flow-matching
Args:
sample (`torch.FloatTensor`):
The input sample.
timestep (`int`, *optional*):
The current timestep in the diffusion chain.
Returns:
`torch.FloatTensor`:
A scaled input sample.
"""
# Make sure sigmas and timesteps have the same device and dtype as original_samples
sigmas = self.sigmas.to(device=sample.device, dtype=sample.dtype)

View File

@@ -78,16 +78,7 @@ def modulate(
shift: torch.Tensor | None = None,
scale: torch.Tensor | None = None,
) -> torch.Tensor:
"""modulate by shift and scale
Args:
x (torch.Tensor): input tensor.
shift (torch.Tensor, optional): shift tensor. Defaults to None.
scale (torch.Tensor, optional): scale tensor. Defaults to None.
Returns:
torch.Tensor: the output tensor after modulate.
"""
"""modulate by shift and scale"""
if scale is None and shift is None:
return x
elif shift is None:

View File

@@ -544,8 +544,6 @@ class AutoencoderKL(nn.Module):
sample (`torch.Tensor`): Input sample.
sample_posterior (`bool`, *optional*, defaults to `False`):
Whether to sample from the posterior.
return_dict (`bool`, *optional*, defaults to `True`):
Whether or not to return a [`DecoderOutput`] instead of a plain tuple.
"""
x = sample
posterior = self.encode(x).latent_dist

View File

@@ -101,10 +101,6 @@ def load_image(
convert_method (Callable[[PIL.Image.Image], PIL.Image.Image], *optional*):
A conversion method to apply to the image after loading it. When set to `None` the image will be converted
"RGB".
Returns:
`PIL.Image.Image`:
A PIL Image.
"""
if isinstance(image, str):
if image.startswith("http://") or image.startswith("https://"):

View File

@@ -82,13 +82,6 @@ class PipelineStage(ABC):
result.add_check("image_latent", batch.image_latent, V.is_tensor)
return result
Args:
batch: The current batch information.
server_args: The inference arguments.
Returns:
A VerificationResult containing the verification status.
"""
# Default implementation - no verification
return VerificationResult()
@@ -119,9 +112,7 @@ class PipelineStage(ABC):
"""
Verify the output for the stage.
Args:
batch: The current batch information.
server_args: The inference arguments.
Returns:
A VerificationResult containing the verification status.
@@ -182,9 +173,7 @@ class PipelineStage(ABC):
Execute the stage's processing on the batch with optional verification and logging.
Should not be overridden by subclasses.
Args:
batch: The current batch information.
server_args: The inference arguments.
Returns:
The updated batch information after this stage's processing.
@@ -232,9 +221,7 @@ class PipelineStage(ABC):
This method should be implemented by subclasses to provide the forward
processing logic for the stage.
Args:
batch: The current batch information.
server_args: The inference arguments.
Returns:
The updated batch information after this stage's processing.

View File

@@ -38,9 +38,7 @@ class ConditioningStage(PipelineStage):
"""
Apply conditioning to the diffusion process.
Args:
batch: The current batch information.
server_args: The inference arguments.
Returns:
The batch with applied conditioning.

View File

@@ -193,21 +193,6 @@ class DecodingStage(PipelineStage):
representations to pixel-space video/images. It also optionally decodes
trajectory latents for visualization purposes.
Args:
batch: The current batch containing:
- latents: Tensor to decode (batch, channels, frames, height_latents, width_latents)
- return_trajectory_decoded (optional): Flag to decode trajectory latents
- trajectory_latents (optional): Latents at different timesteps
- trajectory_timesteps (optional): Corresponding timesteps
server_args: Configuration containing:
- vae_cpu_offload: Whether to offload VAE to CPU after decoding
- model_loaded: Track VAE loading state
- model_paths: Path to VAE model if loading needed
Returns:
Modified batch with:
- output: Decoded frames (batch, channels, frames, height, width) as CPU float32
- trajectory_decoded (if requested): List of decoded frames per timestep
"""
# load vae if not already loaded (used for memory constrained devices)
self.load_model()

View File

@@ -478,10 +478,6 @@ class DenoisingStage(PipelineStage):
"""
Prepare all necessary invariant variables for the denoising loop.
Args:
batch: The current batch information.
server_args: The inference arguments.
Returns:
A dictionary containing all the prepared variables for the denoising loop.
"""
@@ -940,13 +936,6 @@ class DenoisingStage(PipelineStage):
) -> Req:
"""
Run the denoising loop.
Args:
batch: The current batch information.
server_args: The inference arguments.
Returns:
The batch with denoised latents.
"""
# Prepare variables for the denoising loop
@@ -1127,13 +1116,6 @@ class DenoisingStage(PipelineStage):
) -> tqdm:
"""
Create a progress bar for the denoising process.
Args:
iterable: The iterable to iterate over.
total: The total number of items.
Returns:
A tqdm progress bar.
"""
local_rank = get_world_group().local_rank
disable = local_rank != 0
@@ -1176,11 +1158,6 @@ class DenoisingStage(PipelineStage):
Args:
i: The current timestep index.
batch: The current batch information.
server_args: The inference arguments.
Returns:
The attention metadata, or None if not applicable.
"""
attn_metadata = None
self.attn_metadata_builder = None
@@ -1375,10 +1352,6 @@ class DenoisingStage(PipelineStage):
def prepare_sta_param(self, batch: Req, server_args: ServerArgs):
"""
Prepare Sliding Tile Attention (STA) parameters and settings.
Args:
batch: The current batch information.
server_args: The inference arguments.
"""
# TODO(kevin): STA mask search, currently only support Wan2.1 with 69x768x1280
STA_mode = server_args.STA_mode

View File

@@ -60,9 +60,7 @@ class EncodingStage(PipelineStage):
"""
Encode pixel space representations into latent space.
Args:
batch: The current batch information.
server_args: The inference arguments.
Returns:
The batch with encoded latents.

View File

@@ -96,13 +96,6 @@ class ImageEncodingStage(PipelineStage):
) -> Req:
"""
Encode the prompt into image encoder hidden states.
Args:
batch: The current batch information.
server_args: The inference arguments.
Returns:
The batch with encoded prompt embeddings.
"""
if batch.condition_image is None:
@@ -212,13 +205,6 @@ class ImageVAEEncodingStage(PipelineStage):
) -> Req:
"""
Encode pixel representations into latent space.
Args:
batch: The current batch information.
server_args: The inference arguments.
Returns:
The batch with encoded outputs.
"""
if batch.condition_image is None:

View File

@@ -182,13 +182,6 @@ class InputValidationStage(PipelineStage):
) -> Req:
"""
Validate and prepare inputs.
Args:
batch: The current batch information.
server_args: The inference arguments.
Returns:
The validated batch information.
"""
self._generate_seeds(batch, server_args)

View File

@@ -42,9 +42,7 @@ class LatentPreparationStage(PipelineStage):
"""
Prepare initial latent variables for the diffusion process.
Args:
batch: The current batch information.
server_args: The inference arguments.
Returns:
The batch with prepared latent variables.
@@ -111,9 +109,7 @@ class LatentPreparationStage(PipelineStage):
"""
Adjust video length based on VAE version.
Args:
batch: The current batch information.
server_args: The inference arguments.
Returns:
The batch with adjusted video length.

View File

@@ -53,13 +53,6 @@ class TextEncodingStage(PipelineStage):
) -> Req:
"""
Encode the prompt into text encoder hidden states.
Args:
batch: The current batch information.
server_args: The inference arguments.
Returns:
The batch with encoded prompt embeddings.
"""
assert len(self.tokenizers) == len(self.text_encoders)
assert len(self.text_encoders) == len(

View File

@@ -61,9 +61,7 @@ class TimestepPreparationStage(PipelineStage):
"""
Prepare timesteps for the diffusion process.
Args:
batch: The current batch information.
server_args: The inference arguments.
Returns:
The batch with prepared timesteps.

View File

@@ -738,17 +738,6 @@ class ServerArgs:
) -> int:
"""
Find an available port with retry logic.
Args:
port: Initial port to check
port_inc: Port increment for each attempt
max_attempts: Maximum number of attempts to find an available port
Returns:
An available port number
Raises:
RuntimeError: If no available port is found after max_attempts
"""
attempts = 0
original_port = port
@@ -1088,13 +1077,6 @@ _global_server_args = None
def prepare_server_args(argv: list[str]) -> ServerArgs:
"""
Prepare the inference arguments from the command line arguments.
Args:
argv: The command line arguments. Typically, it should be `sys.argv[1:]`
to ensure compatibility with `parse_args` when no arguments are passed.
Returns:
The inference arguments.
"""
parser = FlexibleArgumentParser()
ServerArgs.add_cli_args(parser)

View File

@@ -296,14 +296,7 @@ def load_dict(file_path):
def get_diffusers_component_config(
model_path: str,
) -> dict[str, Any]:
"""Gets a configuration of a submodule for the given diffusers model.
Args:
model_path: the path of the submodule (can be local path or HuggingFace model ID)
Returns:
The loaded configuration.
"""
"""Gets a configuration of a submodule for the given diffusers model."""
# Download from HuggingFace Hub if path doesn't exist locally
if not os.path.exists(model_path):

View File

@@ -83,12 +83,7 @@ def parse_args():
def collect_test_items(files, filter_expr=None):
"""Collect test item node IDs from the given files using pytest --collect-only.
Raises:
RuntimeError: If pytest collection fails due to errors (e.g., syntax errors,
import errors, or other collection failures).
"""
"""Collect test item node IDs from the given files using pytest --collect-only."""
cmd = [sys.executable, "-m", "pytest", "--collect-only", "-q"]
if filter_expr:
cmd.extend(["-k", filter_expr])

View File

@@ -277,7 +277,7 @@ def _get_video_dimensions_from_metadata(
if width == 0 or height == 0:
return None
return (int(width), int(height))
return int(width), int(height)
def _get_video_dimensions_from_frame(cap: cv2.VideoCapture) -> tuple[int, int]:
@@ -289,8 +289,6 @@ def _get_video_dimensions_from_frame(cap: cv2.VideoCapture) -> tuple[int, int]:
Returns:
Tuple of (width, height)
Raises:
ValueError: If unable to read a frame from the video
"""
ret, frame = cap.read()
if not ret or frame is None:
@@ -298,7 +296,7 @@ def _get_video_dimensions_from_frame(cap: cv2.VideoCapture) -> tuple[int, int]:
# frame.shape is (height, width, channels)
height, width = frame.shape[:2]
return (int(width), int(height))
return int(width), int(height)
def get_video_dimensions(file_path: str) -> tuple[int, int]:
@@ -306,14 +304,9 @@ def get_video_dimensions(file_path: str) -> tuple[int, int]:
Tries to get dimensions from metadata first, falls back to reading first frame.
Args:
file_path: Path to the video file
Returns:
Tuple of (width, height)
Raises:
ValueError: If unable to get video dimensions
"""
cap = cv2.VideoCapture(file_path)
try:

View File

@@ -596,11 +596,7 @@ def set_mixed_precision_policy(
def get_compute_dtype() -> torch.dtype:
"""Get the current compute dtype from mixed precision policy.
Returns:
torch.dtype: The compute dtype to use, defaults to get_default_dtype() if no policy set
"""
"""Get the current compute dtype from mixed precision policy."""
if not hasattr(_mixed_precision_state, "state"):
return torch.get_default_dtype()
else: