[diffusion] chore: remove useless params (#15925)
This commit is contained in:
@@ -668,10 +668,6 @@ class TransformerLoader(ComponentLoader):
|
|||||||
"Only diffusers format is supported."
|
"Only diffusers format is supported."
|
||||||
)
|
)
|
||||||
|
|
||||||
if server_args.override_transformer_cls_name is not None:
|
|
||||||
cls_name = server_args.override_transformer_cls_name
|
|
||||||
logger.info("Overriding transformer cls_name to %s", cls_name)
|
|
||||||
|
|
||||||
server_args.model_paths["transformer"] = component_model_path
|
server_args.model_paths["transformer"] = component_model_path
|
||||||
|
|
||||||
# Config from Diffusers supersedes sgl_diffusion's model config
|
# Config from Diffusers supersedes sgl_diffusion's model config
|
||||||
|
|||||||
@@ -191,31 +191,26 @@ class PipelineStage(ABC):
|
|||||||
"""
|
"""
|
||||||
stage_name = self.__class__.__name__
|
stage_name = self.__class__.__name__
|
||||||
# Check if verification is enabled (simple approach for prototype)
|
# Check if verification is enabled (simple approach for prototype)
|
||||||
enable_verification = getattr(server_args, "enable_stage_verification", False)
|
|
||||||
|
|
||||||
if enable_verification:
|
# Pre-execution input verification
|
||||||
# Pre-execution input verification
|
try:
|
||||||
try:
|
input_result = self.verify_input(batch, server_args)
|
||||||
input_result = self.verify_input(batch, server_args)
|
self._run_verification(input_result, stage_name, "input")
|
||||||
self._run_verification(input_result, stage_name, "input")
|
except Exception as e:
|
||||||
except Exception as e:
|
logger.error("Input verification failed for %s: %s", stage_name, str(e))
|
||||||
logger.error("Input verification failed for %s: %s", stage_name, str(e))
|
raise
|
||||||
raise
|
|
||||||
|
|
||||||
# Execute the actual stage logic with unified profiling
|
# Execute the actual stage logic with unified profiling
|
||||||
with StageProfiler(stage_name, logger=logger, timings=batch.timings):
|
with StageProfiler(stage_name, logger=logger, timings=batch.timings):
|
||||||
result = self.forward(batch, server_args)
|
result = self.forward(batch, server_args)
|
||||||
|
|
||||||
if enable_verification:
|
# Post-execution output verification
|
||||||
# Post-execution output verification
|
try:
|
||||||
try:
|
output_result = self.verify_output(result, server_args)
|
||||||
output_result = self.verify_output(result, server_args)
|
self._run_verification(output_result, stage_name, "output")
|
||||||
self._run_verification(output_result, stage_name, "output")
|
except Exception as e:
|
||||||
except Exception as e:
|
logger.error("Output verification failed for %s: %s", stage_name, str(e))
|
||||||
logger.error(
|
raise
|
||||||
"Output verification failed for %s: %s", stage_name, str(e)
|
|
||||||
)
|
|
||||||
raise
|
|
||||||
|
|
||||||
return result
|
return result
|
||||||
|
|
||||||
|
|||||||
@@ -200,7 +200,6 @@ class DecodingStage(PipelineStage):
|
|||||||
- trajectory_latents (optional): Latents at different timesteps
|
- trajectory_latents (optional): Latents at different timesteps
|
||||||
- trajectory_timesteps (optional): Corresponding timesteps
|
- trajectory_timesteps (optional): Corresponding timesteps
|
||||||
server_args: Configuration containing:
|
server_args: Configuration containing:
|
||||||
- output_type: "latent" to skip decoding, otherwise decode to pixels
|
|
||||||
- vae_cpu_offload: Whether to offload VAE to CPU after decoding
|
- vae_cpu_offload: Whether to offload VAE to CPU after decoding
|
||||||
- model_loaded: Track VAE loading state
|
- model_loaded: Track VAE loading state
|
||||||
- model_paths: Path to VAE model if loading needed
|
- model_paths: Path to VAE model if loading needed
|
||||||
@@ -213,10 +212,7 @@ class DecodingStage(PipelineStage):
|
|||||||
# load vae if not already loaded (used for memory constrained devices)
|
# load vae if not already loaded (used for memory constrained devices)
|
||||||
self.load_model()
|
self.load_model()
|
||||||
|
|
||||||
if server_args.output_type == "latent":
|
frames = self.decode(batch.latents, server_args)
|
||||||
frames = batch.latents
|
|
||||||
else:
|
|
||||||
frames = self.decode(batch.latents, server_args)
|
|
||||||
|
|
||||||
# decode trajectory latents if needed
|
# decode trajectory latents if needed
|
||||||
if batch.return_trajectory_decoded:
|
if batch.return_trajectory_decoded:
|
||||||
|
|||||||
@@ -145,31 +145,6 @@ def _sanitize_for_logging(obj: Any, key_hint: str | None = None) -> Any:
|
|||||||
return "<unserializable>"
|
return "<unserializable>"
|
||||||
|
|
||||||
|
|
||||||
class ExecutionMode(str, Enum):
|
|
||||||
"""
|
|
||||||
Enumeration for different pipeline modes.
|
|
||||||
|
|
||||||
Inherits from str to allow string comparison for backward compatibility.
|
|
||||||
"""
|
|
||||||
|
|
||||||
INFERENCE = "inference"
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def from_string(cls, value: str) -> "ExecutionMode":
|
|
||||||
"""Convert string to ExecutionMode enum."""
|
|
||||||
try:
|
|
||||||
return cls(value.lower())
|
|
||||||
except ValueError:
|
|
||||||
raise ValueError(
|
|
||||||
f"Invalid mode: {value}. Must be one of: {', '.join([m.value for m in cls])}"
|
|
||||||
) from None
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def choices(cls) -> list[str]:
|
|
||||||
"""Get all available choices as strings for argparse."""
|
|
||||||
return [mode.value for mode in cls]
|
|
||||||
|
|
||||||
|
|
||||||
@dataclasses.dataclass
|
@dataclasses.dataclass
|
||||||
class ServerArgs:
|
class ServerArgs:
|
||||||
# Model and path configuration (for convenience)
|
# Model and path configuration (for convenience)
|
||||||
@@ -178,14 +153,7 @@ class ServerArgs:
|
|||||||
# Attention
|
# Attention
|
||||||
attention_backend: str = None
|
attention_backend: str = None
|
||||||
|
|
||||||
# Running mode
|
|
||||||
mode: ExecutionMode = ExecutionMode.INFERENCE
|
|
||||||
|
|
||||||
# Cache strategy
|
|
||||||
cache_strategy: str = "none"
|
|
||||||
|
|
||||||
# Distributed executor backend
|
# Distributed executor backend
|
||||||
distributed_executor_backend: str = "mp"
|
|
||||||
nccl_port: Optional[int] = None
|
nccl_port: Optional[int] = None
|
||||||
|
|
||||||
# HuggingFace specific parameters
|
# HuggingFace specific parameters
|
||||||
@@ -224,8 +192,6 @@ class ServerArgs:
|
|||||||
# Will adapt only q, k, v, o by default.
|
# Will adapt only q, k, v, o by default.
|
||||||
lora_target_modules: list[str] | None = None
|
lora_target_modules: list[str] | None = None
|
||||||
|
|
||||||
output_type: str = "pil"
|
|
||||||
|
|
||||||
# CPU offload parameters
|
# CPU offload parameters
|
||||||
dit_cpu_offload: bool = True
|
dit_cpu_offload: bool = True
|
||||||
dit_layerwise_offload: bool = False
|
dit_layerwise_offload: bool = False
|
||||||
@@ -266,9 +232,6 @@ class ServerArgs:
|
|||||||
|
|
||||||
scheduler_port: int = 5555
|
scheduler_port: int = 5555
|
||||||
|
|
||||||
# Stage verification
|
|
||||||
enable_stage_verification: bool = True
|
|
||||||
|
|
||||||
# Prompt text file for batch processing
|
# Prompt text file for batch processing
|
||||||
prompt_file_path: str | None = None
|
prompt_file_path: str | None = None
|
||||||
|
|
||||||
@@ -280,7 +243,6 @@ class ServerArgs:
|
|||||||
"vae": True,
|
"vae": True,
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
override_transformer_cls_name: str | None = None
|
|
||||||
|
|
||||||
# # DMD parameters
|
# # DMD parameters
|
||||||
# dmd_denoising_steps: List[int] | None = field(default=None)
|
# dmd_denoising_steps: List[int] | None = field(default=None)
|
||||||
@@ -369,24 +331,6 @@ class ServerArgs:
|
|||||||
help="The attention backend to use. If not specified, the backend is automatically selected based on hardware and installed packages.",
|
help="The attention backend to use. If not specified, the backend is automatically selected based on hardware and installed packages.",
|
||||||
)
|
)
|
||||||
|
|
||||||
# Running mode
|
|
||||||
parser.add_argument(
|
|
||||||
"--mode",
|
|
||||||
type=str,
|
|
||||||
choices=ExecutionMode.choices(),
|
|
||||||
default=ServerArgs.mode.value,
|
|
||||||
help="The mode to run SGLang-diffusion",
|
|
||||||
)
|
|
||||||
|
|
||||||
# distributed_executor_backend
|
|
||||||
parser.add_argument(
|
|
||||||
"--distributed-executor-backend",
|
|
||||||
type=str,
|
|
||||||
choices=["mp"],
|
|
||||||
default=ServerArgs.distributed_executor_backend,
|
|
||||||
help="The distributed executor backend to use",
|
|
||||||
)
|
|
||||||
|
|
||||||
# HuggingFace specific parameters
|
# HuggingFace specific parameters
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--trust-remote-code",
|
"--trust-remote-code",
|
||||||
@@ -466,15 +410,6 @@ class ServerArgs:
|
|||||||
help="Set timeout for torch.distributed initialization.",
|
help="Set timeout for torch.distributed initialization.",
|
||||||
)
|
)
|
||||||
|
|
||||||
# Output type
|
|
||||||
parser.add_argument(
|
|
||||||
"--output-type",
|
|
||||||
type=str,
|
|
||||||
default=ServerArgs.output_type,
|
|
||||||
choices=["pil"],
|
|
||||||
help="Output type for the generated video",
|
|
||||||
)
|
|
||||||
|
|
||||||
# Prompt text file for batch processing
|
# Prompt text file for batch processing
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--prompt-file-path",
|
"--prompt-file-path",
|
||||||
@@ -600,19 +535,6 @@ class ServerArgs:
|
|||||||
help="Whether to use webui for better display",
|
help="Whether to use webui for better display",
|
||||||
)
|
)
|
||||||
|
|
||||||
# Stage verification
|
|
||||||
parser.add_argument(
|
|
||||||
"--enable-stage-verification",
|
|
||||||
action=StoreBoolean,
|
|
||||||
default=ServerArgs.enable_stage_verification,
|
|
||||||
help="Enable input/output verification for pipeline stages",
|
|
||||||
)
|
|
||||||
parser.add_argument(
|
|
||||||
"--override-transformer-cls-name",
|
|
||||||
type=str,
|
|
||||||
default=ServerArgs.override_transformer_cls_name,
|
|
||||||
help="Override transformer cls name",
|
|
||||||
)
|
|
||||||
# LoRA
|
# LoRA
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--lora-path",
|
"--lora-path",
|
||||||
@@ -760,10 +682,6 @@ class ServerArgs:
|
|||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_kwargs(cls, **kwargs: Any) -> "ServerArgs":
|
def from_kwargs(cls, **kwargs: Any) -> "ServerArgs":
|
||||||
# Convert mode string to enum if necessary
|
|
||||||
if "mode" in kwargs and isinstance(kwargs["mode"], str):
|
|
||||||
kwargs["mode"] = ExecutionMode.from_string(kwargs["mode"])
|
|
||||||
|
|
||||||
kwargs["pipeline_config"] = PipelineConfig.from_kwargs(kwargs)
|
kwargs["pipeline_config"] = PipelineConfig.from_kwargs(kwargs)
|
||||||
return cls(**kwargs)
|
return cls(**kwargs)
|
||||||
|
|
||||||
@@ -885,14 +803,6 @@ class ServerArgs:
|
|||||||
else:
|
else:
|
||||||
self.disable_autocast = False
|
self.disable_autocast = False
|
||||||
|
|
||||||
# Validate mode consistency
|
|
||||||
assert isinstance(
|
|
||||||
self.mode, ExecutionMode
|
|
||||||
), f"Mode must be an ExecutionMode enum, got {type(self.mode)}"
|
|
||||||
assert (
|
|
||||||
self.mode in ExecutionMode.choices()
|
|
||||||
), f"Invalid execution mode: {self.mode}"
|
|
||||||
|
|
||||||
if self.tp_size == -1:
|
if self.tp_size == -1:
|
||||||
self.tp_size = 1
|
self.tp_size = 1
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user