Split #32584 into 1/2: [LoRA] Guard DP-attention idle forwards against stale LoRA batch state (#32707)

This commit is contained in:
Ethan (Yusheng) Su
2026-07-31 14:23:01 -07:00
committed by GitHub
parent 301ea43f35
commit 7e996a5d0d
8 changed files with 67 additions and 17 deletions
+14 -1
View File
@@ -1,4 +1,4 @@
from typing import Tuple, Union from typing import Optional, Tuple, Union
import torch import torch
import triton import triton
@@ -26,6 +26,10 @@ class BaseLoRABackend(LoRABackendLmHeadMixing):
def __init__(self, max_loras_per_batch: int, device: torch.device): def __init__(self, max_loras_per_batch: int, device: torch.device):
self.max_loras_per_batch = max_loras_per_batch self.max_loras_per_batch = max_loras_per_batch
self.device = device self.device = device
# Set by prepare_lora_batch() before each forward; cleared by
# reset_batch_state() on DP-attention idle forwards. None means "no
# batch prepared" — the LoRA layers read it to skip LoRA application.
self.batch_info: Optional[LoRABatchInfo] = None
self.init_lm_head_config() self.init_lm_head_config()
self._is_moe_lora = False self._is_moe_lora = False
# Static metadata read by prefill-CUDA-graph kernels, refreshed in # Static metadata read by prefill-CUDA-graph kernels, refreshed in
@@ -35,6 +39,15 @@ class BaseLoRABackend(LoRABackendLmHeadMixing):
self.prefill_cuda_graph_max_bs: int | None = None self.prefill_cuda_graph_max_bs: int | None = None
self.prefill_cuda_graph_max_tokens: int | None = None self.prefill_cuda_graph_max_tokens: int | None = None
def reset_batch_state(self):
"""Idle-forward counterpart of prepare_lora_batch(): clears all
per-batch metadata. batch_info=None is the master "no batch
prepared" signal that the layer guards (lora_active) read."""
self.batch_info = None
self.lm_head_batch_info = None
self.lm_head_pass_batch_infos = None
self._lm_head_pass_idx = None
def run_lora_a_embedding( def run_lora_a_embedding(
self, self,
input_ids: torch.Tensor, input_ids: torch.Tensor,
@@ -32,6 +32,13 @@ class TritonLoRABackend(BaseLoRABackend):
**kwargs, **kwargs,
): ):
super().__init__(max_loras_per_batch, device) super().__init__(max_loras_per_batch, device)
# Merged-segment variant of batch_info; set alongside it in
# prepare_lora_batch and cleared together in reset_batch_state.
self.sgemm_batch_info: Optional[LoRABatchInfo] = None
def reset_batch_state(self):
super().reset_batch_state()
self.sgemm_batch_info = None
def run_lora_a_embedding( def run_lora_a_embedding(
self, self,
@@ -55,7 +62,12 @@ class TritonLoRABackend(BaseLoRABackend):
"""Return the sgemm batch_info (merged segments when available).""" """Return the sgemm batch_info (merged segments when available)."""
if pruned_batch_info is not None: if pruned_batch_info is not None:
return pruned_batch_info return pruned_batch_info
return getattr(self, "sgemm_batch_info", None) or self.batch_info assert self.batch_info is not None, (
"LoRA kernel invoked with no prepared batch (DP-attention idle "
"forward?). Gate the caller on lora_active, as in "
"sglang/srt/lora/layers.py forwards."
)
return self.sgemm_batch_info or self.batch_info
def run_lora_a_sgemm( def run_lora_a_sgemm(
self, self,
@@ -46,8 +46,6 @@ def _get_state(
if not hasattr(attn_module.kv_b_proj, "A_buffer"): if not hasattr(attn_module.kv_b_proj, "A_buffer"):
return None return None
lora_backend = attn_module.kv_b_proj.lora_backend lora_backend = attn_module.kv_b_proj.lora_backend
if not hasattr(lora_backend, "batch_info"):
return None
batch_info = lora_backend.batch_info batch_info = lora_backend.batch_info
if batch_info is None: if batch_info is None:
return None return None
+26 -8
View File
@@ -45,10 +45,24 @@ class BaseLayerWithLoRA(nn.Module):
self.weight = self.base_layer.weight self.weight = self.base_layer.weight
if hasattr(self.base_layer, "bias") and self.base_layer.bias is not None: if hasattr(self.base_layer, "bias") and self.base_layer.bias is not None:
self.bias = self.base_layer.bias self.bias = self.base_layer.bias
# Forward reduce_results so model code that inspects it on the module
# (e.g. DeepseekV2AttentionMLA's `assert not self.o_proj.reduce_results`
# on DP-attention idle forwards) keeps working when the layer is
# LoRA-wrapped.
if hasattr(self.base_layer, "reduce_results"):
self.reduce_results = self.base_layer.reduce_results
def forward(self, x: torch.Tensor): def forward(self, x: torch.Tensor):
return self.base_layer.forward(x) return self.base_layer.forward(x)
@property
def lora_active(self) -> bool:
"""True when this layer has LoRA buffers set AND the current forward
has LoRA batch metadata. batch_info is None on DP-attention idle
forwards (see LoRAManager.prepare_lora_batch), so idle forwards take
the base path."""
return self.set_lora and self.lora_backend.batch_info is not None
def set_lora_info(self, *args): def set_lora_info(self, *args):
pass pass
@@ -211,8 +225,9 @@ class VocabParallelEmbeddingWithLoRA(BaseLayerWithLoRA):
): ):
base_output = self.extra_token_embedding(input_, base_output) base_output = self.extra_token_embedding(input_, base_output)
# Apply LoRA if configured # Apply LoRA if configured; DP-attention idle forwards take the base
if self.set_lora: # path (see lora_active).
if self.lora_active:
# The backend's run_lora_a_embedding now handles both regular # The backend's run_lora_a_embedding now handles both regular
# and extra tokens efficiently with CUDA graph support # and extra tokens efficiently with CUDA graph support
base_output = self.apply_lora(base_output, input_, batch_info) base_output = self.apply_lora(base_output, input_, batch_info)
@@ -377,8 +392,7 @@ class ParallelLMHeadWithLoRA(BaseLayerWithLoRA):
hidden_states, self.weight, bias=getattr(self.base_layer, "bias", None) hidden_states, self.weight, bias=getattr(self.base_layer, "bias", None)
) )
# Apply LoRA if set if self.lora_active:
if self.set_lora:
base_output = self.apply_lora(base_output, hidden_states) base_output = self.apply_lora(base_output, hidden_states)
return base_output return base_output
@@ -467,7 +481,7 @@ class ColumnParallelLinearWithLoRA(BaseLayerWithLoRA):
self.base_layer, input_, bias self.base_layer, input_, bias
) )
if self.set_lora: if self.lora_active:
output_parallel = self.apply_lora(output_parallel, input_) output_parallel = self.apply_lora(output_parallel, input_)
if self.base_layer.gather_output: if self.base_layer.gather_output:
@@ -773,7 +787,8 @@ class RowParallelLinearWithLoRA(BaseLayerWithLoRA):
and not should_skip_mlp_all_reduce() and not should_skip_mlp_all_reduce()
) )
if self.set_lora and should_reduce: lora_active = self.lora_active
if lora_active and should_reduce:
lora_a_output = self.lora_backend.run_lora_a_sgemm( lora_a_output = self.lora_backend.run_lora_a_sgemm(
input_parallel, self.A_buffer input_parallel, self.A_buffer
) )
@@ -787,7 +802,7 @@ class RowParallelLinearWithLoRA(BaseLayerWithLoRA):
base_output=output_, base_output=output_,
) )
else: else:
if self.set_lora: if lora_active:
output_parallel = self.apply_lora(output_parallel, input_parallel) output_parallel = self.apply_lora(output_parallel, input_parallel)
if should_reduce: if should_reduce:
output_ = tensor_model_parallel_all_reduce(output_parallel) output_ = tensor_model_parallel_all_reduce(output_parallel)
@@ -886,7 +901,7 @@ class ReplicatedLinearWithLoRA(BaseLayerWithLoRA):
def forward(self, x: torch.Tensor): def forward(self, x: torch.Tensor):
bias = self.base_layer.bias if not self.base_layer.skip_bias_add else None bias = self.base_layer.bias if not self.base_layer.skip_bias_add else None
output = self.base_layer.quant_method.apply(self.base_layer, x, bias) output = self.base_layer.quant_method.apply(self.base_layer, x, bias)
if self.set_lora: if self.lora_active:
output = self.apply_lora(output, x) output = self.apply_lora(output, x)
output_bias = self.base_layer.bias if self.base_layer.skip_bias_add else None output_bias = self.base_layer.bias if self.base_layer.skip_bias_add else None
return output, output_bias return output, output_bias
@@ -1090,6 +1105,9 @@ class FusedMoEWithLoRA(BaseLayerWithLoRA):
1. After gate_up projection, before activation 1. After gate_up projection, before activation
2. After down projection, before final reduction 2. After down projection, before final reduction
""" """
# DP-attention idle forward: no batch_info, run the base MoE path.
if self.lora_backend.batch_info is None:
return self.base_layer.forward(hidden_states, topk_output, **kwargs)
# Build LoRA info for this batch # Build LoRA info for this batch
lora_info = self._get_lora_info() lora_info = self._get_lora_info()
+7
View File
@@ -414,6 +414,13 @@ class LoRAManager:
if callable(notify): if callable(notify):
notify(slot_ids) notify(slot_ids)
def reset_lora_batch(self):
"""Clear per-batch LoRA state. Called instead of prepare_lora_batch()
on DP-attention idle forwards (zero local tokens), so the LoRA layers
take the base path instead of reading the previous batch's stale
metadata."""
self.lora_backend.reset_batch_state()
def prepare_lora_batch(self, forward_batch: ForwardBatch): def prepare_lora_batch(self, forward_batch: ForwardBatch):
# set up batch info shared by all lora modules # set up batch info shared by all lora modules
bs = forward_batch.batch_size bs = forward_batch.batch_size
@@ -36,7 +36,7 @@ def qkv_proj_lora_forward(self, input_: torch.Tensor):
AND base_output, so it runs after the rejoin on the main stream. AND base_output, so it runs after the rejoin on the main stream.
""" """
if ( if (
not self.set_lora not self.lora_active
or not is_two_stream_active(input_) or not is_two_stream_active(input_)
or not supports_two_stream_dense_lora(self.A_buffer_qkv, self.B_buffer_qkv) or not supports_two_stream_dense_lora(self.A_buffer_qkv, self.B_buffer_qkv)
): ):
@@ -104,7 +104,7 @@ def row_parallel_lora_forward(
input_parallel = splitted_input[tp_rank].contiguous() input_parallel = splitted_input[tp_rank].contiguous()
if ( if (
not self.set_lora not self.lora_active
or not is_two_stream_active(input_parallel) or not is_two_stream_active(input_parallel)
or not supports_two_stream_dense_lora(self.A_buffer, self.B_buffer) or not supports_two_stream_dense_lora(self.A_buffer, self.B_buffer)
): ):
@@ -176,7 +176,7 @@ def column_parallel_lora_forward(self, input_: torch.Tensor):
for non-decode batches or when LoRA isn't set on this layer. for non-decode batches or when LoRA isn't set on this layer.
""" """
if ( if (
not self.set_lora not self.lora_active
or not is_two_stream_active(input_) or not is_two_stream_active(input_)
or not supports_two_stream_dense_lora(self.A_buffer, self.B_buffer) or not supports_two_stream_dense_lora(self.A_buffer, self.B_buffer)
): ):
@@ -231,7 +231,7 @@ def replicated_lora_forward(self, x: torch.Tensor):
the main after the rejoin. Falls back to the saved-original otherwise. the main after the rejoin. Falls back to the saved-original otherwise.
""" """
if ( if (
not self.set_lora not self.lora_active
or not is_two_stream_active(x) or not is_two_stream_active(x)
or not supports_two_stream_dense_lora(self.A_buffer, self.B_buffer) or not supports_two_stream_dense_lora(self.A_buffer, self.B_buffer)
): ):
@@ -31,7 +31,7 @@ from sglang.srt.lora.trtllm_lora_temp import (
def merged_column_lora_forward(self, input_: torch.Tensor): def merged_column_lora_forward(self, input_: torch.Tensor):
"""O9 — side-stream LoRA-A shrink ‖ base merged-column GEMM.""" """O9 — side-stream LoRA-A shrink ‖ base merged-column GEMM."""
if ( if (
not self.set_lora not self.lora_active
or not is_two_stream_active(input_) or not is_two_stream_active(input_)
or not supports_two_stream_dense_lora(self.A_buffer, self.B_buffer) or not supports_two_stream_dense_lora(self.A_buffer, self.B_buffer)
): ):
@@ -849,6 +849,8 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
if ret.forward_mode.is_idle(): if ret.forward_mode.is_idle():
ret.positions = torch.empty((0,), dtype=torch.int64, device=device) ret.positions = torch.empty((0,), dtype=torch.int64, device=device)
if model_runner.server_args.enable_lora:
model_runner.lora_manager.reset_lora_batch()
return ret return ret
# Override the positions with diffusion LLM or spec_info # Override the positions with diffusion LLM or spec_info