[diffusion] model: support ERNIE-Image (#22439)
This commit is contained in:
@@ -0,0 +1,50 @@
|
|||||||
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from typing import Tuple
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||||
|
|
||||||
|
|
||||||
|
def _is_transformer_layer(n: str, m) -> bool:
|
||||||
|
return "layers" in n and str.isdigit(n.split(".")[-1])
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class ErnieImageArchConfig(DiTArchConfig):
|
||||||
|
patch_size: int = 1
|
||||||
|
in_channels: int = 128
|
||||||
|
out_channels: int = 128
|
||||||
|
num_layers: int = 36
|
||||||
|
attention_head_dim: int = 128
|
||||||
|
num_attention_heads: int = 32
|
||||||
|
ffn_hidden_size: int = 12288
|
||||||
|
text_in_dim: int = 3072
|
||||||
|
rope_theta: int = 256
|
||||||
|
rope_axes_dim: Tuple[int, int, int] = (32, 48, 48)
|
||||||
|
eps: float = 1e-6
|
||||||
|
qk_layernorm: bool = True
|
||||||
|
|
||||||
|
stacked_params_mapping: list[tuple[str, str, str]] = field(default_factory=list)
|
||||||
|
|
||||||
|
param_names_mapping: dict = field(
|
||||||
|
default_factory=lambda: {
|
||||||
|
r"(.*)\.mlp\.gate_proj\.(.*)": (r"\1.mlp.gate_up_proj.\2", 0, 2),
|
||||||
|
r"(.*)\.mlp\.up_proj\.(.*)": (r"\1.mlp.gate_up_proj.\2", 1, 2),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
_fsdp_shard_conditions: list = field(
|
||||||
|
default_factory=lambda: [_is_transformer_layer]
|
||||||
|
)
|
||||||
|
|
||||||
|
def __post_init__(self):
|
||||||
|
super().__post_init__()
|
||||||
|
self.hidden_size = self.num_attention_heads * self.attention_head_dim
|
||||||
|
self.num_channels_latents = self.out_channels
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class ErnieImageDitConfig(DiTConfig):
|
||||||
|
arch_config: DiTArchConfig = field(default_factory=ErnieImageArchConfig)
|
||||||
|
prefix: str = "ernieimage"
|
||||||
@@ -0,0 +1,72 @@
|
|||||||
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
"""Mistral3 text encoder configuration for SGLang diffusion models."""
|
||||||
|
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.configs.models.encoders.base import (
|
||||||
|
TextEncoderArchConfig,
|
||||||
|
TextEncoderConfig,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _is_transformer_layer(n: str, m) -> bool:
|
||||||
|
return "layers" in n and str.isdigit(n.split(".")[-1])
|
||||||
|
|
||||||
|
|
||||||
|
def _is_embeddings(n: str, m) -> bool:
|
||||||
|
return n.endswith("embed_tokens")
|
||||||
|
|
||||||
|
|
||||||
|
def _is_final_norm(n: str, m) -> bool:
|
||||||
|
return n.endswith("norm")
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class Mistral3EncoderArchConfig(TextEncoderArchConfig):
|
||||||
|
"""Mistral3 text encoder architecture config for ErnieImage.
|
||||||
|
|
||||||
|
Uses Mistral3Model (vision-language model) as text encoder,
|
||||||
|
extracting the second-to-last hidden state layer.
|
||||||
|
"""
|
||||||
|
|
||||||
|
vocab_size: int = 131072
|
||||||
|
hidden_size: int = 3072
|
||||||
|
intermediate_size: int = 9216
|
||||||
|
num_hidden_layers: int = 26
|
||||||
|
num_attention_heads: int = 32
|
||||||
|
num_key_value_heads: int = 8
|
||||||
|
hidden_act: str = "silu"
|
||||||
|
max_position_embeddings: int = 262144
|
||||||
|
rms_norm_eps: float = 1e-5
|
||||||
|
pad_token_id: int = 11
|
||||||
|
bos_token_id: int = 1
|
||||||
|
eos_token_id: int = 2
|
||||||
|
tie_word_embeddings: bool = True
|
||||||
|
head_dim: int = 128
|
||||||
|
hidden_state_skip_layer: int = 2 # Use second-to-last hidden state
|
||||||
|
text_len: int = 0
|
||||||
|
|
||||||
|
stacked_params_mapping: list[tuple[str, str, str]] = field(
|
||||||
|
default_factory=lambda: [
|
||||||
|
(".qkv_proj", ".q_proj", "q"),
|
||||||
|
(".qkv_proj", ".k_proj", "k"),
|
||||||
|
(".qkv_proj", ".v_proj", "v"),
|
||||||
|
(".gate_up_proj", ".gate_proj", 0),
|
||||||
|
(".gate_up_proj", ".up_proj", 1),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
_fsdp_shard_conditions: list = field(
|
||||||
|
default_factory=lambda: [_is_transformer_layer, _is_embeddings, _is_final_norm]
|
||||||
|
)
|
||||||
|
|
||||||
|
def __post_init__(self):
|
||||||
|
# Let the parent populate tokenizer_kwargs["max_length"] = self.text_len
|
||||||
|
super().__post_init__()
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class Mistral3EncoderConfig(TextEncoderConfig):
|
||||||
|
arch_config: TextEncoderArchConfig = field(
|
||||||
|
default_factory=Mistral3EncoderArchConfig
|
||||||
|
)
|
||||||
@@ -0,0 +1,57 @@
|
|||||||
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.configs.models.vaes.base import VAEArchConfig, VAEConfig
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class ErnieImageVAEArchConfig(VAEArchConfig):
|
||||||
|
spatial_compression_ratio: int = 8
|
||||||
|
|
||||||
|
base_dim: int = 96
|
||||||
|
decoder_base_dim: int | None = None
|
||||||
|
z_dim: int = 32
|
||||||
|
dim_mult: tuple[int, ...] = (1, 2, 4, 4)
|
||||||
|
num_res_blocks: int = 2
|
||||||
|
attn_scales: tuple[float, ...] = ()
|
||||||
|
temperal_downsample: tuple[bool, ...] = (False, True, True)
|
||||||
|
dropout: float = 0.0
|
||||||
|
|
||||||
|
is_residual: bool = False
|
||||||
|
in_channels: int = 3
|
||||||
|
out_channels: int = 3
|
||||||
|
patch_size: int | None = None
|
||||||
|
scale_factor_temporal: int = 4
|
||||||
|
scale_factor_spatial: int = 8
|
||||||
|
clip_output: bool = True
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class ErnieImageVAEConfig(VAEConfig):
|
||||||
|
arch_config: ErnieImageVAEArchConfig = field(
|
||||||
|
default_factory=ErnieImageVAEArchConfig
|
||||||
|
)
|
||||||
|
|
||||||
|
use_feature_cache: bool = True
|
||||||
|
|
||||||
|
use_tiling: bool = False
|
||||||
|
use_temporal_tiling: bool = False
|
||||||
|
use_parallel_tiling: bool = False
|
||||||
|
|
||||||
|
def get_vae_scale_factor(self):
|
||||||
|
# 8 spatial compression (VAE) * 2 patch = 16 total, consistent with pipeline config
|
||||||
|
return self.arch_config.scale_factor_spatial
|
||||||
|
|
||||||
|
def __post_init__(self):
|
||||||
|
self.blend_num_frames = (
|
||||||
|
self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames
|
||||||
|
) * 2
|
||||||
|
|
||||||
|
def post_init(self):
|
||||||
|
if self.arch_config.dim_mult:
|
||||||
|
self.arch_config.vae_scale_factor = 2 ** (
|
||||||
|
len(self.arch_config.dim_mult) - 1
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
self.arch_config.vae_scale_factor = self.arch_config.scale_factor_spatial
|
||||||
|
self.arch_config.spatial_compression_ratio = self.arch_config.vae_scale_factor
|
||||||
@@ -0,0 +1,206 @@
|
|||||||
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from typing import Callable
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||||
|
from sglang.multimodal_gen.configs.models.dits.ernie_image import ErnieImageDitConfig
|
||||||
|
from sglang.multimodal_gen.configs.models.encoders.mistral3 import Mistral3EncoderConfig
|
||||||
|
from sglang.multimodal_gen.configs.models.vaes.ernie_image import ErnieImageVAEConfig
|
||||||
|
from sglang.multimodal_gen.configs.pipeline_configs.base import (
|
||||||
|
ImagePipelineConfig,
|
||||||
|
ModelTaskType,
|
||||||
|
shard_rotary_emb_for_sp,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
|
|
||||||
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def ernie_image_postprocess_text(outputs, _text_inputs, hidden_layer_index=-2):
|
||||||
|
hidden_states = outputs.hidden_states[hidden_layer_index]
|
||||||
|
return hidden_states
|
||||||
|
|
||||||
|
|
||||||
|
def _patchify_latents(latents: torch.Tensor) -> torch.Tensor:
|
||||||
|
b, c, h, w = latents.shape
|
||||||
|
latents = latents.view(b, c, h // 2, 2, w // 2, 2)
|
||||||
|
latents = latents.permute(0, 1, 3, 5, 2, 4).reshape(b, c * 4, h // 2, w // 2)
|
||||||
|
return latents
|
||||||
|
|
||||||
|
|
||||||
|
def _unpatchify_latents(latents: torch.Tensor) -> torch.Tensor:
|
||||||
|
b, c, h, w = latents.shape
|
||||||
|
latents = latents.reshape(b, c // 4, 2, 2, h, w)
|
||||||
|
latents = latents.permute(0, 1, 4, 2, 5, 3).reshape(b, c // 4, h * 2, w * 2)
|
||||||
|
return latents
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class ErnieImagePipelineConfig(ImagePipelineConfig):
|
||||||
|
"""Configuration for the ErnieImage text-to-image pipeline."""
|
||||||
|
|
||||||
|
should_use_guidance: bool = False
|
||||||
|
task_type: ModelTaskType = ModelTaskType.T2I
|
||||||
|
|
||||||
|
pe_model_max_length: int = None
|
||||||
|
|
||||||
|
vae_tiling: bool = False
|
||||||
|
vae_sp: bool = False
|
||||||
|
|
||||||
|
dit_config: DiTConfig = field(default_factory=ErnieImageDitConfig)
|
||||||
|
vae_config: VAEConfig = field(default_factory=ErnieImageVAEConfig)
|
||||||
|
|
||||||
|
enable_autocast: bool = False
|
||||||
|
|
||||||
|
text_encoder_configs: tuple[EncoderConfig, ...] = field(
|
||||||
|
default_factory=lambda: (Mistral3EncoderConfig(),)
|
||||||
|
)
|
||||||
|
|
||||||
|
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16",))
|
||||||
|
|
||||||
|
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
|
||||||
|
default_factory=lambda: (None,)
|
||||||
|
)
|
||||||
|
|
||||||
|
postprocess_text_funcs: tuple[Callable, ...] = field(
|
||||||
|
default_factory=lambda: (ernie_image_postprocess_text,)
|
||||||
|
)
|
||||||
|
|
||||||
|
text_encoder_extra_args: list[dict] = field(
|
||||||
|
default_factory=lambda: [
|
||||||
|
dict(
|
||||||
|
padding=False,
|
||||||
|
truncation=True,
|
||||||
|
max_length=None,
|
||||||
|
add_special_tokens=True,
|
||||||
|
),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
def tokenize_prompt(self, prompt: list[str], tokenizer, tok_kwargs) -> dict:
|
||||||
|
max_length = tok_kwargs.get("max_length")
|
||||||
|
if max_length is not None:
|
||||||
|
check = tokenizer(
|
||||||
|
prompt,
|
||||||
|
truncation=False,
|
||||||
|
return_tensors="pt",
|
||||||
|
add_special_tokens=tok_kwargs.get("add_special_tokens", True),
|
||||||
|
)
|
||||||
|
for i, ids in enumerate(check["input_ids"]):
|
||||||
|
if ids.shape[-1] > max_length:
|
||||||
|
logger.warning(
|
||||||
|
"Prompt #%d has %d tokens, exceeds max_length=%d. "
|
||||||
|
"The tail will be silently truncated.",
|
||||||
|
i,
|
||||||
|
ids.shape[-1],
|
||||||
|
max_length,
|
||||||
|
)
|
||||||
|
return tokenizer(prompt, **tok_kwargs)
|
||||||
|
|
||||||
|
def prepare_sigmas(self, sigmas, num_inference_steps):
|
||||||
|
return self._prepare_sigmas(sigmas, num_inference_steps)
|
||||||
|
|
||||||
|
def get_vae_scale_factor(self):
|
||||||
|
return 16
|
||||||
|
|
||||||
|
def prepare_latent_shape(self, batch, batch_size, num_frames):
|
||||||
|
vae_scale_factor = self.get_vae_scale_factor()
|
||||||
|
latent_h = batch.height // vae_scale_factor
|
||||||
|
latent_w = batch.width // vae_scale_factor
|
||||||
|
num_channels = self.dit_config.arch_config.in_channels # 128
|
||||||
|
shape = (batch_size, num_channels, latent_h, latent_w)
|
||||||
|
return shape
|
||||||
|
|
||||||
|
def maybe_pack_latents(self, latents, batch_size, batch):
|
||||||
|
return latents
|
||||||
|
|
||||||
|
def get_decode_scale_and_shift(self, device, dtype, vae):
|
||||||
|
if hasattr(vae, "bn") and vae.bn is not None:
|
||||||
|
bn_mean = vae.bn.running_mean.view(1, -1, 1, 1).to(device, dtype)
|
||||||
|
bn_var = vae.bn.running_var.view(1, -1, 1, 1).to(device, dtype)
|
||||||
|
bn_std = torch.sqrt(bn_var + 1e-5)
|
||||||
|
return 1.0 / bn_std, bn_mean
|
||||||
|
return 1.0, None
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def get_freqs_cis(img_shapes, txt_seq_lens, rotary_emb, device, dtype):
|
||||||
|
freqs = rotary_emb(img_shapes, txt_seq_lens, device=device)
|
||||||
|
|
||||||
|
if isinstance(freqs, tuple) and len(freqs) == 2:
|
||||||
|
img_freqs, txt_freqs = freqs
|
||||||
|
img_cos = img_freqs.real.to(dtype=torch.float32).contiguous()
|
||||||
|
img_sin = img_freqs.imag.to(dtype=torch.float32).contiguous()
|
||||||
|
txt_cos = txt_freqs.real.to(dtype=torch.float32).contiguous()
|
||||||
|
txt_sin = txt_freqs.imag.to(dtype=torch.float32).contiguous()
|
||||||
|
img_cache = torch.cat([img_cos, img_sin], dim=-1)
|
||||||
|
txt_cache = torch.cat([txt_cos, txt_sin], dim=-1)
|
||||||
|
return img_cache, txt_cache
|
||||||
|
|
||||||
|
cos = freqs.real.to(dtype=torch.float32).contiguous()
|
||||||
|
sin = freqs.imag.to(dtype=torch.float32).contiguous()
|
||||||
|
return torch.cat([cos, sin], dim=-1)
|
||||||
|
|
||||||
|
def _prepare_cond_kwargs(self, batch, prompt_embeds, rotary_emb, device, dtype):
|
||||||
|
batch_size = prompt_embeds[0].shape[0]
|
||||||
|
height = batch.height
|
||||||
|
width = batch.width
|
||||||
|
vae_scale_factor = self.get_vae_scale_factor()
|
||||||
|
|
||||||
|
img_shapes = [
|
||||||
|
[
|
||||||
|
(
|
||||||
|
1,
|
||||||
|
height // vae_scale_factor,
|
||||||
|
width // vae_scale_factor,
|
||||||
|
)
|
||||||
|
]
|
||||||
|
] * batch_size
|
||||||
|
txt_seq_lens = [prompt_embeds[0].shape[1]]
|
||||||
|
|
||||||
|
if rotary_emb is None:
|
||||||
|
return {
|
||||||
|
"img_shapes": img_shapes,
|
||||||
|
"txt_seq_lens": txt_seq_lens,
|
||||||
|
"freqs_cis": None,
|
||||||
|
}
|
||||||
|
|
||||||
|
freqs_cis = self.get_freqs_cis(
|
||||||
|
img_shapes, txt_seq_lens, rotary_emb, device, dtype
|
||||||
|
)
|
||||||
|
|
||||||
|
if isinstance(freqs_cis, tuple):
|
||||||
|
img_cache, txt_cache = freqs_cis
|
||||||
|
img_cache = shard_rotary_emb_for_sp(img_cache)
|
||||||
|
freqs_cis = (img_cache, txt_cache)
|
||||||
|
|
||||||
|
return {
|
||||||
|
"txt_seq_lens": txt_seq_lens,
|
||||||
|
"freqs_cis": freqs_cis,
|
||||||
|
"img_shapes": img_shapes,
|
||||||
|
}
|
||||||
|
|
||||||
|
def prepare_pos_cond_kwargs(self, batch, device, rotary_emb, dtype):
|
||||||
|
return self._prepare_cond_kwargs(
|
||||||
|
batch, batch.prompt_embeds, rotary_emb, device, dtype
|
||||||
|
)
|
||||||
|
|
||||||
|
def prepare_neg_cond_kwargs(self, batch, device, rotary_emb, dtype):
|
||||||
|
return self._prepare_cond_kwargs(
|
||||||
|
batch, batch.negative_prompt_embeds, rotary_emb, device, dtype
|
||||||
|
)
|
||||||
|
|
||||||
|
def _check_vae_has_bn(self, vae):
|
||||||
|
if not hasattr(self, "_vae_has_bn_cache"):
|
||||||
|
self._vae_has_bn_cache = hasattr(vae, "bn") and vae.bn is not None
|
||||||
|
return self._vae_has_bn_cache
|
||||||
|
|
||||||
|
def preprocess_decoding(self, latents, server_args=None, vae=None):
|
||||||
|
if vae is not None and self._check_vae_has_bn(vae):
|
||||||
|
latents = _unpatchify_latents(latents)
|
||||||
|
return latents
|
||||||
|
|
||||||
|
def post_denoising_loop(self, latents, batch):
|
||||||
|
return latents
|
||||||
@@ -0,0 +1,15 @@
|
|||||||
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
"""Sampling parameters for ErnieImage."""
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class ErnieImageSamplingParams(SamplingParams):
|
||||||
|
negative_prompt: str = " "
|
||||||
|
num_frames: int = 1
|
||||||
|
guidance_scale: float = 5.0
|
||||||
|
num_inference_steps: int = 50
|
||||||
|
use_pe: bool = True
|
||||||
@@ -197,6 +197,9 @@ class SamplingParams:
|
|||||||
return_file_paths_only: bool = True
|
return_file_paths_only: bool = True
|
||||||
enable_sequence_shard: bool | None = None
|
enable_sequence_shard: bool | None = None
|
||||||
|
|
||||||
|
# Prompt enhancement (ErnieImage)
|
||||||
|
use_pe: bool | None = None
|
||||||
|
|
||||||
def _set_output_file_ext(self):
|
def _set_output_file_ext(self):
|
||||||
# add extension if needed
|
# add extension if needed
|
||||||
if not any(
|
if not any(
|
||||||
|
|||||||
@@ -42,6 +42,9 @@ from sglang.multimodal_gen.configs.pipeline_configs import (
|
|||||||
ZImagePipelineConfig,
|
ZImagePipelineConfig,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.configs.pipeline_configs.base import PipelineConfig
|
from sglang.multimodal_gen.configs.pipeline_configs.base import PipelineConfig
|
||||||
|
from sglang.multimodal_gen.configs.pipeline_configs.ernie_image import (
|
||||||
|
ErnieImagePipelineConfig,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.configs.pipeline_configs.flux import (
|
from sglang.multimodal_gen.configs.pipeline_configs.flux import (
|
||||||
Flux2KleinPipelineConfig,
|
Flux2KleinPipelineConfig,
|
||||||
Flux2PipelineConfig,
|
Flux2PipelineConfig,
|
||||||
@@ -74,6 +77,7 @@ from sglang.multimodal_gen.configs.pipeline_configs.wan import (
|
|||||||
Wan2_2_T2V_A14B_Config,
|
Wan2_2_T2V_A14B_Config,
|
||||||
Wan2_2_TI2V_5B_Config,
|
Wan2_2_TI2V_5B_Config,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.configs.sample.ernie_image import ErnieImageSamplingParams
|
||||||
from sglang.multimodal_gen.configs.sample.flux import (
|
from sglang.multimodal_gen.configs.sample.flux import (
|
||||||
Flux2KleinSamplingParams,
|
Flux2KleinSamplingParams,
|
||||||
Flux2SamplingParams,
|
Flux2SamplingParams,
|
||||||
@@ -913,6 +917,19 @@ def _register_configs():
|
|||||||
],
|
],
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# ErnieImage
|
||||||
|
register_configs(
|
||||||
|
sampling_param_cls=ErnieImageSamplingParams,
|
||||||
|
pipeline_config_cls=ErnieImagePipelineConfig,
|
||||||
|
hf_model_paths=[
|
||||||
|
"baidu/ERNIE-Image",
|
||||||
|
"baidu/ERNIE-Image-Turbo",
|
||||||
|
],
|
||||||
|
model_detectors=[
|
||||||
|
lambda hf_id: "ernie-image" in hf_id.lower(),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
_register_configs()
|
_register_configs()
|
||||||
|
|
||||||
|
|||||||
@@ -36,6 +36,15 @@ router = APIRouter(prefix="/v1/images", tags=["images"])
|
|||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def _get_extra_field(request, field_name):
|
||||||
|
"""Get a field from model_extra, with fallback to nested extra_body dict."""
|
||||||
|
extra = request.model_extra or {}
|
||||||
|
value = extra.get(field_name)
|
||||||
|
if value is None and isinstance(extra.get("extra_body"), dict):
|
||||||
|
value = extra["extra_body"].get(field_name)
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
def _read_b64_for_paths(paths: list[str]) -> list[str]:
|
def _read_b64_for_paths(paths: list[str]) -> list[str]:
|
||||||
"""Read and base64-encode each file. Must be called before cloud upload deletes them."""
|
"""Read and base64-encode each file. Must be called before cloud upload deletes them."""
|
||||||
result = []
|
result = []
|
||||||
@@ -137,6 +146,7 @@ async def generations(
|
|||||||
upscaling_model_path=request.upscaling_model_path,
|
upscaling_model_path=request.upscaling_model_path,
|
||||||
upscaling_scale=request.upscaling_scale,
|
upscaling_scale=request.upscaling_scale,
|
||||||
perf_dump_path=request.perf_dump_path,
|
perf_dump_path=request.perf_dump_path,
|
||||||
|
use_pe=_get_extra_field(request, "use_pe"),
|
||||||
)
|
)
|
||||||
batch = prepare_request(
|
batch = prepare_request(
|
||||||
server_args=server_args,
|
server_args=server_args,
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ from abc import ABC
|
|||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from typing import Any, Dict, List, Optional, Union
|
from typing import Any, Dict, List, Optional, Union
|
||||||
|
|
||||||
from pydantic import BaseModel, Field
|
from pydantic import BaseModel, ConfigDict, Field
|
||||||
|
|
||||||
|
|
||||||
# Image API protocol models
|
# Image API protocol models
|
||||||
@@ -24,6 +24,8 @@ class ImageResponse(BaseModel):
|
|||||||
|
|
||||||
|
|
||||||
class ImageGenerationsRequest(BaseModel):
|
class ImageGenerationsRequest(BaseModel):
|
||||||
|
model_config = ConfigDict(extra="allow")
|
||||||
|
|
||||||
prompt: str
|
prompt: str
|
||||||
model: Optional[str] = None
|
model: Optional[str] = None
|
||||||
n: Optional[int] = 1
|
n: Optional[int] = 1
|
||||||
|
|||||||
@@ -0,0 +1,162 @@
|
|||||||
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from transformers import AutoModelForCausalLM, AutoTokenizer
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
|
||||||
|
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
|
||||||
|
ComponentLoader,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||||
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
|
|
||||||
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def _read_model_max_length(model_path: str) -> int | None:
|
||||||
|
"""Read model_max_length from tokenizer_config.json in the given directory."""
|
||||||
|
config_path = os.path.join(model_path, "tokenizer_config.json")
|
||||||
|
if os.path.exists(config_path):
|
||||||
|
try:
|
||||||
|
with open(config_path, encoding="utf-8") as f:
|
||||||
|
config = json.load(f)
|
||||||
|
val = config.get("model_max_length")
|
||||||
|
if val is not None:
|
||||||
|
return int(val)
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning(
|
||||||
|
"Failed to read tokenizer_config.json from %s: %s", model_path, e
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
class PEModelWrapper:
|
||||||
|
|
||||||
|
def __init__(self, model, tokenizer, device, model_max_length: int):
|
||||||
|
self.model = model
|
||||||
|
self.pe_tokenizer = tokenizer
|
||||||
|
self.device = device
|
||||||
|
self.model_max_length = model_max_length
|
||||||
|
|
||||||
|
def generate(self, prompt: str, sampling_params: dict) -> dict:
|
||||||
|
inputs = self.pe_tokenizer(
|
||||||
|
prompt,
|
||||||
|
return_tensors="pt",
|
||||||
|
truncation=True,
|
||||||
|
max_length=self.model_max_length,
|
||||||
|
).to(self.device)
|
||||||
|
|
||||||
|
input_len = inputs["input_ids"].shape[1]
|
||||||
|
|
||||||
|
generate_kwargs = dict(
|
||||||
|
**inputs,
|
||||||
|
max_new_tokens=sampling_params.get("max_new_tokens", self.model_max_length),
|
||||||
|
do_sample=True,
|
||||||
|
)
|
||||||
|
temperature = sampling_params.get("temperature")
|
||||||
|
top_p = sampling_params.get("top_p")
|
||||||
|
if temperature is not None:
|
||||||
|
generate_kwargs["temperature"] = temperature
|
||||||
|
if top_p is not None:
|
||||||
|
generate_kwargs["top_p"] = top_p
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
output_ids = self.model.generate(**generate_kwargs)
|
||||||
|
|
||||||
|
new_tokens = output_ids[0, input_len:]
|
||||||
|
text = self.pe_tokenizer.decode(new_tokens, skip_special_tokens=True)
|
||||||
|
return {"text": text}
|
||||||
|
|
||||||
|
def to(self, *args, **kwargs):
|
||||||
|
"""Move underlying model to device."""
|
||||||
|
self.model = self.model.to(*args, **kwargs)
|
||||||
|
if args:
|
||||||
|
device = args[0]
|
||||||
|
if isinstance(device, (str, torch.device)):
|
||||||
|
self.device = torch.device(device)
|
||||||
|
return self
|
||||||
|
|
||||||
|
|
||||||
|
class PELoader(ComponentLoader):
|
||||||
|
"""Loader for prompt-enhancement causal LM (Ministral-3 based)."""
|
||||||
|
|
||||||
|
component_names = ["pe"]
|
||||||
|
expected_library = "transformers"
|
||||||
|
|
||||||
|
def load_customized(
|
||||||
|
self, component_model_path: str, server_args: ServerArgs, component_name: str
|
||||||
|
):
|
||||||
|
logger.info("Loading PE model from %s ...", component_model_path)
|
||||||
|
|
||||||
|
pe_tokenizer_dir = os.path.join(
|
||||||
|
os.path.dirname(component_model_path), "pe_tokenizer"
|
||||||
|
)
|
||||||
|
if not os.path.exists(
|
||||||
|
os.path.join(component_model_path, "tokenizer_config.json")
|
||||||
|
) and os.path.exists(os.path.join(pe_tokenizer_dir, "tokenizer_config.json")):
|
||||||
|
tokenizer_path = pe_tokenizer_dir
|
||||||
|
logger.info(
|
||||||
|
"PE tokenizer files not found in %s, using %s",
|
||||||
|
component_model_path,
|
||||||
|
tokenizer_path,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
tokenizer_path = component_model_path
|
||||||
|
|
||||||
|
model_max_length = _read_model_max_length(tokenizer_path)
|
||||||
|
if model_max_length is None:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Cannot load PE model: 'model_max_length' not found in "
|
||||||
|
f"{os.path.join(tokenizer_path, 'tokenizer_config.json')}. "
|
||||||
|
"Please ensure the PE component directory (or its sibling "
|
||||||
|
"pe_tokenizer/ directory) contains a valid tokenizer_config.json "
|
||||||
|
"with a 'model_max_length' field."
|
||||||
|
)
|
||||||
|
logger.info(
|
||||||
|
"PE model_max_length=%d (from tokenizer_config.json)", model_max_length
|
||||||
|
)
|
||||||
|
|
||||||
|
tokenizer = AutoTokenizer.from_pretrained(
|
||||||
|
tokenizer_path,
|
||||||
|
trust_remote_code=server_args.trust_remote_code,
|
||||||
|
)
|
||||||
|
if tokenizer.pad_token_id is None:
|
||||||
|
tokenizer.pad_token_id = tokenizer.eos_token_id
|
||||||
|
|
||||||
|
attn_impl = "flash_attention_2"
|
||||||
|
try:
|
||||||
|
model = AutoModelForCausalLM.from_pretrained(
|
||||||
|
component_model_path,
|
||||||
|
torch_dtype=torch.bfloat16,
|
||||||
|
trust_remote_code=server_args.trust_remote_code,
|
||||||
|
attn_implementation=attn_impl,
|
||||||
|
)
|
||||||
|
logger.info("PE model: using Flash Attention 2")
|
||||||
|
except (ValueError, ImportError):
|
||||||
|
logger.warning("Flash Attention 2 not available, falling back to SDPA")
|
||||||
|
attn_impl = "sdpa"
|
||||||
|
model = AutoModelForCausalLM.from_pretrained(
|
||||||
|
component_model_path,
|
||||||
|
torch_dtype=torch.bfloat16,
|
||||||
|
trust_remote_code=server_args.trust_remote_code,
|
||||||
|
attn_implementation=attn_impl,
|
||||||
|
)
|
||||||
|
|
||||||
|
device = get_local_torch_device()
|
||||||
|
model = model.to(device).eval()
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
"PE model loaded on %s: %s (attn=%s)",
|
||||||
|
device,
|
||||||
|
model.__class__.__name__,
|
||||||
|
attn_impl,
|
||||||
|
)
|
||||||
|
|
||||||
|
return PEModelWrapper(
|
||||||
|
model=model,
|
||||||
|
tokenizer=tokenizer,
|
||||||
|
device=device,
|
||||||
|
model_max_length=model_max_length,
|
||||||
|
)
|
||||||
@@ -0,0 +1,477 @@
|
|||||||
|
# Copyright 2026 Baidu ERNIE-Image Team and The HuggingFace Team. All rights reserved.
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
from typing import Any, Optional, Tuple
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import torch.nn.functional as F
|
||||||
|
from diffusers.models.embeddings import TimestepEmbedding, Timesteps
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.configs.models.dits.ernie_image import (
|
||||||
|
ErnieImageDitConfig,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.distributed import (
|
||||||
|
get_tp_world_size,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.layers.attention.layer import USPAttention
|
||||||
|
from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm, apply_qk_norm
|
||||||
|
from sglang.multimodal_gen.runtime.layers.linear import (
|
||||||
|
ColumnParallelLinear,
|
||||||
|
MergedColumnParallelLinear,
|
||||||
|
RowParallelLinear,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.layers.quantization import QuantizationConfig
|
||||||
|
from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT
|
||||||
|
from sglang.multimodal_gen.runtime.utils.layerwise_offload import OffloadableDiTMixin
|
||||||
|
|
||||||
|
|
||||||
|
def _rope(pos: torch.Tensor, dim: int, theta: int) -> torch.Tensor:
|
||||||
|
assert dim % 2 == 0
|
||||||
|
scale = torch.arange(0, dim, 2, dtype=torch.float64, device=pos.device) / dim
|
||||||
|
omega = 1.0 / (theta**scale)
|
||||||
|
out = torch.einsum("...n,d->...nd", pos, omega) # codespell:ignore nd
|
||||||
|
return out.float()
|
||||||
|
|
||||||
|
|
||||||
|
class EmbedND3(nn.Module):
|
||||||
|
"""3D rotary positional embedding for (temporal/batch_idx, height, width)."""
|
||||||
|
|
||||||
|
def __init__(self, dim: int, theta: int, axes_dim: Tuple[int, int, int]):
|
||||||
|
super().__init__()
|
||||||
|
self.dim = dim
|
||||||
|
self.theta = theta
|
||||||
|
self.axes_dim = list(axes_dim)
|
||||||
|
|
||||||
|
def forward(self, ids: torch.Tensor) -> torch.Tensor:
|
||||||
|
emb = torch.cat(
|
||||||
|
[_rope(ids[..., i], self.axes_dim[i], self.theta) for i in range(3)],
|
||||||
|
dim=-1,
|
||||||
|
)
|
||||||
|
emb = emb.unsqueeze(1).permute(2, 0, 1, 3)
|
||||||
|
return torch.stack([emb, emb], dim=-1).reshape(*emb.shape[:-1], -1)
|
||||||
|
|
||||||
|
|
||||||
|
class ErnieImageSelfAttention(nn.Module):
|
||||||
|
"""Self-attention with separate Q/K/V projections and QK LayerNorm.
|
||||||
|
|
||||||
|
Module name hierarchy matches diffusers Attention naming convention:
|
||||||
|
self_attention.to_q, self_attention.to_k, self_attention.to_v,
|
||||||
|
self_attention.to_out.0, self_attention.norm_q, self_attention.norm_k.
|
||||||
|
|
||||||
|
Supports tensor parallelism: Q/K/V projections use ColumnParallelLinear
|
||||||
|
(output dim sharded by heads), output projection uses RowParallelLinear
|
||||||
|
(input dim sharded, all-reduce after matmul).
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
hidden_size: int,
|
||||||
|
num_heads: int,
|
||||||
|
head_dim: int,
|
||||||
|
eps: float = 1e-6,
|
||||||
|
qk_layernorm: bool = True,
|
||||||
|
prefix: str = "",
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
self.num_heads = num_heads
|
||||||
|
self.head_dim = head_dim
|
||||||
|
|
||||||
|
tp_size = get_tp_world_size()
|
||||||
|
self.num_local_heads = num_heads // tp_size
|
||||||
|
assert (
|
||||||
|
num_heads % tp_size == 0
|
||||||
|
), f"num_heads ({num_heads}) must be divisible by tp_size ({tp_size})"
|
||||||
|
|
||||||
|
self.to_q = ColumnParallelLinear(
|
||||||
|
hidden_size,
|
||||||
|
hidden_size,
|
||||||
|
bias=False,
|
||||||
|
gather_output=False,
|
||||||
|
prefix=f"{prefix}.to_q",
|
||||||
|
)
|
||||||
|
self.to_k = ColumnParallelLinear(
|
||||||
|
hidden_size,
|
||||||
|
hidden_size,
|
||||||
|
bias=False,
|
||||||
|
gather_output=False,
|
||||||
|
prefix=f"{prefix}.to_k",
|
||||||
|
)
|
||||||
|
self.to_v = ColumnParallelLinear(
|
||||||
|
hidden_size,
|
||||||
|
hidden_size,
|
||||||
|
bias=False,
|
||||||
|
gather_output=False,
|
||||||
|
prefix=f"{prefix}.to_v",
|
||||||
|
)
|
||||||
|
self.to_out = nn.ModuleList(
|
||||||
|
[
|
||||||
|
RowParallelLinear(
|
||||||
|
hidden_size,
|
||||||
|
hidden_size,
|
||||||
|
bias=False,
|
||||||
|
input_is_parallel=True,
|
||||||
|
prefix=f"{prefix}.to_out.0",
|
||||||
|
),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
self.qk_layernorm = qk_layernorm
|
||||||
|
if qk_layernorm:
|
||||||
|
self.norm_q = RMSNorm(head_dim, eps=eps)
|
||||||
|
self.norm_k = RMSNorm(head_dim, eps=eps)
|
||||||
|
|
||||||
|
self.attn = USPAttention(
|
||||||
|
num_heads=self.num_local_heads,
|
||||||
|
head_size=head_dim,
|
||||||
|
prefix=f"{prefix}.attn",
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
x: torch.Tensor,
|
||||||
|
rotary_pos_emb: torch.Tensor,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
B, S, H = x.shape
|
||||||
|
|
||||||
|
q, _ = self.to_q(x)
|
||||||
|
k, _ = self.to_k(x)
|
||||||
|
v, _ = self.to_v(x)
|
||||||
|
|
||||||
|
q = q.view(B, S, self.num_local_heads, self.head_dim)
|
||||||
|
k = k.view(B, S, self.num_local_heads, self.head_dim)
|
||||||
|
v = v.view(B, S, self.num_local_heads, self.head_dim)
|
||||||
|
|
||||||
|
if self.qk_layernorm:
|
||||||
|
q, k = apply_qk_norm(
|
||||||
|
q,
|
||||||
|
k,
|
||||||
|
self.norm_q,
|
||||||
|
self.norm_k,
|
||||||
|
self.head_dim,
|
||||||
|
)
|
||||||
|
|
||||||
|
q = _apply_rotary_bshd(q, rotary_pos_emb)
|
||||||
|
k = _apply_rotary_bshd(k, rotary_pos_emb)
|
||||||
|
|
||||||
|
attn_out = self.attn(q, k, v)
|
||||||
|
attn_out = attn_out.reshape(B, S, self.num_local_heads * self.head_dim)
|
||||||
|
out, _ = self.to_out[0](attn_out)
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
class ErnieImageMLP(nn.Module):
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
hidden_size: int,
|
||||||
|
ffn_hidden_size: int,
|
||||||
|
prefix: str = "",
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
self.gate_up_proj = MergedColumnParallelLinear(
|
||||||
|
hidden_size,
|
||||||
|
[ffn_hidden_size, ffn_hidden_size],
|
||||||
|
bias=False,
|
||||||
|
gather_output=False,
|
||||||
|
prefix=f"{prefix}.gate_up_proj",
|
||||||
|
)
|
||||||
|
self.linear_fc2 = RowParallelLinear(
|
||||||
|
ffn_hidden_size,
|
||||||
|
hidden_size,
|
||||||
|
bias=False,
|
||||||
|
input_is_parallel=True,
|
||||||
|
prefix=f"{prefix}.linear_fc2",
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
gate_up, _ = self.gate_up_proj(x)
|
||||||
|
gate, up = gate_up.chunk(2, dim=-1)
|
||||||
|
x = up * F.gelu(gate)
|
||||||
|
x, _ = self.linear_fc2(x)
|
||||||
|
return x
|
||||||
|
|
||||||
|
|
||||||
|
class ErnieImageSharedAdaLNBlock(nn.Module):
|
||||||
|
"""Single-stream transformer block with externally-computed Shared AdaLN."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
hidden_size: int,
|
||||||
|
num_heads: int,
|
||||||
|
head_dim: int,
|
||||||
|
ffn_hidden_size: int,
|
||||||
|
eps: float = 1e-6,
|
||||||
|
qk_layernorm: bool = True,
|
||||||
|
prefix: str = "",
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
self.adaLN_sa_ln = RMSNorm(hidden_size, eps=eps)
|
||||||
|
self.self_attention = ErnieImageSelfAttention(
|
||||||
|
hidden_size,
|
||||||
|
num_heads,
|
||||||
|
head_dim,
|
||||||
|
eps,
|
||||||
|
qk_layernorm,
|
||||||
|
prefix=f"{prefix}.self_attention",
|
||||||
|
)
|
||||||
|
self.adaLN_mlp_ln = RMSNorm(hidden_size, eps=eps)
|
||||||
|
self.mlp = ErnieImageMLP(hidden_size, ffn_hidden_size, prefix=f"{prefix}.mlp")
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
x: torch.Tensor,
|
||||||
|
rotary_pos_emb: torch.Tensor,
|
||||||
|
shift_msa: torch.Tensor,
|
||||||
|
scale_msa: torch.Tensor,
|
||||||
|
gate_msa: torch.Tensor,
|
||||||
|
shift_mlp: torch.Tensor,
|
||||||
|
scale_mlp: torch.Tensor,
|
||||||
|
gate_mlp: torch.Tensor,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
residual = x
|
||||||
|
x = self.adaLN_sa_ln(x) * (1 + scale_msa) + shift_msa
|
||||||
|
x = residual + gate_msa * self.self_attention(x, rotary_pos_emb)
|
||||||
|
|
||||||
|
residual = x
|
||||||
|
x = self.adaLN_mlp_ln(x) * (1 + scale_mlp) + shift_mlp
|
||||||
|
x = residual + gate_mlp * self.mlp(x)
|
||||||
|
|
||||||
|
return x
|
||||||
|
|
||||||
|
|
||||||
|
def _apply_rotary_bshd(x: torch.Tensor, freqs: torch.Tensor) -> torch.Tensor:
|
||||||
|
freqs = freqs.permute(1, 0, 2, 3)
|
||||||
|
rot_dim = freqs.shape[-1]
|
||||||
|
x_rot, x_pass = x[..., :rot_dim], x[..., rot_dim:]
|
||||||
|
|
||||||
|
cos_ = torch.cos(freqs).to(x.dtype)
|
||||||
|
sin_ = torch.sin(freqs).to(x.dtype)
|
||||||
|
|
||||||
|
x1, x2 = x_rot.chunk(2, dim=-1)
|
||||||
|
x_rotated = torch.cat((-x2, x1), dim=-1)
|
||||||
|
|
||||||
|
x_rot = x_rot * cos_ + x_rotated * sin_
|
||||||
|
return torch.cat((x_rot, x_pass), dim=-1)
|
||||||
|
|
||||||
|
|
||||||
|
class ErnieImageTransformer2DModel(CachableDiT, OffloadableDiTMixin):
|
||||||
|
"""ErnieImage DiT: Single-stream transformer with Shared AdaLN."""
|
||||||
|
|
||||||
|
_supports_gradient_checkpointing = True
|
||||||
|
_no_split_modules = ["ErnieImageSharedAdaLNBlock"]
|
||||||
|
_skip_layerwise_casting_patterns = ["pos_embed", "norm"]
|
||||||
|
|
||||||
|
_fsdp_shard_conditions = ErnieImageDitConfig().arch_config._fsdp_shard_conditions
|
||||||
|
_compile_conditions = []
|
||||||
|
param_names_mapping = ErnieImageDitConfig().arch_config.param_names_mapping
|
||||||
|
reverse_param_names_mapping = {}
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
config: ErnieImageDitConfig,
|
||||||
|
hf_config: dict[str, Any],
|
||||||
|
quant_config: Optional[QuantizationConfig] = None,
|
||||||
|
):
|
||||||
|
super().__init__(config=config, hf_config=hf_config)
|
||||||
|
|
||||||
|
arch = config.arch_config
|
||||||
|
self.hidden_size = arch.hidden_size
|
||||||
|
self.num_attention_heads = arch.num_attention_heads
|
||||||
|
self.num_channels_latents = arch.out_channels
|
||||||
|
self.head_dim = arch.attention_head_dim
|
||||||
|
self.num_layers = arch.num_layers
|
||||||
|
self.patch_size = arch.patch_size
|
||||||
|
self.out_channels = arch.out_channels
|
||||||
|
self.inner_dim = self.hidden_size
|
||||||
|
|
||||||
|
tp_size = get_tp_world_size()
|
||||||
|
|
||||||
|
self.x_embedder = nn.ModuleDict(
|
||||||
|
{
|
||||||
|
"proj": nn.Conv2d(
|
||||||
|
arch.in_channels,
|
||||||
|
self.inner_dim,
|
||||||
|
kernel_size=arch.patch_size,
|
||||||
|
stride=arch.patch_size,
|
||||||
|
bias=True,
|
||||||
|
),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
if arch.text_in_dim != self.inner_dim:
|
||||||
|
self.text_proj = nn.Linear(arch.text_in_dim, self.inner_dim, bias=False)
|
||||||
|
else:
|
||||||
|
self.text_proj = None
|
||||||
|
|
||||||
|
self.time_proj = Timesteps(
|
||||||
|
self.inner_dim,
|
||||||
|
flip_sin_to_cos=False,
|
||||||
|
downscale_freq_shift=0,
|
||||||
|
)
|
||||||
|
self.time_embedding = TimestepEmbedding(
|
||||||
|
in_channels=self.inner_dim,
|
||||||
|
time_embed_dim=self.inner_dim,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.pos_embed = EmbedND3(
|
||||||
|
dim=self.head_dim,
|
||||||
|
theta=arch.rope_theta,
|
||||||
|
axes_dim=arch.rope_axes_dim,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.adaLN_modulation = nn.Sequential(
|
||||||
|
nn.SiLU(),
|
||||||
|
nn.Linear(self.inner_dim, 6 * self.inner_dim),
|
||||||
|
)
|
||||||
|
|
||||||
|
self.layers = nn.ModuleList(
|
||||||
|
[
|
||||||
|
ErnieImageSharedAdaLNBlock(
|
||||||
|
hidden_size=self.inner_dim,
|
||||||
|
num_heads=self.num_attention_heads,
|
||||||
|
head_dim=self.head_dim,
|
||||||
|
ffn_hidden_size=arch.ffn_hidden_size,
|
||||||
|
eps=arch.eps,
|
||||||
|
qk_layernorm=arch.qk_layernorm,
|
||||||
|
prefix=f"layers.{i}",
|
||||||
|
)
|
||||||
|
for i in range(self.num_layers)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
self.final_norm = nn.ModuleDict(
|
||||||
|
{
|
||||||
|
"norm": nn.LayerNorm(
|
||||||
|
self.inner_dim, elementwise_affine=False, eps=arch.eps
|
||||||
|
),
|
||||||
|
"linear": nn.Linear(self.inner_dim, self.inner_dim * 2),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
self.final_linear = ColumnParallelLinear(
|
||||||
|
self.inner_dim,
|
||||||
|
arch.patch_size * arch.patch_size * self.out_channels,
|
||||||
|
bias=True,
|
||||||
|
gather_output=True,
|
||||||
|
prefix="final_linear",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.layer_names = ["layers"]
|
||||||
|
|
||||||
|
self.__post_init__()
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
hidden_states: torch.Tensor,
|
||||||
|
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
|
||||||
|
timestep: torch.LongTensor,
|
||||||
|
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor] | None = None,
|
||||||
|
guidance=None,
|
||||||
|
**kwargs,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""
|
||||||
|
Args:
|
||||||
|
hidden_states: [B, C, H, W] latent images (patchified, 128 channels)
|
||||||
|
encoder_hidden_states: [B, T, text_dim] or list of text embeddings
|
||||||
|
timestep: [B] timestep values
|
||||||
|
Returns:
|
||||||
|
output: [B, C, H, W] predicted noise / denoised output
|
||||||
|
"""
|
||||||
|
device, dtype = hidden_states.device, hidden_states.dtype
|
||||||
|
B, C, H, W = hidden_states.shape
|
||||||
|
p = self.patch_size
|
||||||
|
Hp, Wp = H // p, W // p
|
||||||
|
N_img = Hp * Wp
|
||||||
|
|
||||||
|
img_tokens = self.x_embedder["proj"](hidden_states) # [B, D, Hp, Wp]
|
||||||
|
img_tokens = img_tokens.reshape(B, self.inner_dim, N_img).transpose(
|
||||||
|
1, 2
|
||||||
|
) # [B, N_img, D]
|
||||||
|
|
||||||
|
if isinstance(encoder_hidden_states, (list, tuple)):
|
||||||
|
encoder_hidden_states = encoder_hidden_states[0]
|
||||||
|
text_tokens = encoder_hidden_states # [B, T, text_dim]
|
||||||
|
if self.text_proj is not None and text_tokens.numel() > 0:
|
||||||
|
text_tokens = self.text_proj(text_tokens)
|
||||||
|
Tmax = text_tokens.shape[1]
|
||||||
|
|
||||||
|
x = torch.cat([img_tokens, text_tokens], dim=1) # [B, S, D]
|
||||||
|
|
||||||
|
grid_yx = torch.stack(
|
||||||
|
torch.meshgrid(
|
||||||
|
torch.arange(Hp, device=device, dtype=torch.float32),
|
||||||
|
torch.arange(Wp, device=device, dtype=torch.float32),
|
||||||
|
indexing="ij",
|
||||||
|
),
|
||||||
|
dim=-1,
|
||||||
|
).reshape(-1, 2)
|
||||||
|
|
||||||
|
image_ids = torch.cat(
|
||||||
|
[
|
||||||
|
torch.full((B, N_img, 1), Tmax, device=device, dtype=torch.float32),
|
||||||
|
grid_yx.view(1, N_img, 2).expand(B, -1, -1),
|
||||||
|
],
|
||||||
|
dim=-1,
|
||||||
|
)
|
||||||
|
|
||||||
|
if Tmax > 0:
|
||||||
|
text_ids = torch.cat(
|
||||||
|
[
|
||||||
|
torch.arange(Tmax, device=device, dtype=torch.float32)
|
||||||
|
.view(1, Tmax, 1)
|
||||||
|
.expand(B, -1, -1),
|
||||||
|
torch.zeros((B, Tmax, 2), device=device),
|
||||||
|
],
|
||||||
|
dim=-1,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
text_ids = torch.zeros((B, 0, 3), device=device)
|
||||||
|
|
||||||
|
all_ids = torch.cat([image_ids, text_ids], dim=1)
|
||||||
|
rotary_pos_emb = self.pos_embed(all_ids)
|
||||||
|
|
||||||
|
t_emb = self.time_proj(timestep.to(dtype))
|
||||||
|
c = self.time_embedding(t_emb.to(dtype=dtype))
|
||||||
|
|
||||||
|
mod_params = self.adaLN_modulation(c)
|
||||||
|
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
|
||||||
|
t.unsqueeze(1) for t in mod_params.chunk(6, dim=-1)
|
||||||
|
)
|
||||||
|
|
||||||
|
for layer in self.layers:
|
||||||
|
x = layer(
|
||||||
|
x,
|
||||||
|
rotary_pos_emb,
|
||||||
|
shift_msa,
|
||||||
|
scale_msa,
|
||||||
|
gate_msa,
|
||||||
|
shift_mlp,
|
||||||
|
scale_mlp,
|
||||||
|
gate_mlp,
|
||||||
|
)
|
||||||
|
|
||||||
|
scale, shift = self.final_norm["linear"](c).chunk(2, dim=-1)
|
||||||
|
x = self.final_norm["norm"](x) * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)
|
||||||
|
|
||||||
|
patches, _ = self.final_linear(x[:, :N_img, :])
|
||||||
|
|
||||||
|
output = patches.view(B, Hp, Wp, p, p, self.out_channels)
|
||||||
|
output = output.permute(0, 5, 1, 3, 2, 4).contiguous()
|
||||||
|
output = output.view(B, self.out_channels, H, W)
|
||||||
|
|
||||||
|
return output
|
||||||
|
|
||||||
|
|
||||||
|
EntryClass = ErnieImageTransformer2DModel
|
||||||
@@ -0,0 +1,232 @@
|
|||||||
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
"""ErnieImage text-to-image pipeline."""
|
||||||
|
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
|
||||||
|
ComposedPipelineBase,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline
|
||||||
|
from sglang.multimodal_gen.runtime.pipelines_core.stages.input_validation import (
|
||||||
|
InputValidationStage,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.ernie_image_pe import (
|
||||||
|
PromptEnhancementStage,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.pipelines_core.stages.text_encoding import (
|
||||||
|
TextEncodingStage,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
|
||||||
|
maybe_download_model,
|
||||||
|
maybe_download_model_index,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
|
|
||||||
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class ErnieImagePipeline(LoRAPipeline, ComposedPipelineBase):
|
||||||
|
|
||||||
|
pipeline_name = "ErnieImagePipeline"
|
||||||
|
|
||||||
|
_required_config_modules = [
|
||||||
|
"text_encoder",
|
||||||
|
"tokenizer",
|
||||||
|
"vae",
|
||||||
|
"transformer",
|
||||||
|
"scheduler",
|
||||||
|
]
|
||||||
|
|
||||||
|
def _has_pe_in_model_index(self, server_args) -> bool:
|
||||||
|
try:
|
||||||
|
model_index = maybe_download_model_index(server_args.model_path)
|
||||||
|
return "pe" in model_index and model_index["pe"] is not None
|
||||||
|
except Exception:
|
||||||
|
return False
|
||||||
|
|
||||||
|
def _read_tokenizer_model_max_length(self, model_path: str):
|
||||||
|
"""Read model_max_length from tokenizer/tokenizer_config.json.
|
||||||
|
|
||||||
|
Supports both local paths and HuggingFace Hub model IDs.
|
||||||
|
Returns None if the value cannot be determined.
|
||||||
|
"""
|
||||||
|
tokenizer_config_subpath = os.path.join("tokenizer", "tokenizer_config.json")
|
||||||
|
|
||||||
|
# Local path
|
||||||
|
if os.path.exists(model_path):
|
||||||
|
config_path = os.path.join(model_path, tokenizer_config_subpath)
|
||||||
|
if os.path.exists(config_path):
|
||||||
|
with open(config_path, encoding="utf-8") as f:
|
||||||
|
config = json.load(f)
|
||||||
|
return config.get("model_max_length")
|
||||||
|
return None
|
||||||
|
|
||||||
|
# Remote HuggingFace Hub model ID
|
||||||
|
try:
|
||||||
|
import tempfile
|
||||||
|
|
||||||
|
from huggingface_hub import hf_hub_download
|
||||||
|
|
||||||
|
with tempfile.TemporaryDirectory() as tmp_dir:
|
||||||
|
config_path = hf_hub_download(
|
||||||
|
repo_id=model_path,
|
||||||
|
filename=tokenizer_config_subpath,
|
||||||
|
local_dir=tmp_dir,
|
||||||
|
)
|
||||||
|
with open(config_path, encoding="utf-8") as f:
|
||||||
|
config = json.load(f)
|
||||||
|
return config.get("model_max_length")
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning(
|
||||||
|
"Failed to read tokenizer_config.json from %s: %s", model_path, e
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
|
||||||
|
def _resolve_pe_tokenizer_path(self, model_path: str, server_args) -> str:
|
||||||
|
"""Resolve the directory that contains the PE tokenizer files."""
|
||||||
|
pe_component_path = server_args.component_paths.get(
|
||||||
|
"pe", os.path.join(model_path, "pe")
|
||||||
|
)
|
||||||
|
if os.path.exists(os.path.join(pe_component_path, "tokenizer_config.json")):
|
||||||
|
return pe_component_path
|
||||||
|
pe_tokenizer_dir = os.path.join(model_path, "pe_tokenizer")
|
||||||
|
if os.path.exists(os.path.join(pe_tokenizer_dir, "tokenizer_config.json")):
|
||||||
|
return pe_tokenizer_dir
|
||||||
|
return pe_component_path
|
||||||
|
|
||||||
|
def _read_pe_model_max_length(self, model_path: str, server_args) -> int | None:
|
||||||
|
# If model_path is a Hub ID, download the full model first (or use cache)
|
||||||
|
# so that pe/tokenizer_config.json is available locally.
|
||||||
|
if not os.path.exists(model_path):
|
||||||
|
try:
|
||||||
|
model_path = maybe_download_model(
|
||||||
|
model_path, force_diffusers_model=True
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning(
|
||||||
|
"Failed to download model to read pe/tokenizer_config.json: %s", e
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
|
||||||
|
tokenizer_path = self._resolve_pe_tokenizer_path(model_path, server_args)
|
||||||
|
config_path = os.path.join(tokenizer_path, "tokenizer_config.json")
|
||||||
|
if os.path.exists(config_path):
|
||||||
|
try:
|
||||||
|
with open(config_path, encoding="utf-8") as f:
|
||||||
|
config = json.load(f)
|
||||||
|
val = config.get("model_max_length")
|
||||||
|
if val is not None:
|
||||||
|
return int(val)
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning(
|
||||||
|
"Failed to read tokenizer_config.json from %s: %s",
|
||||||
|
tokenizer_path,
|
||||||
|
e,
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
|
||||||
|
def load_modules(self, server_args, loaded_modules=None):
|
||||||
|
has_pe = self._has_pe_in_model_index(server_args)
|
||||||
|
if has_pe:
|
||||||
|
if "pe" not in self._required_config_modules:
|
||||||
|
self._required_config_modules.insert(0, "pe")
|
||||||
|
logger.info("PE model detected in model_index.json, will load PE module.")
|
||||||
|
|
||||||
|
pipeline_config = server_args.pipeline_config
|
||||||
|
|
||||||
|
# --- Text encoder max_length ---
|
||||||
|
text_model_max_length = self._read_tokenizer_model_max_length(
|
||||||
|
server_args.model_path
|
||||||
|
)
|
||||||
|
if text_model_max_length is not None:
|
||||||
|
# 1. Update arch_config.text_len so the model knows the true sequence length
|
||||||
|
if (
|
||||||
|
hasattr(pipeline_config, "text_encoder_configs")
|
||||||
|
and pipeline_config.text_encoder_configs
|
||||||
|
):
|
||||||
|
arch_config = pipeline_config.text_encoder_configs[0].arch_config
|
||||||
|
arch_config.text_len = text_model_max_length
|
||||||
|
arch_config.tokenizer_kwargs["max_length"] = text_model_max_length
|
||||||
|
# 2. Update text_encoder_extra_args used by TextEncodingStage tokenization
|
||||||
|
if (
|
||||||
|
hasattr(pipeline_config, "text_encoder_extra_args")
|
||||||
|
and pipeline_config.text_encoder_extra_args
|
||||||
|
):
|
||||||
|
pipeline_config.text_encoder_extra_args[0][
|
||||||
|
"max_length"
|
||||||
|
] = text_model_max_length
|
||||||
|
logger.info(
|
||||||
|
"Set text encoder model_max_length=%d from tokenizer/tokenizer_config.json",
|
||||||
|
text_model_max_length,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
logger.warning(
|
||||||
|
"Could not read model_max_length from tokenizer/tokenizer_config.json, "
|
||||||
|
"text encoder will use the default text_len from arch config."
|
||||||
|
)
|
||||||
|
|
||||||
|
# --- PE model_max_length ---
|
||||||
|
if has_pe:
|
||||||
|
pe_model_max_length = self._read_pe_model_max_length(
|
||||||
|
server_args.model_path, server_args
|
||||||
|
)
|
||||||
|
if pe_model_max_length is not None:
|
||||||
|
pipeline_config.pe_model_max_length = pe_model_max_length
|
||||||
|
logger.info(
|
||||||
|
"Set PE model_max_length=%d from pe/tokenizer_config.json",
|
||||||
|
pe_model_max_length,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
raise RuntimeError(
|
||||||
|
"PE model is present but 'model_max_length' could not be read from "
|
||||||
|
"pe/tokenizer_config.json. Please ensure the PE component directory "
|
||||||
|
"contains a valid tokenizer_config.json with a 'model_max_length' field."
|
||||||
|
)
|
||||||
|
|
||||||
|
return super().load_modules(server_args, loaded_modules)
|
||||||
|
|
||||||
|
def create_pipeline_stages(self, server_args):
|
||||||
|
self.add_stage(InputValidationStage())
|
||||||
|
|
||||||
|
pe_model = self.get_module("pe")
|
||||||
|
if pe_model is not None:
|
||||||
|
pe_tokenizer = getattr(pe_model, "pe_tokenizer", None)
|
||||||
|
if pe_tokenizer is None:
|
||||||
|
from transformers import AutoTokenizer
|
||||||
|
|
||||||
|
pe_tokenizer_path = self._resolve_pe_tokenizer_path(
|
||||||
|
self.model_path, server_args
|
||||||
|
)
|
||||||
|
logger.warning(
|
||||||
|
"pe_tokenizer not found on pe_model (%s), loading from %s",
|
||||||
|
type(pe_model).__name__,
|
||||||
|
pe_tokenizer_path,
|
||||||
|
)
|
||||||
|
pe_tokenizer = AutoTokenizer.from_pretrained(
|
||||||
|
pe_tokenizer_path,
|
||||||
|
trust_remote_code=server_args.trust_remote_code,
|
||||||
|
)
|
||||||
|
self.add_stage(
|
||||||
|
PromptEnhancementStage(
|
||||||
|
pe_model=pe_model,
|
||||||
|
pe_tokenizer=pe_tokenizer,
|
||||||
|
),
|
||||||
|
"prompt_enhancement_stage",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.add_stage(
|
||||||
|
TextEncodingStage(
|
||||||
|
text_encoders=[self.get_module("text_encoder")],
|
||||||
|
tokenizers=[self.get_module("tokenizer")],
|
||||||
|
),
|
||||||
|
"prompt_encoding_stage_primary",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.add_standard_timestep_preparation_stage()
|
||||||
|
self.add_standard_latent_preparation_stage()
|
||||||
|
self.add_standard_denoising_stage()
|
||||||
|
self.add_standard_decoding_stage()
|
||||||
|
|
||||||
|
|
||||||
|
EntryClass = ErnieImagePipeline
|
||||||
+98
@@ -0,0 +1,98 @@
|
|||||||
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
"""
|
||||||
|
Prompt enhancement stage for ErnieImage pipeline.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import json
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
|
||||||
|
from sglang.multimodal_gen.runtime.pipelines_core.stages.base import PipelineStage
|
||||||
|
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||||
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
|
|
||||||
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class PromptEnhancementStage(PipelineStage):
|
||||||
|
|
||||||
|
def __init__(self, pe_model, pe_tokenizer):
|
||||||
|
super().__init__()
|
||||||
|
self.pe_model = pe_model
|
||||||
|
self.pe_tokenizer = pe_tokenizer
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def forward(self, batch: Req, server_args: ServerArgs) -> Req:
|
||||||
|
# Skip if use_pe is disabled or tokenizer unavailable
|
||||||
|
use_pe = getattr(batch, "use_pe", True)
|
||||||
|
if not use_pe or self.pe_model is None:
|
||||||
|
return batch
|
||||||
|
|
||||||
|
if self.pe_tokenizer is None:
|
||||||
|
logger.warning(
|
||||||
|
"pe_tokenizer is None, skipping prompt enhancement. "
|
||||||
|
"Check PE model loading logs for errors."
|
||||||
|
)
|
||||||
|
return batch
|
||||||
|
|
||||||
|
# Read max_new_tokens from pipeline config (injected from tokenizer_config.json at load time)
|
||||||
|
max_new_tokens = server_args.pipeline_config.pe_model_max_length
|
||||||
|
|
||||||
|
prompt = batch.prompt
|
||||||
|
if isinstance(prompt, str):
|
||||||
|
prompts = [prompt]
|
||||||
|
else:
|
||||||
|
prompts = list(prompt)
|
||||||
|
|
||||||
|
height = getattr(batch, "height", 1024)
|
||||||
|
width = getattr(batch, "width", 1024)
|
||||||
|
|
||||||
|
enhanced = []
|
||||||
|
for p in prompts:
|
||||||
|
enhanced_p = self._enhance_single_prompt(
|
||||||
|
p, width, height, max_new_tokens=max_new_tokens
|
||||||
|
)
|
||||||
|
enhanced.append(enhanced_p)
|
||||||
|
|
||||||
|
if isinstance(batch.prompt, str):
|
||||||
|
batch.prompt = enhanced[0]
|
||||||
|
else:
|
||||||
|
batch.prompt = enhanced
|
||||||
|
|
||||||
|
logger.info("PE enhanced prompt: %s", batch.prompt)
|
||||||
|
return batch
|
||||||
|
|
||||||
|
def _enhance_single_prompt(
|
||||||
|
self,
|
||||||
|
prompt: str,
|
||||||
|
width: int,
|
||||||
|
height: int,
|
||||||
|
max_new_tokens: int,
|
||||||
|
temperature: float = None,
|
||||||
|
top_p: float = None,
|
||||||
|
) -> str:
|
||||||
|
user_content = json.dumps(
|
||||||
|
{"prompt": prompt, "width": width, "height": height},
|
||||||
|
ensure_ascii=False,
|
||||||
|
)
|
||||||
|
messages = [{"role": "user", "content": user_content}]
|
||||||
|
|
||||||
|
input_text = self.pe_tokenizer.apply_chat_template(
|
||||||
|
messages,
|
||||||
|
tokenize=False,
|
||||||
|
add_generation_prompt=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
sampling_params = {"max_new_tokens": max_new_tokens}
|
||||||
|
if temperature is not None:
|
||||||
|
sampling_params["temperature"] = temperature
|
||||||
|
if top_p is not None:
|
||||||
|
sampling_params["top_p"] = top_p
|
||||||
|
|
||||||
|
output = self.pe_model.generate(
|
||||||
|
prompt=input_text,
|
||||||
|
sampling_params=sampling_params,
|
||||||
|
)
|
||||||
|
|
||||||
|
return output["text"].strip()
|
||||||
Reference in New Issue
Block a user