[diffusion] chore: remove redundant identity preprocess_text functions(#20633)
Co-authored-by: Fengyuan Yu <15fengyuan@gmail.com>
This commit is contained in:
co-authored by
Fengyuan Yu
parent
e570ca96f6
commit
9fa7b974fd
@@ -97,10 +97,6 @@ class STA_Mode(str, Enum):
|
|||||||
NONE = None
|
NONE = None
|
||||||
|
|
||||||
|
|
||||||
def preprocess_text(prompt: str) -> str:
|
|
||||||
return prompt
|
|
||||||
|
|
||||||
|
|
||||||
def postprocess_text(output: BaseEncoderOutput, _text_inputs) -> torch.tensor:
|
def postprocess_text(output: BaseEncoderOutput, _text_inputs) -> torch.tensor:
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
@@ -206,8 +202,8 @@ class PipelineConfig:
|
|||||||
def postprocess_image(self, image):
|
def postprocess_image(self, image):
|
||||||
return image.last_hidden_state
|
return image.last_hidden_state
|
||||||
|
|
||||||
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
|
preprocess_text_funcs: tuple[Callable[[str], str] | None, ...] = field(
|
||||||
default_factory=lambda: (preprocess_text,)
|
default_factory=lambda: (None,)
|
||||||
)
|
)
|
||||||
|
|
||||||
# get prompt_embeds from encoder output
|
# get prompt_embeds from encoder output
|
||||||
|
|||||||
@@ -23,12 +23,10 @@ from sglang.multimodal_gen.configs.models.vaes.flux import Flux2VAEConfig, FluxV
|
|||||||
from sglang.multimodal_gen.configs.pipeline_configs.base import (
|
from sglang.multimodal_gen.configs.pipeline_configs.base import (
|
||||||
ImagePipelineConfig,
|
ImagePipelineConfig,
|
||||||
ModelTaskType,
|
ModelTaskType,
|
||||||
preprocess_text,
|
|
||||||
shard_rotary_emb_for_sp,
|
shard_rotary_emb_for_sp,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.configs.pipeline_configs.hunyuan import (
|
from sglang.multimodal_gen.configs.pipeline_configs.hunyuan import (
|
||||||
clip_postprocess_text,
|
clip_postprocess_text,
|
||||||
clip_preprocess_text,
|
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import _pack_latents
|
from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import _pack_latents
|
||||||
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
|
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
|
||||||
@@ -65,8 +63,8 @@ class FluxPipelineConfig(ImagePipelineConfig):
|
|||||||
default_factory=lambda: ("bf16", "bf16")
|
default_factory=lambda: ("bf16", "bf16")
|
||||||
)
|
)
|
||||||
|
|
||||||
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
|
preprocess_text_funcs: tuple[Callable[[str], str] | None, ...] = field(
|
||||||
default_factory=lambda: (clip_preprocess_text, preprocess_text),
|
default_factory=lambda: (None, None),
|
||||||
)
|
)
|
||||||
|
|
||||||
postprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
|
postprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
|
||||||
@@ -650,8 +648,8 @@ class Flux2KleinPipelineConfig(Flux2PipelineConfig):
|
|||||||
default_factory=lambda: (Qwen3TextConfig(),)
|
default_factory=lambda: (Qwen3TextConfig(),)
|
||||||
)
|
)
|
||||||
|
|
||||||
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
|
preprocess_text_funcs: tuple[Callable[[str], str] | None, ...] = field(
|
||||||
default_factory=lambda: (preprocess_text,),
|
default_factory=lambda: (None,),
|
||||||
)
|
)
|
||||||
|
|
||||||
postprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
|
postprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
|
||||||
|
|||||||
@@ -56,10 +56,6 @@ def llama_postprocess_text(outputs: BaseEncoderOutput, _text_inputs) -> torch.te
|
|||||||
return last_hidden_state
|
return last_hidden_state
|
||||||
|
|
||||||
|
|
||||||
def clip_preprocess_text(prompt: str) -> str:
|
|
||||||
return prompt
|
|
||||||
|
|
||||||
|
|
||||||
def clip_postprocess_text(outputs: BaseEncoderOutput, _text_inputs) -> torch.tensor:
|
def clip_postprocess_text(outputs: BaseEncoderOutput, _text_inputs) -> torch.tensor:
|
||||||
pooler_output: torch.tensor = outputs.pooler_output
|
pooler_output: torch.tensor = outputs.pooler_output
|
||||||
return pooler_output
|
return pooler_output
|
||||||
@@ -84,8 +80,8 @@ class HunyuanConfig(PipelineConfig):
|
|||||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(
|
text_encoder_configs: tuple[EncoderConfig, ...] = field(
|
||||||
default_factory=lambda: (LlamaConfig(), CLIPTextConfig())
|
default_factory=lambda: (LlamaConfig(), CLIPTextConfig())
|
||||||
)
|
)
|
||||||
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
|
preprocess_text_funcs: tuple[Callable[[str], str] | None, ...] = field(
|
||||||
default_factory=lambda: (llama_preprocess_text, clip_preprocess_text)
|
default_factory=lambda: (llama_preprocess_text, None)
|
||||||
)
|
)
|
||||||
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.tensor], ...] = (
|
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.tensor], ...] = (
|
||||||
field(default_factory=lambda: (llama_postprocess_text, clip_postprocess_text))
|
field(default_factory=lambda: (llama_postprocess_text, clip_postprocess_text))
|
||||||
|
|||||||
@@ -14,7 +14,6 @@ from sglang.multimodal_gen.configs.models.vaes.ltx_audio import LTXAudioVAEConfi
|
|||||||
from sglang.multimodal_gen.configs.pipeline_configs.base import (
|
from sglang.multimodal_gen.configs.pipeline_configs.base import (
|
||||||
ModelTaskType,
|
ModelTaskType,
|
||||||
PipelineConfig,
|
PipelineConfig,
|
||||||
preprocess_text,
|
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.distributed import (
|
from sglang.multimodal_gen.runtime.distributed import (
|
||||||
get_sp_parallel_rank,
|
get_sp_parallel_rank,
|
||||||
@@ -189,8 +188,8 @@ class LTX2PipelineConfig(PipelineConfig):
|
|||||||
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16",))
|
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16",))
|
||||||
text_encoder_extra_args: list[dict] = field(default_factory=lambda: [{}])
|
text_encoder_extra_args: list[dict] = field(default_factory=lambda: [{}])
|
||||||
|
|
||||||
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
|
preprocess_text_funcs: tuple[Callable[[str], str] | None, ...] = field(
|
||||||
default_factory=lambda: (preprocess_text,)
|
default_factory=lambda: (None,)
|
||||||
)
|
)
|
||||||
postprocess_text_funcs: tuple[
|
postprocess_text_funcs: tuple[
|
||||||
Callable[[BaseEncoderOutput, dict], torch.Tensor], ...
|
Callable[[BaseEncoderOutput, dict], torch.Tensor], ...
|
||||||
|
|||||||
@@ -30,7 +30,6 @@ from sglang.multimodal_gen.configs.models.vaes.sana import SanaVAEConfig
|
|||||||
from sglang.multimodal_gen.configs.pipeline_configs.base import (
|
from sglang.multimodal_gen.configs.pipeline_configs.base import (
|
||||||
ModelTaskType,
|
ModelTaskType,
|
||||||
SpatialImagePipelineConfig,
|
SpatialImagePipelineConfig,
|
||||||
preprocess_text,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -65,8 +64,8 @@ class SanaPipelineConfig(SpatialImagePipelineConfig):
|
|||||||
|
|
||||||
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16",))
|
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16",))
|
||||||
|
|
||||||
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
|
preprocess_text_funcs: tuple[Callable[[str], str] | None, ...] = field(
|
||||||
default_factory=lambda: (preprocess_text,),
|
default_factory=lambda: (None,),
|
||||||
)
|
)
|
||||||
|
|
||||||
postprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
|
postprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
|
||||||
|
|||||||
@@ -236,10 +236,12 @@ class TextEncodingStage(PipelineStage):
|
|||||||
else {}
|
else {}
|
||||||
)
|
)
|
||||||
|
|
||||||
processed_text_list: list[str] = []
|
if preprocess_func is not None:
|
||||||
for prompt_str in texts:
|
processed_text_list: list[str] = [
|
||||||
preprocessed = preprocess_func(prompt_str)
|
preprocess_func(prompt_str) for prompt_str in texts
|
||||||
processed_text_list.append(preprocessed)
|
]
|
||||||
|
else:
|
||||||
|
processed_text_list = texts
|
||||||
|
|
||||||
# Prepare tokenizer args
|
# Prepare tokenizer args
|
||||||
tok_kwargs = self.prepare_tokenizer_kwargs(
|
tok_kwargs = self.prepare_tokenizer_kwargs(
|
||||||
|
|||||||
Reference in New Issue
Block a user