[diffusion] model: LTX-2 (1/2) (#17495)

Co-authored-by: FlamingoPg <1106310035@qq.com>
This commit is contained in:
GMI Xiao Jin
2026-01-23 19:59:48 -08:00
committed by GitHub
parent 894928a951
commit 797a9811a2
9 changed files with 519 additions and 29 deletions

View File

@@ -14,8 +14,12 @@ import time
from typing import Any, List, Union
import numpy as np
import torch
from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams
from sglang.multimodal_gen.configs.sample.sampling_params import (
DataType,
SamplingParams,
)
from sglang.multimodal_gen.runtime.entrypoints.openai.utils import (
ListLorasReq,
MergeLoraWeightsReq,
@@ -231,8 +235,27 @@ class DiffGenerator:
request_idx + 1,
)
continue
audio_sample_rate = output_batch.audio_sample_rate
for output_idx, sample in enumerate(output_batch.output):
num_outputs = len(output_batch.output)
audio = output_batch.audio
if req.data_type == DataType.VIDEO:
if isinstance(audio, torch.Tensor) and audio.ndim >= 2:
audio = (
audio[output_idx]
if audio.shape[0] > output_idx
else None
)
elif isinstance(audio, np.ndarray) and audio.ndim >= 2:
audio = (
audio[output_idx]
if audio.shape[0] > output_idx
else None
)
if audio is not None and not (
isinstance(sample, (tuple, list)) and len(sample) == 2
):
sample = (sample, audio)
frames = post_process_sample(
sample,
fps=req.fps,
@@ -242,11 +265,13 @@ class DiffGenerator:
num_outputs, output_idx
),
data_type=req.data_type,
audio_sample_rate=audio_sample_rate,
)
result_item: dict[str, Any] = {
"samples": sample,
"frames": frames,
"audio": audio,
"prompts": req.prompt,
"size": (req.height, req.width, req.num_frames),
"generation_time": timer.duration,

View File

@@ -118,12 +118,28 @@ async def forward_to_scheduler(req_obj, sp):
raise RuntimeError("Model generation returned no output.")
output_file_path = sp.output_file_path()
sample = response.output[0]
try:
audio = response.audio
except AttributeError:
audio = None
if isinstance(audio, torch.Tensor) and audio.ndim >= 2:
audio = audio[0]
if audio is not None and not (
isinstance(sample, (tuple, list)) and len(sample) == 2
):
sample = (sample, audio)
post_process_sample(
sample=response.output[0],
sample=sample,
data_type=sp.data_type,
fps=sp.fps or 24,
save_output=True,
save_file_path=output_file_path,
audio_sample_rate=(
response.audio_sample_rate
if hasattr(response, "audio_sample_rate")
else None
),
)
if hasattr(response, "model_dump"):

View File

@@ -7,6 +7,7 @@ import time
from typing import Any, List, Optional, Union
import httpx
import torch
from fastapi import UploadFile
from sglang.multimodal_gen.configs.sample.sampling_params import DataType
@@ -211,22 +212,32 @@ async def process_generation_batch(
result = await scheduler_client.forward([batch])
if result.output is None:
error_msg = getattr(result, "error", "Unknown error")
error_msg = result.error or "Unknown error"
raise RuntimeError(
f"Model generation returned no output. Error from scheduler: {error_msg}"
)
save_file_path_list = []
audio_sample_rate = result.audio_sample_rate
if batch.data_type == DataType.VIDEO:
for idx, output in enumerate(result.output):
save_file_path = str(
os.path.join(batch.output_path, batch.output_file_name)
)
sample = result.output[idx]
audio = result.audio
if isinstance(audio, torch.Tensor) and audio.ndim >= 2:
audio = audio[idx] if audio.shape[0] > idx else None
if audio is not None and not (
isinstance(sample, (tuple, list)) and len(sample) == 2
):
sample = (sample, audio)
post_process_sample(
result.output[idx],
sample,
batch.data_type,
batch.fps,
batch.save_output,
save_file_path,
audio_sample_rate=audio_sample_rate,
)
save_file_path_list.append(save_file_path)
else:
@@ -242,6 +253,7 @@ async def process_generation_batch(
batch.fps,
batch.save_output,
save_file_path,
audio_sample_rate=audio_sample_rate,
)
save_file_path_list.append(save_file_path)

View File

@@ -9,10 +9,25 @@ diffusion models.
"""
import os
import shutil
import subprocess
import tempfile
from typing import Any, Optional
import imageio
import numpy as np
import torch
try:
import scipy.io.wavfile as scipy_wavfile
except ImportError: # pragma: no cover
scipy_wavfile = None
try:
import imageio_ffmpeg as _imageio_ffmpeg
except ImportError: # pragma: no cover
_imageio_ffmpeg = None
from sglang.multimodal_gen.configs.sample.sampling_params import (
DataType,
SamplingParams,
@@ -24,6 +39,171 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import CYAN, RESET, init_
logger = init_logger(__name__)
def _normalize_audio_to_numpy(audio: Any) -> np.ndarray | None:
"""Convert audio (torch / numpy) into a float32 numpy array in [-1, 1], best-effort."""
if audio is None:
return None
if isinstance(audio, torch.Tensor):
audio_np = audio.detach().float().clamp(-1.0, 1.0).cpu().numpy()
elif isinstance(audio, np.ndarray):
audio_np = audio.astype(np.float32, copy=False)
audio_np = np.clip(audio_np, -1.0, 1.0)
else:
return None
# 1. Squeeze leading singleton dimensions (Batch, etc.)
while audio_np.ndim > 1 and audio_np.shape[0] == 1:
audio_np = audio_np.squeeze(0)
# 2. Handle (C, L) -> (L, C)
if audio_np.ndim == 2 and audio_np.shape[0] < audio_np.shape[1]:
audio_np = audio_np.transpose(1, 0)
# 3. Final safety check: if still 2D and channels (dim 1) is huge, something is wrong
if audio_np.ndim == 2 and audio_np.shape[1] > 256 and audio_np.shape[0] == 1:
audio_np = audio_np.flatten()
return audio_np
def _pick_audio_sample_rate(
*,
audio_np: np.ndarray,
audio_sample_rate: Optional[int],
fps: int,
num_frames: int,
) -> int:
"""Pick a plausible sample rate, falling back to inferring from video duration."""
selected_sr = int(audio_sample_rate) if audio_sample_rate is not None else None
if selected_sr is None or not (8000 <= selected_sr <= 192000):
selected_sr = 24000
try:
duration_s = float(num_frames) / float(fps) if fps else 0.0
if duration_s > 0:
audio_len = (
int(audio_np.shape[0])
if audio_np.ndim == 2
else int(audio_np.shape[-1])
)
inferred_sr = int(round(float(audio_len) / duration_s))
if 8000 <= inferred_sr <= 192000:
selected_sr = inferred_sr
except Exception:
pass
return selected_sr
def _resolve_ffmpeg_exe() -> str:
ffmpeg_exe = "ffmpeg"
ffmpeg_on_path = shutil.which("ffmpeg")
if ffmpeg_on_path:
ffmpeg_exe = ffmpeg_on_path
try:
if _imageio_ffmpeg is not None:
ffmpeg_exe = _imageio_ffmpeg.get_ffmpeg_exe()
except Exception:
pass
ffmpeg_ok = False
if ffmpeg_exe:
if os.path.isabs(ffmpeg_exe):
ffmpeg_ok = os.path.exists(ffmpeg_exe)
else:
ffmpeg_ok = shutil.which(ffmpeg_exe) is not None
if not ffmpeg_ok:
raise RuntimeError("ffmpeg not found")
return ffmpeg_exe
def _mux_audio_np_into_mp4(
*,
save_file_path: str,
audio_np: np.ndarray,
sample_rate: int,
ffmpeg_exe: str,
) -> None:
merged_path = save_file_path.rsplit(".", 1)[0] + ".tmp_mux.mp4"
tmp_wav_path = None
try:
if scipy_wavfile is None:
raise RuntimeError(
"scipy is required to mux audio into mp4 (pip install scipy)"
)
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as f:
tmp_wav_path = f.name
scipy_wavfile.write(tmp_wav_path, sample_rate, audio_np)
subprocess.run(
[
ffmpeg_exe,
"-y",
"-i",
save_file_path,
"-i",
tmp_wav_path,
"-c:v",
"copy",
"-c:a",
"aac",
"-strict",
"experimental",
merged_path,
],
check=True,
stdout=subprocess.DEVNULL,
stderr=subprocess.DEVNULL,
)
os.replace(merged_path, save_file_path)
finally:
if tmp_wav_path:
try:
os.remove(tmp_wav_path)
except OSError:
pass
if os.path.exists(merged_path):
try:
os.remove(merged_path)
except OSError:
pass
def _maybe_mux_audio_into_mp4(
*,
save_file_path: str,
audio: Any,
frames: list,
fps: int,
audio_sample_rate: Optional[int],
) -> None:
"""Best-effort mux audio into an already-written mp4 at save_file_path.
Any failure should keep the silent video and only log a warning.
"""
audio_np = _normalize_audio_to_numpy(audio)
if audio_np is None:
return
selected_sr = _pick_audio_sample_rate(
audio_np=audio_np,
audio_sample_rate=audio_sample_rate,
fps=fps,
num_frames=len(frames),
)
try:
ffmpeg_exe = _resolve_ffmpeg_exe()
_mux_audio_np_into_mp4(
save_file_path=save_file_path,
audio_np=audio_np,
sample_rate=selected_sr,
ffmpeg_exe=ffmpeg_exe,
)
logger.info(f"Merged video saved to {CYAN}{save_file_path}{RESET}")
except Exception as e:
logger.warning(
"Failed to mux audio into mp4 (saved silent video): %s",
str(e),
)
def prepare_request(
server_args: ServerArgs,
sampling_params: SamplingParams,
@@ -32,7 +212,10 @@ def prepare_request(
Create a Req object with sampling_params as a parameter.
"""
req = Req(sampling_params=sampling_params, VSA_sparsity=server_args.VSA_sparsity)
diffusers_kwargs = getattr(sampling_params, "diffusers_kwargs", None)
try:
diffusers_kwargs = sampling_params.diffusers_kwargs
except AttributeError:
diffusers_kwargs = None
if diffusers_kwargs:
req.extra["diffusers_kwargs"] = diffusers_kwargs
@@ -49,27 +232,51 @@ def prepare_request(
def post_process_sample(
sample: torch.Tensor,
sample: Any,
data_type: DataType,
fps: int,
save_output: bool = True,
save_file_path: str = None,
save_file_path: Optional[str] = None,
audio_sample_rate: Optional[int] = None,
):
"""
Process sample output and save video if necessary
"""
# 1. Vectorized processing on GPU/CPU tensor
if sample.dim() == 3:
# for images, dim t is missing
sample = sample.unsqueeze(1)
audio = None
if isinstance(sample, (tuple, list)) and len(sample) == 2:
sample, audio = sample
# Convert to uint8 and move to CPU in bulk
# Shape: [C, T, H, W] -> [T, H, W, C]
sample = (sample * 255).clamp(0, 255).to(torch.uint8)
videos = sample.permute(1, 2, 3, 0).cpu().numpy()
frames = None
if isinstance(sample, torch.Tensor):
if sample.dim() == 3:
sample = sample.unsqueeze(1)
sample = (sample * 255).clamp(0, 255).to(torch.uint8)
videos = sample.permute(1, 2, 3, 0).cpu().numpy()
frames = list(videos)
else:
if not isinstance(sample, np.ndarray):
raise TypeError(f"Unsupported sample type: {type(sample)}")
# Convert to list of frames for imageio
frames = list(videos)
arr = sample
if arr.ndim == 3:
if arr.shape[-1] in (1, 3, 4):
arr = arr[None, ...]
else:
arr = arr[..., None]
if arr.ndim != 4:
raise ValueError(f"Unexpected numpy sample shape: {tuple(arr.shape)}")
if arr.shape[-1] not in (1, 3, 4) and arr.shape[0] in (1, 3, 4):
t = torch.from_numpy(arr)
if t.dim() == 3:
t = t.unsqueeze(1)
t = (t * 255).clamp(0, 255).to(torch.uint8)
videos = t.permute(1, 2, 3, 0).cpu().numpy()
frames = list(videos)
else:
if arr.dtype != np.uint8:
arr = (np.clip(arr, 0.0, 1.0) * 255.0).astype(np.uint8)
frames = list(arr)
# 2. Save outputs if requested
if save_output:
@@ -86,6 +293,15 @@ def post_process_sample(
codec="libx264",
quality=quality,
)
_maybe_mux_audio_into_mp4(
save_file_path=save_file_path,
audio=audio,
frames=frames,
fps=fps,
audio_sample_rate=audio_sample_rate,
)
else:
quality = 75
if len(frames) > 1:

View File

@@ -257,6 +257,38 @@ class RotaryEmbedding(CustomOp):
return s
class LinearScalingRotaryEmbedding(RotaryEmbedding):
def __init__(
self,
head_size: int,
rotary_dim: int,
max_position_embeddings: int,
base: int | float,
is_neox_style: bool,
dtype: torch.dtype,
scaling_factor: float,
) -> None:
self.scaling_factor = float(scaling_factor)
super().__init__(
head_size=head_size,
rotary_dim=rotary_dim,
max_position_embeddings=max_position_embeddings,
base=base,
is_neox_style=is_neox_style,
dtype=dtype,
)
def _compute_cos_sin_cache(self) -> torch.Tensor:
inv_freq = self._compute_inv_freq(self.base)
t = torch.arange(self.max_position_embeddings, dtype=torch.float)
t = t / self.scaling_factor
freqs = torch.einsum("i,j -> ij", t, inv_freq)
cos = freqs.cos()
sin = freqs.sin()
cache = torch.cat((cos, sin), dim=-1)
return cache
class OneDRotaryEmbedding(torch.nn.Module):
"""1D rotary positional embedding with caching."""
@@ -949,10 +981,23 @@ def get_rope(
rope_scaling_args = None
if partial_rotary_factor < 1.0:
rotary_dim = int(rotary_dim * partial_rotary_factor)
max_position_embeddings = max_position
rope_type = None
if rope_scaling is not None:
rope_type = rope_scaling.get("rope_type", rope_scaling.get("type", None))
if rope_type in (None, "default"):
rope_scaling = None
elif rope_type == "linear":
factor = float(rope_scaling.get("factor", 1.0))
original_max = rope_scaling.get("original_max_position_embeddings", None)
if original_max is not None:
max_position_embeddings = max(
max_position_embeddings, int(float(original_max) * factor)
)
key = (
head_size,
rotary_dim,
max_position,
max_position_embeddings,
base,
is_neox_style,
rope_scaling_args,
@@ -963,9 +1008,21 @@ def get_rope(
if rope_scaling is None:
rotary_emb = RotaryEmbedding(
head_size, rotary_dim, max_position, base, is_neox_style, dtype
head_size, rotary_dim, max_position_embeddings, base, is_neox_style, dtype
)
else:
raise ValueError(f"Unknown RoPE scaling {rope_scaling}")
if rope_type == "linear":
factor = float(rope_scaling.get("factor", 1.0))
rotary_emb = LinearScalingRotaryEmbedding(
head_size=head_size,
rotary_dim=rotary_dim,
max_position_embeddings=max_position_embeddings,
base=base,
is_neox_style=is_neox_style,
dtype=dtype,
scaling_factor=factor,
)
else:
raise ValueError(f"Unknown RoPE scaling {rope_scaling}")
_ROPE_DICT[key] = rotary_emb
return rotary_emb

View File

@@ -269,6 +269,17 @@ class ComponentLoader(ABC):
"vision_language_encoder": (VisionLanguageEncoderLoader, "transformers"),
}
# Loaders for audio/video specific components that might vary
av_module_loaders = {
"audio_vae": (VAELoader, "diffusers"),
"vocoder": (VocoderLoader, "diffusers"),
"connectors": (AdapterLoader, "diffusers"),
}
# NOTE(FlamingoPg): special for LTX-2 models
if module_type == "vocoder" or module_type == "connectors":
transformers_or_diffusers = "diffusers"
if module_type in module_loaders:
loader_cls, expected_library = module_loaders[module_type]
# Assert that the library matches what's expected for this module type
@@ -277,6 +288,11 @@ class ComponentLoader(ABC):
), f"{module_type} must be loaded from {expected_library}, got {transformers_or_diffusers}"
return loader_cls()
if module_type in av_module_loaders:
loader_cls, expected_library = av_module_loaders[module_type]
if transformers_or_diffusers == expected_library:
return loader_cls()
# For unknown module types, use a generic loader
logger.warning(
"No specific loader found for module type: %s. Using generic loader.",
@@ -484,6 +500,10 @@ class TextEncoderLoader(ComponentLoader):
self._get_all_weights(model, model_path, to_cpu=should_offload)
)
# Explicitly move model to target device after loading weights
if not should_offload:
model = model.to(local_torch_device)
if should_offload:
# Disable FSDP for MPS as it's not compatible
if current_platform.is_mps():
@@ -604,7 +624,7 @@ class VAELoader(ComponentLoader):
return server_args.vae_cpu_offload
def load_customized(
self, component_model_path: str, server_args: ServerArgs, *args
self, component_model_path: str, server_args: ServerArgs, module_name: str
):
"""Load the VAE based on the model path, and inference args."""
config = get_diffusers_component_config(model_path=component_model_path)
@@ -613,10 +633,16 @@ class VAELoader(ComponentLoader):
class_name is not None
), "Model config does not contain a _class_name attribute. Only diffusers format is supported."
server_args.model_paths["vae"] = component_model_path
server_args.model_paths[module_name] = component_model_path
logger.debug("HF model config: %s", config)
vae_config = server_args.pipeline_config.vae_config
if module_name == "audio_vae":
vae_config = server_args.pipeline_config.audio_vae_config
vae_precision = server_args.pipeline_config.audio_vae_precision
else:
vae_config = server_args.pipeline_config.vae_config
vae_precision = server_args.pipeline_config.vae_precision
vae_config.update_model_arch(config)
# NOTE: some post init logics are only available after updated with config
@@ -635,7 +661,7 @@ class VAELoader(ComponentLoader):
custom_module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(custom_module)
vae_cls = getattr(custom_module, cls_name)
vae_dtype = PRECISION_TO_TYPE[server_args.pipeline_config.vae_precision]
vae_dtype = PRECISION_TO_TYPE[vae_precision]
with set_default_torch_dtype(vae_dtype):
vae = vae_cls.from_pretrained(
component_model_path,
@@ -647,9 +673,7 @@ class VAELoader(ComponentLoader):
# Load from ModelRegistry (standard VAE classes)
with (
set_default_torch_dtype(
PRECISION_TO_TYPE[server_args.pipeline_config.vae_precision]
),
set_default_torch_dtype(PRECISION_TO_TYPE[vae_precision]),
skip_init_modules(),
):
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
@@ -664,6 +688,71 @@ class VAELoader(ComponentLoader):
return vae.eval()
class VocoderLoader(ComponentLoader):
def should_offload(
self, server_args: ServerArgs, model_config: ModelConfig | None = None
):
return server_args.vae_cpu_offload
def load_customized(
self, component_model_path: str, server_args: ServerArgs, module_name: str
):
config = get_diffusers_component_config(model_path=component_model_path)
class_name = config.pop("_class_name", None)
assert (
class_name is not None
), "Model config does not contain a _class_name attribute. Only diffusers format is supported."
server_args.model_paths[module_name] = component_model_path
from sglang.multimodal_gen.configs.models.vocoder.ltx_vocoder import (
LTXVocoderConfig,
)
vocoder_config = LTXVocoderConfig()
vocoder_config.update_model_arch(config)
try:
vocoder_precision = server_args.pipeline_config.audio_vae_precision
except AttributeError:
vocoder_precision = "fp32"
vocoder_dtype = PRECISION_TO_TYPE[vocoder_precision]
should_offload = self.should_offload(server_args)
target_device = self.target_device(should_offload)
with set_default_torch_dtype(vocoder_dtype), skip_init_modules():
vocoder_cls, _ = ModelRegistry.resolve_model_cls(class_name)
vocoder = vocoder_cls(vocoder_config).to(target_device)
safetensors_list = _list_safetensors_files(component_model_path)
assert (
len(safetensors_list) == 1
), f"Found {len(safetensors_list)} safetensors files in {component_model_path}"
loaded = safetensors_load_file(safetensors_list[0])
incompatible = vocoder.load_state_dict(loaded, strict=False)
missing_keys = []
unexpected_keys = []
try:
missing_keys = incompatible.missing_keys
unexpected_keys = incompatible.unexpected_keys
except AttributeError:
# Best-effort fallback in case older torch returns a tuple-like.
try:
missing_keys = incompatible[0]
unexpected_keys = incompatible[1]
except Exception:
pass
if missing_keys or unexpected_keys:
logger.warning(
"Loaded vocoder with missing_keys=%d unexpected_keys=%d",
len(missing_keys),
len(unexpected_keys),
)
return vocoder.eval()
class TransformerLoader(ComponentLoader):
"""Loader for transformer."""
@@ -755,6 +844,58 @@ class TransformerLoader(ComponentLoader):
return model
class AdapterLoader(ComponentLoader):
"""Loader for small adapter-style modules (e.g., LTX-2 connectors).
This loader intentionally avoids FSDP sharding and just:
1) Instantiates the module from `config.json`.
2) Loads a single safetensors state_dict.
"""
def load_customized(
self, component_model_path: str, server_args: ServerArgs, *args
):
config = get_diffusers_component_config(model_path=component_model_path)
cls_name = config.pop("_class_name", None)
if cls_name is None:
raise ValueError(
"Model config does not contain a _class_name attribute. "
"Only diffusers format is supported."
)
config.pop("_diffusers_version", None)
config.pop("_name_or_path", None)
server_args.model_paths["connectors"] = component_model_path
model_cls, _ = ModelRegistry.resolve_model_cls(cls_name)
target_device = get_local_torch_device()
default_dtype = PRECISION_TO_TYPE[server_args.pipeline_config.dit_precision]
from types import SimpleNamespace
with set_default_torch_dtype(default_dtype), skip_init_modules():
connector_cfg = SimpleNamespace(**config)
model = model_cls(connector_cfg).to(
device=target_device, dtype=default_dtype
)
safetensors_list = _list_safetensors_files(component_model_path)
if not safetensors_list:
raise ValueError(f"No safetensors files found in {component_model_path}")
if len(safetensors_list) != 1:
raise ValueError(
f"Found {len(safetensors_list)} safetensors files in {component_model_path}, expected 1"
)
loaded = safetensors_load_file(safetensors_list[0])
model.load_state_dict(loaded, strict=False)
return model.eval()
class SchedulerLoader(ComponentLoader):
"""Loader for scheduler."""

View File

@@ -82,11 +82,25 @@ class Req:
# Tracking if embeddings are already processed
is_prompt_processed: bool = False
# Audio Embeddings (LTX-2)
audio_prompt_embeds: list[torch.Tensor] | torch.Tensor = field(default_factory=list)
negative_audio_prompt_embeds: list[torch.Tensor] | torch.Tensor = field(
default_factory=list
)
# Latent tensors
latents: torch.Tensor | None = None
# Flux-2
latent_ids: torch.Tensor | None = None
# Audio Latents (LTX-2)
audio_latents: torch.Tensor | None = None
raw_audio_latent_shape: tuple[int, ...] | None = None
# Audio Parameters
fps: float = 24.0
generate_audio: bool = True
raw_latent_shape: torch.Tensor | None = None
noise_pred: torch.Tensor | None = None
# vae-encoded condition image
@@ -115,6 +129,7 @@ class Req:
trajectory_timesteps: list[torch.Tensor] | None = None
trajectory_latents: torch.Tensor | None = None
trajectory_audio_latents: torch.Tensor | None = None
# Extra parameters that might be needed by specific pipeline implementations
extra: dict[str, Any] = field(default_factory=dict)
@@ -297,6 +312,8 @@ class OutputBatch:
"""
output: torch.Tensor | None = None
audio: torch.Tensor | None = None
audio_sample_rate: int | None = None
trajectory_timesteps: list[torch.Tensor] | None = None
trajectory_latents: torch.Tensor | None = None
trajectory_decoded: list[torch.Tensor] | None = None

View File

@@ -77,7 +77,9 @@ class TextEncodingStage(PipelineStage):
for pe in pooler_embeds_list:
batch.pooled_embeds.append(pe)
if batch.prompt_attention_mask is not None:
if batch.prompt_attention_mask is None:
batch.prompt_attention_mask = []
for am in prompt_masks_list:
batch.prompt_attention_mask.append(am)
@@ -98,7 +100,8 @@ class TextEncodingStage(PipelineStage):
for pe in neg_pooler_embeds_list:
batch.neg_pooled_embeds.append(pe)
if batch.negative_attention_mask is not None:
if batch.negative_attention_mask is None:
batch.negative_attention_mask = []
for nm in neg_masks_list:
batch.negative_attention_mask.append(nm)

View File

@@ -74,6 +74,7 @@ class TimestepPreparationStage(PipelineStage):
n_tokens = batch.n_tokens
sigmas = server_args.pipeline_config.prepare_sigmas(sigmas, num_inference_steps)
batch.sigmas = sigmas
# Prepare extra kwargs for set_timesteps
extra_set_timesteps_kwargs = {}
@@ -87,6 +88,8 @@ class TimestepPreparationStage(PipelineStage):
key, value = callee(batch, server_args)
assert isinstance(key, str)
extra_set_timesteps_kwargs[key] = value
if key == "mu":
batch.extra["mu"] = value
# Handle custom timesteps or sigmas
if timesteps is not None and sigmas is not None: