diff --git a/docs/cookbook/diffusion/SenseNova/SenseNova-U1.5-8B-MoT.mdx b/docs/cookbook/diffusion/SenseNova/SenseNova-U1.5-8B-MoT.mdx
index 94ffb59c7..f5006040c 100644
--- a/docs/cookbook/diffusion/SenseNova/SenseNova-U1.5-8B-MoT.mdx
+++ b/docs/cookbook/diffusion/SenseNova/SenseNova-U1.5-8B-MoT.mdx
@@ -78,11 +78,38 @@ with open("output_sensenova_u15.png", "wb") as f:
f.write(image_bytes)
```
-## 5. Benchmark
+## 5. Advanced usage
-### 5.1 Speedup Benchmark
+### 5.1 Dynamic request batching
-#### 5.1.1 Single-run Profile
+Start the server with dynamic batching enabled:
+
+```bash Command
+sglang serve \
+ --model-path sensenova/SenseNova-U1.5-8B-MoT \
+ --port 30000 \
+ --batching-max-size 2 \
+ --batching-delay-ms 100
+```
+
+Submit requests concurrently so they arrive within the batching delay. Requests
+in one batch must use the same resolution, inference steps, guidance settings,
+and output options. Requests with multiple outputs or `think_mode=true` are
+executed sequentially.
+
+### 5.2 Ascend NPU optimizations
+
+On Ascend NPU, SenseNova-U1 automatically uses fused inference attention,
+RMSNorm, and SwiGLU MLP operators for supported inputs. FIA and SwiGLU fall back
+when their requirements are not met. Run one complete warmup request before
+benchmarking because the fused MLP packs its gate and up projection weights on
+first use.
+
+## 6. Benchmark
+
+### 6.1 Speedup benchmark
+
+#### 6.1.1 Single-run profile
@@ -137,4 +164,21 @@ with open("output_sensenova_u15.png", "wb") as f:
------------------------------------------------------------
```
+
+
+ Environment: one Ascend 910C; workload: four concurrent 2048 x 2048
+ text-to-image requests, 50 denoising steps, CFG 4, BF16.
+
+ | Batch size | NPU operators | Duration (s) | Throughput (images/s) | Mean latency (s) | Peak reserved memory (MB) |
+ |---:|---|---:|---:|---:|---:|
+ | 1 | Disabled | 234.28 | 0.01707 | 146.47 | 35698 |
+ | 2 | Disabled | 234.83 | 0.01703 | 176.06 | 37870 |
+ | 1 | FIA + RMSNorm + SwiGLU | 205.94 | 0.01942 | 128.80 | 35718 |
+ | 2 | FIA + RMSNorm + SwiGLU | 195.02 | 0.02051 | 146.33 | 37850 |
+
+ Optimized B1 improves throughput by 13.76% over unoptimized B1. Optimized
+ B2 improves throughput by 20.42% over unoptimized B2 and by 5.60% over
+ optimized B1.
+
+
diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/base.py b/python/sglang/multimodal_gen/configs/pipeline_configs/base.py
index 4c51a5397..24e5fed0c 100644
--- a/python/sglang/multimodal_gen/configs/pipeline_configs/base.py
+++ b/python/sglang/multimodal_gen/configs/pipeline_configs/base.py
@@ -426,6 +426,10 @@ class PipelineConfig:
"""
return self.task_type in (ModelTaskType.T2I, ModelTaskType.T2V)
+ def supports_dynamic_batching_for_request(self, batch) -> bool:
+ """Return whether one request may participate in dynamic batching."""
+ return True
+
def supports_disaggregation(self) -> bool:
"""Return whether multi-service disaggregated deployment is supported."""
diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/sensenova_u1.py b/python/sglang/multimodal_gen/configs/pipeline_configs/sensenova_u1.py
index 9b4bb3481..6a5cf96c2 100644
--- a/python/sglang/multimodal_gen/configs/pipeline_configs/sensenova_u1.py
+++ b/python/sglang/multimodal_gen/configs/pipeline_configs/sensenova_u1.py
@@ -8,6 +8,10 @@ from sglang.multimodal_gen.configs.pipeline_configs.base import (
from sglang.multimodal_gen.configs.pipeline_configs.model_deployment_config import (
ModelDeploymentConfig,
)
+from sglang.multimodal_gen.configs.sensenova_u1 import (
+ DEFAULT_THINK_MODE,
+ RESOLUTION_ALIGNMENT,
+)
def _is_runtime_option_requested(value) -> bool:
@@ -81,7 +85,23 @@ class SenseNovaU1PipelineConfig(PipelineConfig):
supports_cfg_parallel: bool = False
def supports_dynamic_batching(self):
- return False
+ return True
+
+ def supports_dynamic_batching_for_request(self, batch) -> bool:
+ sampling_params = getattr(batch, "sampling_params", None)
+ return not bool(getattr(sampling_params, "think_mode", DEFAULT_THINK_MODE))
+
+ def estimate_request_cost(self, batch) -> float:
+ image_tokens = (int(batch.width) // RESOLUTION_ALIGNMENT) * (
+ int(batch.height) // RESOLUTION_ALIGNMENT
+ )
+ cfg_branches = 2 if float(batch.guidance_scale) > 1 else 1
+ return float(
+ image_tokens
+ * int(batch.num_inference_steps)
+ * cfg_branches
+ * int(batch.num_outputs_per_prompt)
+ )
def supports_disaggregation(self) -> bool:
return False
diff --git a/python/sglang/multimodal_gen/configs/sample/sensenova_u1.py b/python/sglang/multimodal_gen/configs/sample/sensenova_u1.py
index 9aeb37e2b..9add7ace2 100644
--- a/python/sglang/multimodal_gen/configs/sample/sensenova_u1.py
+++ b/python/sglang/multimodal_gen/configs/sample/sensenova_u1.py
@@ -13,9 +13,9 @@ from sglang.multimodal_gen.configs.sensenova_u1 import (
DEFAULT_T_EPS,
DEFAULT_THINK_MODE,
DEFAULT_TIMESTEP_SHIFT,
+ RESOLUTION_ALIGNMENT,
SENSENOVA_U1_CFG_NORM_CHOICES,
SENSENOVA_U1_REQUEST_EXTRA_KEY,
- SENSENOVA_U1_RESOLUTION_ALIGNMENT,
)
_PUBLIC_OVERRIDE_FIELDS = {
@@ -74,12 +74,12 @@ class SenseNovaU1SamplingParams(SamplingParams):
def _validate(self) -> None:
super()._validate()
if (
- self.width % SENSENOVA_U1_RESOLUTION_ALIGNMENT != 0
- or self.height % SENSENOVA_U1_RESOLUTION_ALIGNMENT != 0
+ self.width % RESOLUTION_ALIGNMENT != 0
+ or self.height % RESOLUTION_ALIGNMENT != 0
):
raise ValueError(
"SenseNova-U1 requires width and height to be divisible by "
- f"{SENSENOVA_U1_RESOLUTION_ALIGNMENT}, got "
+ f"{RESOLUTION_ALIGNMENT}, got "
f"{self.width}x{self.height}."
)
if self.num_frames != 1:
diff --git a/python/sglang/multimodal_gen/configs/sensenova_u1.py b/python/sglang/multimodal_gen/configs/sensenova_u1.py
index 1056ee9ab..aa731201a 100644
--- a/python/sglang/multimodal_gen/configs/sensenova_u1.py
+++ b/python/sglang/multimodal_gen/configs/sensenova_u1.py
@@ -19,7 +19,7 @@ SENSENOVA_U1_CFG_NORM_CHOICES = (
"channel",
"cfg_zero_star",
)
-SENSENOVA_U1_RESOLUTION_ALIGNMENT = 32
+RESOLUTION_ALIGNMENT = 32
DEFAULT_CFG_NORM = "none"
DEFAULT_TIMESTEP_SHIFT = 3.0
diff --git a/python/sglang/multimodal_gen/runtime/managers/scheduler.py b/python/sglang/multimodal_gen/runtime/managers/scheduler.py
index 6653b60d5..ddd92c46b 100644
--- a/python/sglang/multimodal_gen/runtime/managers/scheduler.py
+++ b/python/sglang/multimodal_gen/runtime/managers/scheduler.py
@@ -44,6 +44,9 @@ from sglang.multimodal_gen.runtime.managers.dynamic_batch_admission import (
)
from sglang.multimodal_gen.runtime.managers.gpu_worker import GPUWorker
from sglang.multimodal_gen.runtime.pipelines_core import Req
+from sglang.multimodal_gen.runtime.pipelines_core.request_utils import (
+ normalize_output_seeds,
+)
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import (
BatchMetricsWindow,
OutputBatch,
@@ -559,6 +562,12 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
if base_req.is_warmup or candidate_req.is_warmup:
return "warmup"
+ if self._requires_sequential_multi_output(base_req, candidate_req):
+ return "sequential_multi_output"
+ if not self._pipeline_supports_dynamic_batching_for_request(
+ base_req, candidate_req
+ ):
+ return "pipeline_request_unsupported"
if self._has_realtime_session(base_req) or self._has_realtime_session(
candidate_req
):
@@ -589,11 +598,35 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
def _has_realtime_session(req: Req) -> bool:
return bool(req.realtime_session_id) or req.session is not None
+ def _requires_sequential_multi_output(self, *reqs: Req) -> bool:
+ pipeline_config = self.server_args.pipeline_config
+ return (
+ pipeline_config.supports_sequential_multi_output_inference()
+ and not pipeline_config.supports_sequential_dit_inference()
+ and any(max(1, int(req.num_outputs_per_prompt or 1)) > 1 for req in reqs)
+ )
+
+ def _pipeline_supports_dynamic_batching_for_request(self, *reqs: Req) -> bool:
+ checker = getattr(
+ self.server_args.pipeline_config,
+ "supports_dynamic_batching_for_request",
+ None,
+ )
+ return not callable(checker) or all(checker(req) for req in reqs)
+
def _can_dynamic_batch(self, base_req: Req, candidate_req: Req) -> bool:
"""Return whether `candidate_req` can be merged into a batch with `base_req`."""
if base_req.is_warmup or candidate_req.is_warmup:
return False
+ if self._requires_sequential_multi_output(base_req, candidate_req):
+ return False
+
+ if not self._pipeline_supports_dynamic_batching_for_request(
+ base_req, candidate_req
+ ):
+ return False
+
if self._has_realtime_session(base_req) or self._has_realtime_session(
candidate_req
):
@@ -842,11 +875,26 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
if not self._can_dynamic_batch(base_req, req):
return None
+ dynamic_batch_seeds: list[int | list[int]] = []
+ try:
+ for req in reqs:
+ if max(1, int(req.num_outputs_per_prompt or 1)) == 1:
+ dynamic_batch_seeds.append(
+ normalize_output_seeds(
+ req.seed,
+ num_outputs_per_prompt=1,
+ )[0]
+ )
+ else:
+ dynamic_batch_seeds.append(req.seed)
+ except (TypeError, ValueError):
+ return None
+
merged_req = deepcopy(base_req)
merged_req.prompt = [req.prompt for req in reqs]
merged_req.extra = deepcopy(merged_req.extra)
- merged_req.extra["dynamic_batch_seeds"] = [req.seed for req in reqs]
+ merged_req.extra["dynamic_batch_seeds"] = dynamic_batch_seeds
merged_req.return_file_paths_only = base_req.return_file_paths_only
if merged_req.return_file_paths_only:
dynamic_output_paths: list[str] = []
diff --git a/python/sglang/multimodal_gen/runtime/models/sensenova_u1/neo_unify/modeling_neo_chat.py b/python/sglang/multimodal_gen/runtime/models/sensenova_u1/neo_unify/modeling_neo_chat.py
index 81e3f05c5..94f06163f 100644
--- a/python/sglang/multimodal_gen/runtime/models/sensenova_u1/neo_unify/modeling_neo_chat.py
+++ b/python/sglang/multimodal_gen/runtime/models/sensenova_u1/neo_unify/modeling_neo_chat.py
@@ -21,7 +21,11 @@ from .modeling_fm_modules import (
TimestepEmbedder,
)
from .modeling_neo_vit import NEOVisionModel
-from .modeling_qwen3 import Qwen3ForCausalLM, create_block_causal_mask
+from .modeling_qwen3 import (
+ Qwen3ForCausalLM,
+ create_block_causal_mask,
+ npu_fia_available,
+)
from .modeling_qwen3_moe import Qwen3MoeForCausalLM
from .utils import SYSTEM_MESSAGE_FOR_GEN, load_image_native
@@ -37,20 +41,39 @@ def version_cmp(v1, v2, op="eq"):
return op_func(version.parse(v1), version.parse(v2))
+def _copy_right_aligned_prefix_bnsd(destination, source, lengths):
+ prefix_width = source.shape[2]
+ for batch_index, length in enumerate(lengths):
+ destination[batch_index, :, prefix_width - length : prefix_width].copy_(
+ source[batch_index, :, :length]
+ )
+
+
def prepare_flash_kv_cache(
past_key_values,
current_len: int,
batch_size: int,
+ prefix_lengths: Optional[torch.Tensor] = None,
):
"""
- Convert prefix cache from [B, H, S, D] to flash-attn friendly [B, S, H, D],
- and preallocate full KV buffer for [prefix + current].
+ Preallocate the full KV buffer for [prefix + current]. CUDA/SDPA use
+ [B, S, H, D], while NPU FIA keeps the source [B, H, S, D] layout.
This is done once before denoising loop.
"""
if past_key_values is None:
return
+ lengths = None
+ if prefix_lengths is not None:
+ lengths = [int(length) for length in prefix_lengths.reshape(-1).tolist()]
+ if len(lengths) == 1 and batch_size > 1:
+ lengths *= batch_size
+ if len(lengths) != batch_size:
+ raise ValueError(
+ f"Expected {batch_size} prefix lengths, got {len(lengths)}"
+ )
+
for layer in past_key_values.layers:
past_k = layer.keys
past_v = layer.values
@@ -60,28 +83,66 @@ def prepare_flash_kv_cache(
layer.flash_total_len = current_len
layer.flash_k_cache = None
layer.flash_v_cache = None
+ layer.flash_actual_seq_lengths_kv = None
+ layer.flash_kv_padding_size = None
+ layer.flash_cache_layout = None
continue
- # original cache layout assumed: [B, H, S, D]
- past_k_flash = past_k.transpose(1, 2).contiguous() # [B, S, H, D]
- past_v_flash = past_v.transpose(1, 2).contiguous() # [B, S, H, D]
-
- prefix_len = past_k_flash.shape[1]
+ # original cache layout: [B, H, S, D]
+ prefix_len = past_k.shape[2]
total_len = prefix_len + current_len
+ if lengths is not None and any(
+ length < 0 or length > prefix_len for length in lengths
+ ):
+ raise ValueError(
+ f"Prefix lengths must be between 0 and {prefix_len}, got {lengths}"
+ )
- k_cache = torch.empty(
- (batch_size, total_len, past_k_flash.shape[2], past_k_flash.shape[3]),
- device=past_k_flash.device,
- dtype=past_k_flash.dtype,
+ use_npu_fia = (
+ past_k.device.type == "npu" and lengths is not None and npu_fia_available()
)
- v_cache = torch.empty(
- (batch_size, total_len, past_v_flash.shape[2], past_v_flash.shape[3]),
- device=past_v_flash.device,
- dtype=past_v_flash.dtype,
- )
-
- k_cache[:, :prefix_len].copy_(past_k_flash)
- v_cache[:, :prefix_len].copy_(past_v_flash)
+ if use_npu_fia:
+ k_cache = torch.empty(
+ (batch_size, past_k.shape[1], total_len, past_k.shape[3]),
+ device=past_k.device,
+ dtype=past_k.dtype,
+ )
+ v_cache = torch.empty(
+ (batch_size, past_v.shape[1], total_len, past_v.shape[3]),
+ device=past_v.device,
+ dtype=past_v.dtype,
+ )
+ _copy_right_aligned_prefix_bnsd(k_cache, past_k, lengths)
+ _copy_right_aligned_prefix_bnsd(v_cache, past_v, lengths)
+ layer.flash_actual_seq_lengths_kv = [
+ length + current_len for length in lengths
+ ]
+ layer.flash_kv_padding_size = torch.zeros(
+ 1, dtype=torch.int64, device=past_k.device
+ )
+ layer.flash_cache_layout = "BNSD"
+ else:
+ past_k_flash = past_k.transpose(1, 2).contiguous()
+ past_v_flash = past_v.transpose(1, 2).contiguous()
+ k_cache = torch.empty(
+ (batch_size, total_len, past_k_flash.shape[2], past_k_flash.shape[3]),
+ device=past_k_flash.device,
+ dtype=past_k_flash.dtype,
+ )
+ v_cache = torch.empty(
+ (batch_size, total_len, past_v_flash.shape[2], past_v_flash.shape[3]),
+ device=past_v_flash.device,
+ dtype=past_v_flash.dtype,
+ )
+ k_cache[:, :prefix_len].copy_(past_k_flash)
+ v_cache[:, :prefix_len].copy_(past_v_flash)
+ layer.flash_actual_seq_lengths_kv = (
+ None
+ if lengths is None or past_k.device.type == "npu"
+ else [length + current_len for length in lengths]
+ )
+ layer.flash_kv_padding_size = None
+ layer.flash_cache_layout = "BSND"
layer.flash_prefix_len = prefix_len
layer.flash_total_len = total_len
@@ -101,6 +162,12 @@ def clear_flash_kv_cache(past_key_values):
delattr(layer, "flash_k_cache")
if hasattr(layer, "flash_v_cache"):
delattr(layer, "flash_v_cache")
+ if hasattr(layer, "flash_actual_seq_lengths_kv"):
+ delattr(layer, "flash_actual_seq_lengths_kv")
+ if hasattr(layer, "flash_kv_padding_size"):
+ delattr(layer, "flash_kv_padding_size")
+ if hasattr(layer, "flash_cache_layout"):
+ delattr(layer, "flash_cache_layout")
def optimized_scale(positive_flat, negative_flat):
@@ -127,7 +194,21 @@ def optimized_scale(positive_flat, negative_flat):
return st_star
-def _randn_with_seed(shape, *, device, dtype, seed: int) -> torch.Tensor:
+def _randn_with_seed(shape, *, device, dtype, seed: int | list[int]) -> torch.Tensor:
+ if isinstance(seed, list):
+ if len(shape) == 0 or len(seed) != shape[0]:
+ raise ValueError(
+ f"expected one seed per batch item, got {len(seed)} seeds for shape {shape}"
+ )
+ return torch.cat(
+ [
+ _randn_with_seed(
+ (1, *shape[1:]), device=device, dtype=dtype, seed=item_seed
+ )
+ for item_seed in seed
+ ],
+ dim=0,
+ )
try:
generator = torch.Generator(device)
except (RuntimeError, TypeError):
@@ -554,21 +635,61 @@ class NEOChatModel(PreTrainedModel):
return template.get_prompt() + append_text
return template.get_prompt()
- def _build_t2i_text_inputs(self, tokenizer, query: str):
- model_inputs = tokenizer(query, return_tensors="pt")
- input_ids = model_inputs["input_ids"].to(self.device)
+ def _build_t2i_text_inputs(self, tokenizer, query: str | list[str]):
+ queries = [query] if isinstance(query, str) else query
+ if not queries:
+ raise ValueError("query batch must not be empty")
+ encoded = [
+ tokenizer(item, return_tensors="pt")["input_ids"][0] for item in queries
+ ]
+ lengths = [item.shape[0] for item in encoded]
+ prefix_lengths = torch.tensor(lengths, dtype=torch.long, device=self.device)
+ max_length = max(lengths)
+ pad_token_id = getattr(tokenizer, "pad_token_id", None)
+ if pad_token_id is None:
+ pad_token_id = getattr(tokenizer, "eos_token_id", None)
+ if pad_token_id is None:
+ pad_token_id = 0
+
+ input_ids = torch.full(
+ (len(encoded), max_length),
+ int(pad_token_id),
+ dtype=encoded[0].dtype,
+ device=self.device,
+ )
+ key_valid_mask = torch.zeros(
+ (len(encoded), max_length), dtype=torch.bool, device=self.device
+ )
+ for batch_index, item in enumerate(encoded):
+ item = item.to(self.device)
+ item_length = item.shape[0]
+ input_ids[batch_index, :item_length] = item
+ key_valid_mask[batch_index, :item_length] = True
t_idx = torch.arange(
- 0, input_ids.shape[1], dtype=torch.long, device=input_ids.device
- )
+ max_length, dtype=torch.long, device=input_ids.device
+ ).expand(len(encoded), -1)
h_idx = torch.zeros_like(t_idx)
w_idx = torch.zeros_like(t_idx)
- indexes = torch.stack([t_idx, h_idx, w_idx], dim=0)
+ batched_indexes = torch.stack([t_idx, h_idx, w_idx], dim=1)
+ indexes = batched_indexes[0] if len(encoded) == 1 else batched_indexes
- attention_mask = {"full_attention": create_block_causal_mask(indexes[0])}
- return input_ids, indexes, attention_mask
+ attention_mask = {
+ "full_attention": create_block_causal_mask(t_idx, key_valid_mask)
+ }
+ return input_ids, indexes, attention_mask, key_valid_mask, prefix_lengths
def _build_t2i_image_indexes(self, token_h, token_w, text_len, device):
+ if isinstance(text_len, torch.Tensor):
+ text_len = text_len.to(device=device, dtype=torch.long).reshape(-1)
+ batch_size = text_len.shape[0]
+ image_len = token_h * token_w
+ t_image = text_len[:, None].expand(batch_size, image_len)
+ idx = torch.arange(image_len, device=device, dtype=torch.long)
+ h_image = (idx // token_w).expand(batch_size, -1)
+ w_image = (idx % token_w).expand(batch_size, -1)
+ return torch.stack([t_image, h_image, w_image], dim=1)
+
t_image = torch.full(
(token_h * token_w,), text_len, dtype=torch.long, device=device
)
@@ -2140,6 +2261,16 @@ class NEOChatModel(PreTrainedModel):
timesteps = self._apply_time_schedule(
timesteps, token_h * token_w, timestep_shift
)
+ denoise_embeddings = None
+ if device.type == "npu":
+ denoise_embeddings = self.fm_modules["timestep_embedder"](timesteps[:-1])
+ if self.add_noise_scale_embedding:
+ noise_level = timesteps.new_tensor(
+ [noise_scale / self.noise_scale_max_value]
+ )
+ denoise_embeddings = denoise_embeddings + self.fm_modules[
+ "noise_scale_embedder"
+ ](noise_level)
for step_i in range(num_steps):
t = timesteps[step_i]
@@ -2157,18 +2288,20 @@ class NEOChatModel(PreTrainedModel):
gen_model=True,
grid_hw=grid_hw,
).view(batch_size, token_h * token_w, -1)
- t_expanded = t.expand(batch_size * token_h * token_w)
- timestep_embeddings = self.fm_modules["timestep_embedder"](t_expanded).view(
- batch_size, token_h * token_w, -1
- )
- if self.add_noise_scale_embedding:
- noise_scale_tensor = torch.full_like(
- t_expanded, noise_scale / self.noise_scale_max_value
- )
- noise_embeddings = self.fm_modules["noise_scale_embedder"](
- noise_scale_tensor
+ if denoise_embeddings is not None:
+ timestep_embeddings = denoise_embeddings[step_i].view(1, 1, -1)
+ else:
+ t_expanded = t.expand(batch_size * token_h * token_w)
+ timestep_embeddings = self.fm_modules["timestep_embedder"](
+ t_expanded
).view(batch_size, token_h * token_w, -1)
- timestep_embeddings += noise_embeddings
+ if self.add_noise_scale_embedding:
+ noise_scale_tensor = torch.full_like(
+ t_expanded, noise_scale / self.noise_scale_max_value
+ )
+ timestep_embeddings += self.fm_modules["noise_scale_embedder"](
+ noise_scale_tensor
+ ).view(batch_size, token_h * token_w, -1)
image_embeds = image_embeds + timestep_embeddings
out_cond = self._t2i_predict_v(
@@ -2296,24 +2429,39 @@ class NEOChatModel(PreTrainedModel):
):
assert self.concat_time_token_num == 0
assert cfg_norm in ["cfg_zero_star", "global", "none", "channel"]
+ prompts = prompt if isinstance(prompt, list) else [prompt]
+ if not prompts:
+ raise ValueError("prompt batch must not be empty")
+ if batch_size < 1:
+ raise ValueError(f"batch_size must be positive, got {batch_size}")
+ if isinstance(prompt, list) and len(prompts) != batch_size:
+ raise ValueError(
+ f"batch_size={batch_size} does not match {len(prompts)} prompts"
+ )
+ if len(prompts) > 1 and think_mode:
+ raise ValueError(
+ "batched SenseNova-U1 generation does not support think_mode"
+ )
self._notify_layer_offload_phase("prefix")
merge_size = int(1 / self.downsample_ratio)
self.config.t_eps = t_eps
- # question_condition = f"Please generate an image based on the following description: {prompt}"
- question_condition = f"{prompt}"
- # question_condition += f"\nThe resolution of the image should be {image_size}"
-
think_text = ""
needs_cfg = cfg_scale > 1
think_content = (
"\n" if think_mode else "\n\n\n\n" + IMG_START_TOKEN
)
- query_condition = self._build_t2i_query(
- question_condition,
- system_message=SYSTEM_MESSAGE_FOR_GEN,
- append_text=think_content,
+ query_conditions = [
+ self._build_t2i_query(
+ item,
+ system_message=SYSTEM_MESSAGE_FOR_GEN,
+ append_text=think_content,
+ )
+ for item in prompts
+ ]
+ query_condition = (
+ query_conditions[0] if len(query_conditions) == 1 else query_conditions
)
query_uncondition = (
self._build_t2i_query("", append_text=IMG_START_TOKEN)
@@ -2321,19 +2469,26 @@ class NEOChatModel(PreTrainedModel):
else None
)
- input_ids_condition, indexes_condition, attention_mask_condition_prefix = (
- self._build_t2i_text_inputs(tokenizer, query_condition)
- )
+ (
+ input_ids_condition,
+ indexes_condition,
+ attention_mask_condition_prefix,
+ condition_key_valid_mask,
+ condition_prefix_lengths,
+ ) = self._build_t2i_text_inputs(tokenizer, query_condition)
if query_uncondition is not None:
(
input_ids_uncondition,
indexes_uncondition,
attention_mask_uncondition_prefix,
+ _,
+ uncondition_prefix_lengths,
) = self._build_t2i_text_inputs(tokenizer, query_uncondition)
else:
input_ids_uncondition = indexes_uncondition = (
attention_mask_uncondition_prefix
) = None
+ uncondition_prefix_lengths = None
token_h = image_size[1] // (self.patch_size * merge_size)
token_w = image_size[0] // (self.patch_size * merge_size)
@@ -2341,14 +2496,18 @@ class NEOChatModel(PreTrainedModel):
indexes_image_condition = self._build_t2i_image_indexes(
token_h,
token_w,
- indexes_condition.shape[1],
+ (
+ int(condition_prefix_lengths[0].item())
+ if len(prompts) == 1
+ else condition_prefix_lengths
+ ),
device=input_ids_condition.device,
)
indexes_image_uncondition = (
self._build_t2i_image_indexes(
token_h,
token_w,
- indexes_uncondition.shape[1],
+ int(uncondition_prefix_lengths[0].item()),
device=input_ids_uncondition.device,
)
if indexes_uncondition is not None
@@ -2436,12 +2595,14 @@ class NEOChatModel(PreTrainedModel):
past_key_values_condition,
current_len=token_h * token_w,
batch_size=batch_size,
+ prefix_lengths=condition_prefix_lengths,
)
if past_key_values_uncondition is not None:
prepare_flash_kv_cache(
past_key_values_uncondition,
current_len=token_h * token_w,
batch_size=batch_size,
+ prefix_lengths=uncondition_prefix_lengths,
)
# init noise image tokens
@@ -2465,6 +2626,22 @@ class NEOChatModel(PreTrainedModel):
)
attention_mask_condition = {"full_attention": None}
+ if device.type == "npu" and batch_size > 1 and not npu_fia_available():
+ condition_key_valid_mask = condition_key_valid_mask.expand(batch_size, -1)
+ image_key_valid_mask = torch.ones(
+ (batch_size, token_h * token_w),
+ dtype=torch.bool,
+ device=device,
+ )
+ denoise_key_valid_mask = torch.cat(
+ [condition_key_valid_mask, image_key_valid_mask], dim=1
+ )
+ # Ascend SDPA requires an explicit query dimension.
+ attention_mask_condition["full_attention"] = (
+ denoise_key_valid_mask[:, None, None, :]
+ .expand(-1, -1, token_h * token_w, -1)
+ .contiguous()
+ )
attention_mask_uncondition = {"full_attention": None}
timesteps = torch.linspace(0.0, 1.0, num_steps + 1, device=device)
@@ -2472,6 +2649,16 @@ class NEOChatModel(PreTrainedModel):
timesteps = self._apply_time_schedule(
timesteps, token_h * token_w, timestep_shift
)
+ denoise_embeddings = None
+ if device.type == "npu":
+ denoise_embeddings = self.fm_modules["timestep_embedder"](timesteps[:-1])
+ if self.add_noise_scale_embedding:
+ noise_level = timesteps.new_tensor(
+ [noise_scale / self.noise_scale_max_value]
+ )
+ denoise_embeddings = denoise_embeddings + self.fm_modules[
+ "noise_scale_embedder"
+ ](noise_level)
for step_i in range(num_steps):
t = timesteps[step_i]
@@ -2486,18 +2673,20 @@ class NEOChatModel(PreTrainedModel):
gen_model=True,
grid_hw=grid_hw,
).view(batch_size, token_h * token_w, -1)
- t_expanded = t.expand(batch_size * token_h * token_w)
- timestep_embeddings = self.fm_modules["timestep_embedder"](t_expanded).view(
- batch_size, token_h * token_w, -1
- )
- if self.add_noise_scale_embedding:
- noise_scale_tensor = torch.full_like(
- t_expanded, noise_scale / self.noise_scale_max_value
- )
- noise_embeddings = self.fm_modules["noise_scale_embedder"](
- noise_scale_tensor
+ if denoise_embeddings is not None:
+ timestep_embeddings = denoise_embeddings[step_i].view(1, 1, -1)
+ else:
+ t_expanded = t.expand(batch_size * token_h * token_w)
+ timestep_embeddings = self.fm_modules["timestep_embedder"](
+ t_expanded
).view(batch_size, token_h * token_w, -1)
- timestep_embeddings += noise_embeddings
+ if self.add_noise_scale_embedding:
+ noise_scale_tensor = torch.full_like(
+ t_expanded, noise_scale / self.noise_scale_max_value
+ )
+ timestep_embeddings += self.fm_modules["noise_scale_embedder"](
+ noise_scale_tensor
+ ).view(batch_size, token_h * token_w, -1)
image_embeds = image_embeds + timestep_embeddings
v_pred_condition = self._t2i_predict_v(
diff --git a/python/sglang/multimodal_gen/runtime/models/sensenova_u1/neo_unify/modeling_qwen3.py b/python/sglang/multimodal_gen/runtime/models/sensenova_u1/neo_unify/modeling_qwen3.py
index 7cc39b557..68d4d5396 100644
--- a/python/sglang/multimodal_gen/runtime/models/sensenova_u1/neo_unify/modeling_qwen3.py
+++ b/python/sglang/multimodal_gen/runtime/models/sensenova_u1/neo_unify/modeling_qwen3.py
@@ -5,12 +5,12 @@ from typing import Callable, Optional, Union
import torch
import torch._dynamo
+import torch.nn.functional as F
from torch import nn
from transformers import Qwen3Config
from transformers.activations import ACT2FN
from transformers.cache_utils import Cache, DynamicCache
from transformers.generation import GenerationMixin
-from transformers.integrations import use_kernel_forward_from_hub
from transformers.masking_utils import create_causal_mask
from transformers.modeling_flash_attention_utils import FlashAttentionKwargs
from transformers.modeling_layers import (
@@ -29,6 +29,9 @@ from transformers.processing_utils import Unpack
from transformers.utils import TransformersKwargs, can_return_tuple
from transformers.utils.deprecation import deprecate_kwarg
+from sglang.multimodal_gen.runtime.platforms import current_platform
+from sglang.srt.layers.layernorm import RMSNorm
+
from .transformers_compat import (
causal_mask_kwargs,
model_input_compat,
@@ -55,6 +58,10 @@ _VALID_ATTN_BACKENDS = ("auto", "flash", "sdpa")
_ATTN_BACKEND: str = "auto"
+def npu_fia_available() -> bool:
+ return hasattr(torch.ops.npu, "npu_fused_infer_attention_score")
+
+
def set_attn_backend(backend: str) -> str:
"""Choose the attention kernel used by the Qwen3 layers at runtime.
@@ -92,7 +99,13 @@ def effective_attn_backend() -> str:
def _sdpa_attn_func(
- q, k, v, dropout_p: float = 0.0, softmax_scale=None, causal: bool = False
+ q,
+ k,
+ v,
+ dropout_p: float = 0.0,
+ softmax_scale=None,
+ causal: bool = False,
+ attention_mask: Optional[torch.Tensor] = None,
):
"""Drop-in SDPA fallback for ``flash_attn_func``.
@@ -127,6 +140,7 @@ def _sdpa_attn_func(
q_bhsd,
k_bhsd,
v_bhsd,
+ attn_mask=attention_mask,
dropout_p=dropout_p,
is_causal=causal,
scale=softmax_scale,
@@ -138,6 +152,7 @@ def _sdpa_attn_func(
q_bhsd,
k_bhsd,
v_bhsd,
+ attn_mask=attention_mask,
dropout_p=dropout_p,
is_causal=causal,
)
@@ -146,6 +161,7 @@ def _sdpa_attn_func(
q_bhsd,
k_bhsd,
v_bhsd,
+ attn_mask=attention_mask,
dropout_p=dropout_p,
is_causal=causal,
)
@@ -153,37 +169,167 @@ def _sdpa_attn_func(
def _flash_or_sdpa(
- q, k, v, dropout_p: float = 0.0, softmax_scale=None, causal: bool = False
+ q,
+ k,
+ v,
+ dropout_p: float = 0.0,
+ softmax_scale=None,
+ causal: bool = False,
+ attention_mask: Optional[torch.Tensor] = None,
+ actual_seq_lengths_kv: Optional[list[int]] = None,
+ kv_padding_size: Optional[torch.Tensor] = None,
+ input_layout: str = "BSND",
):
backend = effective_attn_backend()
+ if (
+ q.device.type == "npu"
+ and not causal
+ and dropout_p == 0.0
+ and actual_seq_lengths_kv is not None
+ and npu_fia_available()
+ ):
+ if input_layout == "BNSD":
+ q_bhsd, k_bhsd, v_bhsd = q, k, v
+ batch_size, num_heads, query_length, _ = q.shape
+ num_key_value_heads = k.shape[1]
+ elif input_layout == "BSND":
+ q_bhsd = q.transpose(1, 2).contiguous()
+ k_bhsd = k.transpose(1, 2).contiguous()
+ v_bhsd = v.transpose(1, 2).contiguous()
+ batch_size, query_length, num_heads, _ = q.shape
+ num_key_value_heads = k.shape[2]
+ else:
+ raise ValueError(f"Unsupported attention input layout: {input_layout}")
+ out, _ = torch.ops.npu.npu_fused_infer_attention_score(
+ q_bhsd,
+ k_bhsd,
+ v_bhsd,
+ actual_seq_lengths=[query_length] * batch_size,
+ actual_seq_lengths_kv=actual_seq_lengths_kv,
+ kv_padding_size=kv_padding_size,
+ num_heads=num_heads,
+ num_key_value_heads=num_key_value_heads,
+ scale=q.shape[-1] ** -0.5 if softmax_scale is None else softmax_scale,
+ input_layout="BNSD",
+ sparse_mode=0,
+ )
+ return out.transpose(1, 2).contiguous()
+ if input_layout != "BSND":
+ raise RuntimeError("BNSD attention input requires the NPU FIA path")
+ if actual_seq_lengths_kv is not None:
+ batch_size, query_length = q.shape[:2]
+ padded_key_length = k.shape[1]
+ prefix_width = padded_key_length - query_length
+ if len(actual_seq_lengths_kv) != batch_size:
+ raise ValueError(
+ f"Expected {batch_size} KV lengths, got {len(actual_seq_lengths_kv)}"
+ )
+ if all(length == padded_key_length for length in actual_seq_lengths_kv):
+ return _flash_or_sdpa(
+ q,
+ k,
+ v,
+ dropout_p=dropout_p,
+ softmax_scale=softmax_scale,
+ causal=causal,
+ attention_mask=attention_mask,
+ )
+ outputs = []
+ for batch_index, total_length in enumerate(actual_seq_lengths_kv):
+ prefix_length = total_length - query_length
+ if prefix_length < 0 or prefix_length > prefix_width:
+ raise ValueError(
+ f"KV length {total_length} is incompatible with query length "
+ f"{query_length} and padded key length {padded_key_length}"
+ )
+ if total_length == padded_key_length:
+ compact_k = k[batch_index : batch_index + 1]
+ compact_v = v[batch_index : batch_index + 1]
+ else:
+ compact_k = torch.cat(
+ (
+ k[batch_index : batch_index + 1, :prefix_length],
+ k[batch_index : batch_index + 1, prefix_width:],
+ ),
+ dim=1,
+ )
+ compact_v = torch.cat(
+ (
+ v[batch_index : batch_index + 1, :prefix_length],
+ v[batch_index : batch_index + 1, prefix_width:],
+ ),
+ dim=1,
+ )
+ outputs.append(
+ _flash_or_sdpa(
+ q[batch_index : batch_index + 1],
+ compact_k,
+ compact_v,
+ dropout_p=dropout_p,
+ softmax_scale=softmax_scale,
+ causal=causal,
+ attention_mask=attention_mask,
+ )
+ )
+ return torch.cat(outputs, dim=0)
# flash-attn ships CUDA kernels only. On XPU / CPU we transparently fall
# back to SDPA even if the user asked for ``flash`` — the alternative
# (crashing on first forward) is worse, and ``set_attn_backend('flash')``
# already guarded against the "package missing" case.
- if backend == "flash" and q.device.type == "cuda":
+ if backend == "flash" and q.device.type == "cuda" and attention_mask is None:
return flash_attn_func(
q, k, v, dropout_p=dropout_p, softmax_scale=softmax_scale, causal=causal
)
return _sdpa_attn_func(
- q, k, v, dropout_p=dropout_p, softmax_scale=softmax_scale, causal=causal
+ q,
+ k,
+ v,
+ dropout_p=dropout_p,
+ softmax_scale=softmax_scale,
+ causal=causal,
+ attention_mask=attention_mask,
)
-def create_block_causal_mask(index: torch.Tensor):
- """
- index: (L)
- return: (1, 1, L, L) block-wise causal attention mask
- """
- L = index.size(0)
- idx_i = index.unsqueeze(1).expand(L, L)
- idx_j = index.unsqueeze(0).expand(L, L)
+def position_ids_from_indexes(indexes: torch.Tensor, coordinate: int) -> torch.Tensor:
+ """Return one coordinate as ``[batch, sequence]`` position IDs."""
+ if indexes.ndim == 2:
+ return indexes[coordinate].unsqueeze(0)
+ if indexes.ndim == 3:
+ return indexes[:, coordinate]
+ raise ValueError(f"indexes must have 2 or 3 dimensions, got {indexes.ndim}")
- arange = torch.arange(L, device=index.device)
- mask = (idx_j == idx_i) | (arange.unsqueeze(0) <= arange.unsqueeze(1))
- return torch.where(
- mask[None, None, :, :] > 0, torch.tensor(0.0), torch.tensor(float("-inf"))
- )
+def create_block_causal_mask(
+ index: torch.Tensor, key_valid_mask: Optional[torch.Tensor] = None
+):
+ """
+ index: (L) or (B, L)
+ key_valid_mask: optional (B, L), where True marks a real token
+ return: (B, 1, L, L) block-wise causal attention mask
+ """
+ if index.ndim == 1:
+ index = index.unsqueeze(0)
+ if index.ndim != 2:
+ raise ValueError(f"index must have 1 or 2 dimensions, got {index.ndim}")
+
+ batch_size, seq_len = index.shape
+ idx_i = index.unsqueeze(2)
+ idx_j = index.unsqueeze(1)
+
+ arange = torch.arange(seq_len, device=index.device)
+ mask = (idx_j == idx_i) | (arange.view(1, 1, -1) <= arange.view(1, -1, 1))
+ if key_valid_mask is not None:
+ if key_valid_mask.shape != (batch_size, seq_len):
+ raise ValueError(
+ "key_valid_mask must match the batched index shape; "
+ f"got {tuple(key_valid_mask.shape)} and {(batch_size, seq_len)}"
+ )
+ mask = mask & key_valid_mask.to(torch.bool).unsqueeze(1)
+
+ output = torch.zeros(mask.shape, dtype=torch.float32, device=index.device)
+ output.masked_fill_(~mask, float("-inf"))
+ return output.unsqueeze(1)
def visualize_mask(mask: torch.Tensor, i: int = 0, j: int = 12):
@@ -196,25 +342,13 @@ def visualize_mask(mask: torch.Tensor, i: int = 0, j: int = 12):
print(" ".join(map(str, row)))
-@use_kernel_forward_from_hub("RMSNorm")
-class Qwen3RMSNorm(nn.Module):
- def __init__(self, hidden_size, eps: float = 1e-6) -> None:
- """
- Qwen3RMSNorm is equivalent to T5LayerNorm
- """
- super().__init__()
- self.weight = nn.Parameter(torch.ones(hidden_size))
- self.variance_epsilon = eps
-
- def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
- input_dtype = hidden_states.dtype
- hidden_states = hidden_states.to(torch.float32)
- variance = hidden_states.pow(2).mean(-1, keepdim=True)
- hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)
- return self.weight * hidden_states.to(input_dtype)
-
- def extra_repr(self):
- return f"{tuple(self.weight.shape)}, eps={self.variance_epsilon}"
+def make_qwen3_rms_norm(hidden_size: int, eps: float) -> RMSNorm:
+ return RMSNorm(
+ hidden_size,
+ eps=eps,
+ cast_x_before_out_mul=True,
+ force_native=not current_platform.is_npu(),
+ )
class Qwen3MLP(nn.Module):
@@ -227,8 +361,45 @@ class Qwen3MLP(nn.Module):
self.up_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False)
self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=False)
self.act_fn = ACT2FN[config.hidden_act]
+ self._npu_gate_up_weight = None
+
+ def _pack_npu_gate_up_weights(self) -> torch.Tensor:
+ reference = self.gate_proj.weight
+ packed = self._npu_gate_up_weight
+ if (
+ packed is not None
+ and packed.device == reference.device
+ and packed.dtype == reference.dtype
+ and packed.untyped_storage().data_ptr()
+ == reference.untyped_storage().data_ptr()
+ ):
+ return packed
+
+ with torch.no_grad():
+ packed = torch.cat(
+ [self.gate_proj.weight, self.up_proj.weight], dim=0
+ ).contiguous()
+ self.gate_proj.weight.set_(packed[: self.intermediate_size])
+ self.up_proj.weight.set_(packed[self.intermediate_size :])
+ self._npu_gate_up_weight = packed
+ return packed
+
+ def _use_npu_fused_mlp(self, x: torch.Tensor) -> bool:
+ if (
+ x.device.type != "npu"
+ or self.training
+ or torch.is_grad_enabled()
+ or x.dtype != torch.bfloat16
+ or self.gate_proj.weight.dtype != x.dtype
+ or self.config.hidden_act != "silu"
+ ):
+ return False
+ return hasattr(torch.ops.npu, "npu_swiglu")
def forward(self, x):
+ if self._use_npu_fused_mlp(x):
+ gate_up = F.linear(x, self._pack_npu_gate_up_weights())
+ return self.down_proj(torch.ops.npu.npu_swiglu(gate_up))
down_proj = self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x))
return down_proj
@@ -462,23 +633,29 @@ class Qwen3Attention(nn.Module):
bias=config.attention_bias,
)
- self.q_norm = Qwen3RMSNorm(
+ self.q_norm = make_qwen3_rms_norm(
self.head_dim // 2, eps=config.rms_norm_eps
) # unlike olmo, only on the head dim!
- self.q_norm_mot_gen = Qwen3RMSNorm(self.head_dim // 2, eps=config.rms_norm_eps)
- self.q_norm_hw = Qwen3RMSNorm(self.head_dim // 2, eps=config.rms_norm_eps)
- self.q_norm_hw_mot_gen = Qwen3RMSNorm(
+ self.q_norm_mot_gen = make_qwen3_rms_norm(
+ self.head_dim // 2, eps=config.rms_norm_eps
+ )
+ self.q_norm_hw = make_qwen3_rms_norm(
+ self.head_dim // 2, eps=config.rms_norm_eps
+ )
+ self.q_norm_hw_mot_gen = make_qwen3_rms_norm(
self.head_dim // 2, eps=config.rms_norm_eps
)
- self.k_norm = Qwen3RMSNorm(
+ self.k_norm = make_qwen3_rms_norm(
self.head_dim // 2, eps=config.rms_norm_eps
) # thus post q_norm does not need reshape
- self.k_norm_mot_gen = Qwen3RMSNorm(self.head_dim // 2, eps=config.rms_norm_eps)
- self.k_norm_hw = Qwen3RMSNorm(
+ self.k_norm_mot_gen = make_qwen3_rms_norm(
+ self.head_dim // 2, eps=config.rms_norm_eps
+ )
+ self.k_norm_hw = make_qwen3_rms_norm(
self.head_dim // 2, eps=config.rms_norm_eps
) # thus post q_norm does not need reshape
- self.k_norm_hw_mot_gen = Qwen3RMSNorm(
+ self.k_norm_hw_mot_gen = make_qwen3_rms_norm(
self.head_dim // 2, eps=config.rms_norm_eps
)
@@ -530,17 +707,23 @@ class Qwen3Attention(nn.Module):
value_states = self.v_proj(hidden_states).view(hidden_shape).transpose(1, 2)
- cos_t, sin_t = self.rotary_emb(hidden_states, indexes[0].unsqueeze(0))
+ cos_t, sin_t = self.rotary_emb(
+ hidden_states, position_ids_from_indexes(indexes, 0)
+ )
query_states_t, key_states_t = apply_rotary_pos_emb(
query_states_t, key_states_t, cos_t, sin_t
)
- cos_h, sin_h = self.rotary_emb_hw(hidden_states, indexes[1].unsqueeze(0))
+ cos_h, sin_h = self.rotary_emb_hw(
+ hidden_states, position_ids_from_indexes(indexes, 1)
+ )
query_states_h, key_states_h = apply_rotary_pos_emb(
query_states_h, key_states_h, cos_h, sin_h
)
- cos_w, sin_w = self.rotary_emb_hw(hidden_states, indexes[2].unsqueeze(0))
+ cos_w, sin_w = self.rotary_emb_hw(
+ hidden_states, position_ids_from_indexes(indexes, 2)
+ )
query_states_w, key_states_w = apply_rotary_pos_emb(
query_states_w, key_states_w, cos_w, sin_w
)
@@ -705,17 +888,23 @@ class Qwen3Attention(nn.Module):
) # [B,H,S,D]
# RoPE
- cos_t, sin_t = self.rotary_emb(hidden_states, indexes[0].unsqueeze(0))
+ cos_t, sin_t = self.rotary_emb(
+ hidden_states, position_ids_from_indexes(indexes, 0)
+ )
query_states_t, key_states_t = apply_rotary_pos_emb(
query_states_t, key_states_t, cos_t, sin_t
)
- cos_h, sin_h = self.rotary_emb_hw(hidden_states, indexes[1].unsqueeze(0))
+ cos_h, sin_h = self.rotary_emb_hw(
+ hidden_states, position_ids_from_indexes(indexes, 1)
+ )
query_states_h, key_states_h = apply_rotary_pos_emb(
query_states_h, key_states_h, cos_h, sin_h
)
- cos_w, sin_w = self.rotary_emb_hw(hidden_states, indexes[2].unsqueeze(0))
+ cos_w, sin_w = self.rotary_emb_hw(
+ hidden_states, position_ids_from_indexes(indexes, 2)
+ )
query_states_w, key_states_w = apply_rotary_pos_emb(
query_states_w, key_states_w, cos_w, sin_w
)
@@ -736,61 +925,93 @@ class Qwen3Attention(nn.Module):
# current image tokens attend to [prefix + current image tokens]
# fully bidirectional inside current block => causal=False
# ------------------------------------------------------------------
- if attention_mask is None:
- # Convert current q/k/v to flash layout [B, S, H, D]
- q = query_states.transpose(1, 2).contiguous()
- k_cur = key_states.transpose(1, 2).contiguous()
- v_cur = value_states.transpose(1, 2).contiguous()
+ key_padding_mask = (
+ attention_mask is not None
+ and attention_mask.ndim == 4
+ and (attention_mask.shape[-2] == 1 or attention_mask.dtype == torch.bool)
+ )
+ if attention_mask is None or key_padding_mask:
+ actual_seq_lengths_kv = None
+ kv_padding_size = None
+ input_layout = "BSND"
+ layer = (
+ past_key_values.layers[self.layer_idx]
+ if past_key_values is not None and not update_cache
+ else None
+ )
+ use_bnsd_cache = (
+ layer is not None
+ and getattr(layer, "flash_cache_layout", None) == "BNSD"
+ and getattr(layer, "flash_k_cache", None) is not None
+ and getattr(layer, "flash_v_cache", None) is not None
+ )
- if past_key_values is not None:
- if update_cache:
- # Rare path, keep compatibility.
- # past_key_values.update expects [B,H,S,D]
- key_states, value_states = past_key_values.update(
- key_states, value_states, self.layer_idx, cache_kwargs=None
- )
- k = key_states.transpose(1, 2).contiguous()
- v = value_states.transpose(1, 2).contiguous()
- else:
- # Optimized path:
- # use preallocated flash_k_cache / flash_v_cache
- layer = past_key_values.layers[self.layer_idx]
-
- if (
- hasattr(layer, "flash_k_cache")
- and layer.flash_k_cache is not None
- and hasattr(layer, "flash_v_cache")
- and layer.flash_v_cache is not None
- ):
- prefix_len = layer.flash_prefix_len
- cur_len = k_cur.shape[1]
-
- # overwrite current segment in-place
- layer.flash_k_cache[:, prefix_len : prefix_len + cur_len].copy_(
- k_cur
- )
- layer.flash_v_cache[:, prefix_len : prefix_len + cur_len].copy_(
- v_cur
- )
-
- k = layer.flash_k_cache[:, : prefix_len + cur_len]
- v = layer.flash_v_cache[:, : prefix_len + cur_len]
- else:
- # fallback if user forgot to prepare flash cache
- layer = past_key_values.layers[self.layer_idx]
- past_k, past_v = layer.keys, layer.values
-
- if past_k is not None:
- past_k = past_k.transpose(1, 2).contiguous()
- past_v = past_v.transpose(1, 2).contiguous()
- k = torch.cat([past_k, k_cur], dim=1)
- v = torch.cat([past_v, v_cur], dim=1)
- else:
- k = k_cur
- v = v_cur
+ if use_bnsd_cache:
+ q = query_states
+ k_cur = key_states
+ v_cur = value_states
+ prefix_len = layer.flash_prefix_len
+ cur_len = k_cur.shape[2]
+ layer.flash_k_cache[:, :, prefix_len : prefix_len + cur_len].copy_(
+ k_cur
+ )
+ layer.flash_v_cache[:, :, prefix_len : prefix_len + cur_len].copy_(
+ v_cur
+ )
+ k = layer.flash_k_cache[:, :, : prefix_len + cur_len]
+ v = layer.flash_v_cache[:, :, : prefix_len + cur_len]
+ actual_seq_lengths_kv = layer.flash_actual_seq_lengths_kv
+ kv_padding_size = layer.flash_kv_padding_size
+ input_layout = "BNSD"
else:
- k = k_cur
- v = v_cur
+ q = query_states.transpose(1, 2).contiguous()
+ k_cur = key_states.transpose(1, 2).contiguous()
+ v_cur = value_states.transpose(1, 2).contiguous()
+
+ if past_key_values is not None:
+ if update_cache:
+ # Rare path, keep compatibility.
+ # past_key_values.update expects [B,H,S,D]
+ key_states, value_states = past_key_values.update(
+ key_states, value_states, self.layer_idx, cache_kwargs=None
+ )
+ k = key_states.transpose(1, 2).contiguous()
+ v = value_states.transpose(1, 2).contiguous()
+ else:
+ if (
+ getattr(layer, "flash_k_cache", None) is not None
+ and getattr(layer, "flash_v_cache", None) is not None
+ ):
+ prefix_len = layer.flash_prefix_len
+ cur_len = k_cur.shape[1]
+ layer.flash_k_cache[
+ :, prefix_len : prefix_len + cur_len
+ ].copy_(k_cur)
+ layer.flash_v_cache[
+ :, prefix_len : prefix_len + cur_len
+ ].copy_(v_cur)
+ k = layer.flash_k_cache[:, : prefix_len + cur_len]
+ v = layer.flash_v_cache[:, : prefix_len + cur_len]
+ actual_seq_lengths_kv = getattr(
+ layer, "flash_actual_seq_lengths_kv", None
+ )
+ kv_padding_size = getattr(
+ layer, "flash_kv_padding_size", None
+ )
+ else:
+ # fallback if the cache was not prepared
+ past_k, past_v = layer.keys, layer.values
+ if past_k is not None:
+ past_k = past_k.transpose(1, 2).contiguous()
+ past_v = past_v.transpose(1, 2).contiguous()
+ k = torch.cat([past_k, k_cur], dim=1)
+ v = torch.cat([past_v, v_cur], dim=1)
+ else:
+ k = k_cur
+ v = v_cur
+ else:
+ k = k_cur
+ v = v_cur
# sanity checks
assert q.ndim == 4 and k.ndim == 4 and v.ndim == 4
@@ -806,6 +1027,10 @@ class Qwen3Attention(nn.Module):
dropout_p=0.0 if not self.training else self.attention_dropout,
softmax_scale=self.scaling,
causal=False,
+ attention_mask=attention_mask,
+ actual_seq_lengths_kv=actual_seq_lengths_kv,
+ kv_padding_size=kv_padding_size,
+ input_layout=input_layout,
) # [B, S_q, H_q, D]
attn_output = attn_output.reshape(*input_shape, -1).contiguous()
@@ -982,17 +1207,23 @@ class Qwen3Attention(nn.Module):
)
value_states = value_states.view(hidden_shape).transpose(1, 2)
- cos_t, sin_t = self.rotary_emb(hidden_states, indexes[0].unsqueeze(0))
+ cos_t, sin_t = self.rotary_emb(
+ hidden_states, position_ids_from_indexes(indexes, 0)
+ )
query_states_t, key_states_t = apply_rotary_pos_emb(
query_states_t, key_states_t, cos_t, sin_t
)
- cos_h, sin_h = self.rotary_emb_hw(hidden_states, indexes[1].unsqueeze(0))
+ cos_h, sin_h = self.rotary_emb_hw(
+ hidden_states, position_ids_from_indexes(indexes, 1)
+ )
query_states_h, key_states_h = apply_rotary_pos_emb(
query_states_h, key_states_h, cos_h, sin_h
)
- cos_w, sin_w = self.rotary_emb_hw(hidden_states, indexes[2].unsqueeze(0))
+ cos_w, sin_w = self.rotary_emb_hw(
+ hidden_states, position_ids_from_indexes(indexes, 2)
+ )
query_states_w, key_states_w = apply_rotary_pos_emb(
query_states_w, key_states_w, cos_w, sin_w
)
@@ -1065,14 +1296,16 @@ class Qwen3DecoderLayer(GradientCheckpointingLayer):
self.mlp = Qwen3MLP(config)
self.mlp_mot_gen = Qwen3MLP(config)
- self.input_layernorm = Qwen3RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
- self.input_layernorm_mot_gen = Qwen3RMSNorm(
+ self.input_layernorm = make_qwen3_rms_norm(
config.hidden_size, eps=config.rms_norm_eps
)
- self.post_attention_layernorm = Qwen3RMSNorm(
+ self.input_layernorm_mot_gen = make_qwen3_rms_norm(
config.hidden_size, eps=config.rms_norm_eps
)
- self.post_attention_layernorm_mot_gen = Qwen3RMSNorm(
+ self.post_attention_layernorm = make_qwen3_rms_norm(
+ config.hidden_size, eps=config.rms_norm_eps
+ )
+ self.post_attention_layernorm_mot_gen = make_qwen3_rms_norm(
config.hidden_size, eps=config.rms_norm_eps
)
self.attention_type = config.layer_types[layer_idx]
@@ -1288,8 +1521,10 @@ class Qwen3Model(Qwen3PreTrainedModel):
for layer_idx in range(config.num_hidden_layers)
]
)
- self.norm = Qwen3RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
- self.norm_mot_gen = Qwen3RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
+ self.norm = make_qwen3_rms_norm(config.hidden_size, eps=config.rms_norm_eps)
+ self.norm_mot_gen = make_qwen3_rms_norm(
+ config.hidden_size, eps=config.rms_norm_eps
+ )
self.gradient_checkpointing = False
self.has_sliding_layers = "sliding_attention" in self.config.layer_types
@@ -1375,11 +1610,13 @@ class Qwen3Model(Qwen3PreTrainedModel):
)
else:
causal_mask_mapping = {
- "full_attention": create_block_causal_mask(indexes[0]),
+ "full_attention": create_block_causal_mask(
+ position_ids_from_indexes(indexes, 0)
+ ),
}
- self.current_index = indexes[0].max()
+ self.current_index = position_ids_from_indexes(indexes, 0).max()
else:
- self.current_index = indexes[0].max()
+ self.current_index = position_ids_from_indexes(indexes, 0).max()
# raise NotImplementedError('not isinstance(causal_mask_mapping := attention_mask, dict)')
# The sliding window alternating layers are not always activated depending on the config
diff --git a/python/sglang/multimodal_gen/runtime/models/sensenova_u1/neo_unify/modeling_qwen3_moe.py b/python/sglang/multimodal_gen/runtime/models/sensenova_u1/neo_unify/modeling_qwen3_moe.py
index 6625fab24..946fbaf82 100644
--- a/python/sglang/multimodal_gen/runtime/models/sensenova_u1/neo_unify/modeling_qwen3_moe.py
+++ b/python/sglang/multimodal_gen/runtime/models/sensenova_u1/neo_unify/modeling_qwen3_moe.py
@@ -21,8 +21,9 @@ from transformers.utils.deprecation import deprecate_kwarg
from .configuration_neo_chat import NEOMoELLMConfig
from .modeling_qwen3 import (
Qwen3Attention,
- Qwen3RMSNorm,
create_block_causal_mask,
+ make_qwen3_rms_norm,
+ position_ids_from_indexes,
)
from .transformers_compat import (
causal_mask_kwargs,
@@ -198,14 +199,16 @@ class Qwen3MoeDecoderLayer(GradientCheckpointingLayer):
),
)
- self.input_layernorm = Qwen3RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
- self.input_layernorm_mot_gen = Qwen3RMSNorm(
+ self.input_layernorm = make_qwen3_rms_norm(
config.hidden_size, eps=config.rms_norm_eps
)
- self.post_attention_layernorm = Qwen3RMSNorm(
+ self.input_layernorm_mot_gen = make_qwen3_rms_norm(
config.hidden_size, eps=config.rms_norm_eps
)
- self.post_attention_layernorm_mot_gen = Qwen3RMSNorm(
+ self.post_attention_layernorm = make_qwen3_rms_norm(
+ config.hidden_size, eps=config.rms_norm_eps
+ )
+ self.post_attention_layernorm_mot_gen = make_qwen3_rms_norm(
config.hidden_size, eps=config.rms_norm_eps
)
self.attention_type = config.layer_types[layer_idx]
@@ -422,8 +425,10 @@ class Qwen3MoeModel(Qwen3MoePreTrainedModel):
for layer_idx in range(config.num_hidden_layers)
]
)
- self.norm = Qwen3RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
- self.norm_mot_gen = Qwen3RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
+ self.norm = make_qwen3_rms_norm(config.hidden_size, eps=config.rms_norm_eps)
+ self.norm_mot_gen = make_qwen3_rms_norm(
+ config.hidden_size, eps=config.rms_norm_eps
+ )
self.gradient_checkpointing = False
self.has_sliding_layers = "sliding_attention" in self.config.layer_types
@@ -499,11 +504,13 @@ class Qwen3MoeModel(Qwen3MoePreTrainedModel):
)
else:
causal_mask_mapping = {
- "full_attention": create_block_causal_mask(indexes[0]),
+ "full_attention": create_block_causal_mask(
+ position_ids_from_indexes(indexes, 0)
+ ),
}
- self.current_index = indexes[0].max()
+ self.current_index = position_ids_from_indexes(indexes, 0).max()
else:
- self.current_index = indexes[0].max()
+ self.current_index = position_ids_from_indexes(indexes, 0).max()
hidden_states = inputs_embeds
diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/sensenova_u1/stages/generation.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/sensenova_u1/stages/generation.py
index 432c38cc8..d1a8c03dc 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/sensenova_u1/stages/generation.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/sensenova_u1/stages/generation.py
@@ -71,7 +71,37 @@ class SenseNovaU1GenerationStage(PipelineStage):
"SenseNova-U1 expects output expansion before generation; "
f"got num_outputs_per_prompt={batch.num_outputs_per_prompt}."
)
- seed = batch.seed[0] if isinstance(batch.seed, list) else int(batch.seed)
+ prompts = batch.prompt if isinstance(batch.prompt, list) else [batch.prompt]
+ batch_size = len(prompts)
+ if batch_size == 0:
+ raise ValueError(
+ "SenseNova-U1 dynamic batch must contain at least one prompt"
+ )
+ if batch_size > 1 and options.think_mode:
+ raise ValueError(
+ "SenseNova-U1 dynamic batching does not support think_mode"
+ )
+
+ dynamic_seeds = batch.extra.get("dynamic_batch_seeds")
+ if dynamic_seeds is None:
+ dynamic_seeds = batch.seed if isinstance(batch.seed, list) else [batch.seed]
+ elif not isinstance(dynamic_seeds, list):
+ dynamic_seeds = [dynamic_seeds]
+ seeds = []
+ for seed in dynamic_seeds:
+ if isinstance(seed, list):
+ if len(seed) != 1:
+ raise ValueError(
+ "SenseNova-U1 dynamic batching requires one seed per request"
+ )
+ seed = seed[0]
+ seeds.append(int(seed))
+ if len(seeds) != batch_size:
+ raise ValueError(
+ "SenseNova-U1 dynamic batch requires one seed per prompt; "
+ f"got {len(seeds)} seeds for {batch_size} prompts"
+ )
+ seed = seeds[0] if batch_size == 1 else seeds
out = self.model.t2i_generate(
self.tokenizer,
@@ -83,7 +113,7 @@ class SenseNovaU1GenerationStage(PipelineStage):
enable_timestep_shift=options.enable_timestep_shift,
cfg_interval=options.cfg_interval,
num_steps=int(batch.num_inference_steps),
- batch_size=1,
+ batch_size=batch_size,
t_eps=options.t_eps,
think_mode=options.think_mode,
seed=seed,
diff --git a/python/sglang/multimodal_gen/test/unit/test_sensenova_u1.py b/python/sglang/multimodal_gen/test/unit/test_sensenova_u1.py
index 0707d8a5b..1453f26eb 100644
--- a/python/sglang/multimodal_gen/test/unit/test_sensenova_u1.py
+++ b/python/sglang/multimodal_gen/test/unit/test_sensenova_u1.py
@@ -1,10 +1,14 @@
# SPDX-License-Identifier: Apache-2.0
import asyncio
import json
+import time
+from collections import deque
from types import SimpleNamespace
import pytest
import torch
+import torch.nn.functional as F
+from transformers.cache_utils import DynamicCache
from sglang.multimodal_gen.configs.pipeline_configs.sensenova_u1 import (
SenseNovaU1PipelineConfig,
@@ -26,6 +30,10 @@ from sglang.multimodal_gen.runtime.entrypoints.openai.utils import (
process_generation_batch,
)
from sglang.multimodal_gen.runtime.managers.gpu_worker import GPUWorker
+from sglang.multimodal_gen.runtime.managers.scheduler import Scheduler
+from sglang.multimodal_gen.runtime.models.sensenova_u1.neo_unify.configuration_neo_chat import (
+ NEOLLMConfig,
+)
from sglang.multimodal_gen.runtime.models.sensenova_u1.neo_unify.configuration_neo_vit import (
NEOVisionConfig,
)
@@ -33,7 +41,20 @@ from sglang.multimodal_gen.runtime.models.sensenova_u1.neo_unify.conversation im
get_conv_template,
)
from sglang.multimodal_gen.runtime.models.sensenova_u1.neo_unify.modeling_neo_chat import (
+ NEOChatModel,
+ _copy_right_aligned_prefix_bnsd,
_randn_with_seed,
+ prepare_flash_kv_cache,
+)
+from sglang.multimodal_gen.runtime.models.sensenova_u1.neo_unify.modeling_qwen3 import (
+ Qwen3Attention,
+ Qwen3MLP,
+ _flash_or_sdpa,
+ _sdpa_attn_func,
+ create_block_causal_mask,
+ make_qwen3_rms_norm,
+ npu_fia_available,
+ position_ids_from_indexes,
)
from sglang.multimodal_gen.runtime.pipelines_core.executors.pipeline_executor import (
PipelineExecutor,
@@ -48,6 +69,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.s
from sglang.multimodal_gen.runtime.platforms import current_platform
from sglang.multimodal_gen.runtime.server_args.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.perf_logger import MemorySnapshot
+from sglang.srt.layers.layernorm import RMSNorm
class _FakeSenseNovaModel:
@@ -56,7 +78,7 @@ class _FakeSenseNovaModel:
def t2i_generate(self, tokenizer, prompt, **kwargs):
self.call_kwargs = {"tokenizer": tokenizer, "prompt": prompt, **kwargs}
- return torch.tensor(
+ sample = torch.tensor(
[
[
[[-1.0, 0.0], [0.5, 1.0]],
@@ -65,6 +87,17 @@ class _FakeSenseNovaModel:
]
]
)
+ return sample.repeat(kwargs["batch_size"], 1, 1, 1)
+
+
+class _FakeTokenizer:
+ pad_token_id = None
+ eos_token_id = 2
+
+ def __call__(self, text, return_tensors):
+ del return_tensors
+ token_count = len(text.split()) + 1
+ return {"input_ids": torch.arange(1, token_count + 1).unsqueeze(0)}
class _RecordingTraceContext:
@@ -188,6 +221,299 @@ def test_sensenova_u1_randn_fallback_preserves_cpu_rng(monkeypatch):
assert torch.equal(torch.get_rng_state(), rng_state)
+def test_sensenova_u1_randn_supports_per_sample_seeds():
+ actual = _randn_with_seed(
+ (2, 3, 4), device=torch.device("cpu"), dtype=torch.float32, seed=[7, 19]
+ )
+ expected = torch.cat(
+ [
+ _randn_with_seed(
+ (1, 3, 4),
+ device=torch.device("cpu"),
+ dtype=torch.float32,
+ seed=seed,
+ )
+ for seed in (7, 19)
+ ]
+ )
+
+ assert torch.equal(actual, expected)
+
+
+def test_sensenova_u1_builds_padded_batched_text_inputs():
+ model = SimpleNamespace(device=torch.device("cpu"))
+
+ input_ids, indexes, attention_mask, valid_mask, prefix_lengths = (
+ NEOChatModel._build_t2i_text_inputs(
+ model, _FakeTokenizer(), ["short", "a much longer prompt"]
+ )
+ )
+
+ assert input_ids.shape == (2, 5)
+ assert indexes.shape == (2, 3, 5)
+ assert prefix_lengths.tolist() == [2, 5]
+ assert valid_mask.tolist() == [
+ [True, True, False, False, False],
+ [True, True, True, True, True],
+ ]
+ mask = attention_mask["full_attention"]
+ assert mask.shape == (2, 1, 5, 5)
+ assert torch.isneginf(mask[0, :, :, 2:]).all()
+ assert torch.isfinite(mask[0, :, :, :2]).any()
+
+
+def test_sensenova_u1_position_indexes_support_batched_inputs():
+ indexes = torch.tensor(
+ [
+ [[0, 1], [0, 0], [0, 0]],
+ [[4, 4], [0, 1], [0, 0]],
+ ]
+ )
+
+ assert torch.equal(position_ids_from_indexes(indexes, 0), indexes[:, 0])
+ assert torch.equal(
+ position_ids_from_indexes(indexes[0], 1), indexes[0, 1].unsqueeze(0)
+ )
+
+
+def test_sensenova_u1_singleton_text_matches_valid_batched_tokens():
+ model = SimpleNamespace(device=torch.device("cpu"))
+ tokenizer = _FakeTokenizer()
+ batched = NEOChatModel._build_t2i_text_inputs(
+ model, tokenizer, ["short", "a much longer prompt"]
+ )
+ for i, prompt in enumerate(["short", "a much longer prompt"]):
+ single = NEOChatModel._build_t2i_text_inputs(model, tokenizer, prompt)
+ length = single[0].shape[1]
+ assert torch.equal(batched[0][i, :length], single[0][0])
+ assert torch.equal(batched[1][i, :, :length], single[1])
+ assert torch.equal(
+ batched[2]["full_attention"][i, :, :length, :length],
+ single[2]["full_attention"][0],
+ )
+
+
+def test_sensenova_u1_block_causal_mask_rejects_padded_keys():
+ indexes = torch.tensor([[0, 1, 2], [0, 1, 2]])
+ valid = torch.tensor([[True, True, False], [True, True, True]])
+
+ mask = create_block_causal_mask(indexes, valid)
+
+ assert mask.shape == (2, 1, 3, 3)
+ assert torch.isneginf(mask[0, :, :, 2]).all()
+ assert mask[1, 0, 2, 2] == 0
+
+
+def test_sensenova_u1_builds_per_sample_image_indexes():
+ indexes = NEOChatModel._build_t2i_image_indexes(
+ SimpleNamespace(),
+ token_h=2,
+ token_w=2,
+ text_len=torch.tensor([2, 5]),
+ device=torch.device("cpu"),
+ )
+
+ assert indexes.shape == (2, 3, 4)
+ assert indexes[:, 0].tolist() == [[2, 2, 2, 2], [5, 5, 5, 5]]
+ assert indexes[:, 1].tolist() == [[0, 0, 1, 1], [0, 0, 1, 1]]
+ assert indexes[:, 2].tolist() == [[0, 1, 0, 1], [0, 1, 0, 1]]
+
+
+def test_sensenova_u1_compacts_variable_length_kv_before_attention():
+ generator = torch.Generator().manual_seed(29)
+ q = torch.randn(2, 3, 4, 8, generator=generator)
+ k = torch.randn(2, 8, 2, 8, generator=generator)
+ v = torch.randn(2, 8, 2, 8, generator=generator)
+ actual = _flash_or_sdpa(
+ q,
+ k,
+ v,
+ actual_seq_lengths_kv=[5, 8],
+ )
+
+ expected_short = _sdpa_attn_func(
+ q[:1],
+ torch.cat((k[:1, :2], k[:1, 5:]), dim=1),
+ torch.cat((v[:1, :2], v[:1, 5:]), dim=1),
+ )
+ expected_long = _sdpa_attn_func(q[1:], k[1:], v[1:])
+ torch.testing.assert_close(actual, torch.cat((expected_short, expected_long)))
+
+
+def test_sensenova_u1_sdpa_masks_padded_prefix_keys():
+ q = torch.tensor([[[[1.0, 0.0]]]])
+ k = torch.tensor([[[[1.0, 0.0]], [[0.0, 1.0]], [[1.0, 1.0]]]])
+ v = torch.tensor([[[[2.0, 0.0]], [[0.0, 4.0]], [[100.0, 100.0]]]])
+ key_mask = torch.tensor([[[[True, True, False]]]])
+
+ actual = _sdpa_attn_func(q, k, v, attention_mask=key_mask)
+ expected = _sdpa_attn_func(q, k[:, :2], v[:, :2])
+
+ torch.testing.assert_close(actual, expected)
+
+
+def test_sensenova_u1_right_aligns_bnsd_prefix_for_npu_fia():
+ source = torch.tensor(
+ [
+ [[[1], [2], [99], [99], [99]]],
+ [[[3], [4], [5], [6], [7]]],
+ ]
+ )
+ destination = torch.zeros(2, 1, 8, 1, dtype=source.dtype)
+
+ _copy_right_aligned_prefix_bnsd(destination, source, [2, 5])
+
+ assert destination[:, 0, :5, 0].tolist() == [
+ [0, 0, 0, 1, 2],
+ [3, 4, 5, 6, 7],
+ ]
+ assert destination[:, :, 5:].eq(0).all()
+
+
+@pytest.mark.parametrize("available", [False, True])
+def test_sensenova_u1_npu_fia_checks_operator_availability(monkeypatch, available):
+ namespace = SimpleNamespace()
+ if available:
+ namespace.npu_fused_infer_attention_score = object()
+ monkeypatch.setattr(torch.ops, "npu", namespace, raising=False)
+
+ assert npu_fia_available() is available
+
+
+@pytest.mark.parametrize(
+ ("is_npu", "uses_native"),
+ [(False, True), (True, False)],
+)
+def test_sensenova_u1_shared_rmsnorm_dispatch(monkeypatch, is_npu, uses_native):
+ monkeypatch.setattr(current_platform, "is_npu", lambda: is_npu)
+
+ norm = make_qwen3_rms_norm(64, eps=1e-6)
+
+ assert isinstance(norm, RMSNorm)
+ assert norm.cast_x_before_out_mul
+ assert (norm._forward_method == norm.forward_native) is uses_native
+
+
+@torch.no_grad()
+def test_sensenova_u1_fused_dense_mlp_matches_original(monkeypatch):
+ config = SimpleNamespace(
+ hidden_size=16,
+ intermediate_size=24,
+ hidden_act="silu",
+ )
+ with torch.random.fork_rng():
+ torch.manual_seed(37)
+ mlp = Qwen3MLP(config).eval()
+ hidden_states = torch.randn(2, 5, config.hidden_size)
+ expected = mlp(hidden_states)
+
+ monkeypatch.setattr(mlp, "_use_npu_fused_mlp", lambda _x: True)
+ monkeypatch.setattr(
+ torch.ops,
+ "npu",
+ SimpleNamespace(
+ npu_swiglu=lambda x, dim=-1: (
+ F.silu(x.chunk(2, dim=dim)[0]) * x.chunk(2, dim=dim)[1]
+ )
+ ),
+ raising=False,
+ )
+ actual = mlp(hidden_states)
+
+ torch.testing.assert_close(actual, expected)
+ assert set(mlp.state_dict()) == {
+ "gate_proj.weight",
+ "up_proj.weight",
+ "down_proj.weight",
+ }
+ assert (
+ mlp.gate_proj.weight.untyped_storage().data_ptr()
+ == mlp.up_proj.weight.untyped_storage().data_ptr()
+ )
+
+
+def test_sensenova_u1_batched_gqa_matches_unpadded_singletons():
+ generator = torch.Generator().manual_seed(17)
+ q = torch.randn(2, 3, 4, 8, generator=generator)
+ k = torch.randn(2, 8, 2, 8, generator=generator)
+ v = torch.randn(2, 8, 2, 8, generator=generator)
+ valid = torch.ones(2, 8, dtype=torch.bool)
+ valid[0, 2:5] = False
+ attention_mask = valid[:, None, None, :].expand(-1, -1, q.shape[1], -1)
+
+ actual = _sdpa_attn_func(q, k, v, attention_mask=attention_mask)
+
+ for i in range(2):
+ expected = _sdpa_attn_func(
+ q[i : i + 1], k[i : i + 1, valid[i]], v[i : i + 1, valid[i]]
+ )
+ torch.testing.assert_close(actual[i : i + 1], expected)
+
+
+@torch.no_grad()
+def test_sensenova_u1_prefix_and_denoise_attention_match_singletons():
+ config = NEOLLMConfig(
+ hidden_size=64,
+ intermediate_size=128,
+ num_hidden_layers=1,
+ num_attention_heads=4,
+ num_key_value_heads=2,
+ head_dim=16,
+ max_position_embeddings=128,
+ )
+ config._attn_implementation = "eager"
+ with torch.random.fork_rng():
+ torch.manual_seed(23)
+ attention = Qwen3Attention(config, layer_idx=0).eval()
+ generator = torch.Generator().manual_seed(31)
+ text = torch.randn(2, 5, 64, generator=generator)
+ image = torch.randn(2, 3, 64, generator=generator)
+ helper = SimpleNamespace(device=torch.device("cpu"))
+
+ def run(prefix, lengths, image_states):
+ batch_size, width, _ = prefix.shape
+ positions = torch.arange(width).expand(batch_size, -1)
+ indexes = torch.stack(
+ [positions, torch.zeros_like(positions), torch.zeros_like(positions)], dim=1
+ )
+ valid = positions < torch.tensor(lengths)[:, None]
+ cache = DynamicCache(config=config)
+ attention.forward_und(
+ prefix, indexes, create_block_causal_mask(positions, valid), cache
+ )
+ prefix_keys = cache.layers[0].keys.clone()
+ prepare_flash_kv_cache(
+ cache,
+ current_len=3,
+ batch_size=batch_size,
+ prefix_lengths=torch.tensor(lengths),
+ )
+ image_indexes = NEOChatModel._build_t2i_image_indexes(
+ helper, 1, 3, torch.tensor(lengths), torch.device("cpu")
+ )
+ outputs = []
+ for _ in range(2):
+ image_states, _ = attention.forward_gen(
+ image_states,
+ image_indexes,
+ None,
+ cache,
+ update_cache=False,
+ )
+ outputs.append(image_states)
+ torch.testing.assert_close(cache.layers[0].keys, prefix_keys)
+ return prefix_keys, outputs
+
+ keys, batched = run(text, [2, 5], image)
+ for i, length in enumerate([2, 5]):
+ single_keys, single = run(text[i : i + 1, :length], [length], image[i : i + 1])
+ torch.testing.assert_close(keys[i : i + 1, :, :length], single_keys)
+ for step in range(2):
+ torch.testing.assert_close(
+ batched[step][i : i + 1], single[step], atol=1e-5, rtol=1e-4
+ )
+
+
def test_sensenova_u1_randn_fallback_preserves_device_rng(monkeypatch):
device_type = current_platform.device_type
if not device_type or device_type == "cpu":
@@ -341,10 +667,169 @@ def test_sensenova_u1_accepts_openai_image_api_num_frames():
def test_sensenova_u1_scheduler_capabilities():
config = SenseNovaU1PipelineConfig()
- assert not config.supports_dynamic_batching()
+ assert config.supports_dynamic_batching()
assert config.supports_sequential_multi_output_inference()
+def _make_sensenova_u1_scheduler_request(
+ request_id: str, prompt: str, seed: int | list[int], **sampling_overrides
+) -> Req:
+ sampling = SenseNovaU1SamplingParams(
+ prompt=prompt,
+ seed=seed,
+ **sampling_overrides,
+ )
+ return Req(
+ request_id=request_id,
+ prompt=prompt,
+ seed=seed,
+ sampling_params=sampling,
+ extra=sampling.build_request_extra(),
+ )
+
+
+def test_sensenova_u1_batch_cost_tracks_resolution_steps_and_cfg():
+ config = SenseNovaU1PipelineConfig()
+ batch = SimpleNamespace(
+ width=1024,
+ height=1024,
+ num_inference_steps=5,
+ guidance_scale=4.0,
+ num_outputs_per_prompt=1,
+ )
+
+ assert config.estimate_request_cost(batch) == 32 * 32 * 5 * 2
+
+
+def test_sensenova_u1_multi_output_request_is_not_dynamically_batched():
+ scheduler = object.__new__(Scheduler)
+ scheduler.server_args = SimpleNamespace(pipeline_config=SenseNovaU1PipelineConfig())
+ sampling = SenseNovaU1SamplingParams(
+ prompt="a mountain lake", num_outputs_per_prompt=2
+ )
+ request = SimpleNamespace(
+ is_warmup=False,
+ realtime_session_id=None,
+ session=None,
+ prompt=sampling.prompt,
+ image_path=None,
+ return_file_paths_only=False,
+ num_outputs_per_prompt=2,
+ sampling_params=sampling,
+ )
+
+ assert not scheduler._can_dynamic_batch(request, request)
+ assert (
+ scheduler._get_dynamic_batch_reject_reason(request, request)
+ == "sequential_multi_output"
+ )
+
+
+def test_sensenova_u1_think_mode_request_is_dispatched_without_batching():
+ scheduler = object.__new__(Scheduler)
+ scheduler.server_args = SimpleNamespace(pipeline_config=SenseNovaU1PipelineConfig())
+ scheduler._batch_admission = SimpleNamespace(enabled=True)
+ scheduler._batch_metrics_enabled = False
+ request = _make_sensenova_u1_scheduler_request(
+ "request-0", "a mountain lake", 7, think_mode=True
+ )
+ scheduler.waiting_queue = deque([(b"identity", request, time.monotonic())])
+
+ assert not scheduler._can_dynamic_batch(request, request)
+ assert (
+ scheduler._get_dynamic_batch_reject_reason(request, request)
+ == "pipeline_request_unsupported"
+ )
+ items = scheduler.get_next_batch_to_run()
+ assert items is not None
+ assert items[0][0] == b"identity"
+ assert items[0][1] is request
+ assert not scheduler.waiting_queue
+
+
+@pytest.mark.parametrize(
+ "sampling_overrides",
+ [
+ {"guidance_scale": 1.0},
+ {"num_inference_steps": 25},
+ {"cfg_norm": "global"},
+ {"timestep_shift": 2.0},
+ {"t_eps": 0.01},
+ ],
+)
+def test_sensenova_u1_scheduler_rejects_heterogeneous_generation_options(
+ sampling_overrides,
+):
+ scheduler = object.__new__(Scheduler)
+ scheduler.server_args = SimpleNamespace(pipeline_config=SenseNovaU1PipelineConfig())
+ base = _make_sensenova_u1_scheduler_request("request-0", "first", 7)
+ candidate = _make_sensenova_u1_scheduler_request(
+ "request-1", "second", 19, **sampling_overrides
+ )
+
+ assert not scheduler._can_dynamic_batch(base, candidate)
+
+
+def test_sensenova_u1_scheduler_normalizes_single_output_seed_lists():
+ scheduler = object.__new__(Scheduler)
+ scheduler.server_args = SimpleNamespace(pipeline_config=SenseNovaU1PipelineConfig())
+ requests = [
+ _make_sensenova_u1_scheduler_request("request-0", "first", [7]),
+ _make_sensenova_u1_scheduler_request("request-1", "second", 19),
+ ]
+
+ merged = scheduler._try_merge_generation_reqs(requests)
+
+ assert merged.extra["dynamic_batch_seeds"] == [7, 19]
+
+
+def test_sensenova_u1_scheduler_merge_and_split_preserve_request_order():
+ scheduler = object.__new__(Scheduler)
+ scheduler.server_args = SimpleNamespace(pipeline_config=SenseNovaU1PipelineConfig())
+ requests = [
+ _make_sensenova_u1_scheduler_request(
+ "request-0",
+ "short",
+ 7,
+ output_path="/tmp/first",
+ output_file_name="first.png",
+ ),
+ _make_sensenova_u1_scheduler_request(
+ "request-1",
+ "a longer prompt",
+ 19,
+ output_path="/tmp/second",
+ output_file_name="second.png",
+ ),
+ ]
+ merged = scheduler._try_merge_generation_reqs(requests)
+ assert merged.prompt == ["short", "a longer prompt"]
+ assert merged.extra["dynamic_batch_seeds"] == [7, 19]
+ expected_paths = [request.output_file_path() for request in requests]
+ assert merged.extra["dynamic_batch_output_paths"] == expected_paths
+ assert requests[0].prompt == "short"
+ outputs = scheduler._split_batched_output(
+ OutputBatch(
+ output=[torch.tensor([7]), torch.tensor([19])],
+ output_file_paths=expected_paths,
+ ),
+ requests,
+ )
+ assert [output.output[0].item() for output in outputs] == [7, 19]
+ assert [output.output_file_paths for output in outputs] == [
+ [path] for path in expected_paths
+ ]
+ assert (
+ scheduler._split_batched_output(
+ OutputBatch(output=[torch.tensor([7])]), requests
+ )
+ is None
+ )
+ requests[1].sampling_params.width = 1024
+ del requests[1]._dynamic_batch_sig
+ assert scheduler._try_merge_generation_reqs(requests) is None
+
+
def test_sensenova_u1_rejects_multi_gpu_during_arg_validation():
config = SenseNovaU1PipelineConfig()
@@ -626,6 +1111,71 @@ def test_sensenova_u1_generation_stage_uses_sglang_params_and_single_model_batch
assert model.call_kwargs["seed"] == 123
+def test_sensenova_u1_generation_stage_passes_dynamic_batch_inputs():
+ sampling = SenseNovaU1SamplingParams(
+ prompt="first prompt",
+ width=1024,
+ height=1024,
+ guidance_scale=4.0,
+ num_inference_steps=5,
+ seed=7,
+ )
+ batch = SimpleNamespace(
+ prompt=["first prompt", "a longer second prompt"],
+ width=sampling.width,
+ height=sampling.height,
+ guidance_scale=sampling.guidance_scale,
+ num_inference_steps=sampling.num_inference_steps,
+ seed=sampling.seed,
+ num_outputs_per_prompt=1,
+ extra={
+ **sampling.build_request_extra(),
+ "dynamic_batch_seeds": [7, 19],
+ },
+ metrics=None,
+ )
+ model = _FakeSenseNovaModel()
+ stage = SenseNovaU1GenerationStage(model=model, tokenizer="tok")
+
+ output = stage.forward(batch, server_args=SimpleNamespace())
+
+ assert len(output.output) == 2
+ assert model.call_kwargs["prompt"] == [
+ "first prompt",
+ "a longer second prompt",
+ ]
+ assert model.call_kwargs["batch_size"] == 2
+ assert model.call_kwargs["seed"] == [7, 19]
+
+
+def test_sensenova_u1_generation_stage_rejects_batched_think_mode():
+ sampling = SenseNovaU1SamplingParams(
+ prompt="first prompt",
+ width=1024,
+ height=1024,
+ think_mode=True,
+ )
+ batch = SimpleNamespace(
+ prompt=["first prompt", "second prompt"],
+ width=sampling.width,
+ height=sampling.height,
+ guidance_scale=sampling.guidance_scale,
+ num_inference_steps=sampling.num_inference_steps,
+ seed=sampling.seed,
+ num_outputs_per_prompt=1,
+ extra={
+ **sampling.build_request_extra(),
+ "dynamic_batch_seeds": [7, 19],
+ },
+ metrics=None,
+ )
+
+ with pytest.raises(ValueError, match="think_mode"):
+ SenseNovaU1GenerationStage(
+ model=_FakeSenseNovaModel(), tokenizer="tok"
+ ).forward(batch, server_args=SimpleNamespace())
+
+
def test_sensenova_u1_multi_output_request_expands_before_generation_stage():
sampling = SenseNovaU1SamplingParams(
prompt="a mountain lake",