[2/n] lora - Shared outer experts and support qwen3_30b_a3b_instruct (#21466)
Co-authored-by: Baizhou Zhang <sobereddiezhang@gmail.com>
This commit is contained in:
co-authored by
Baizhou Zhang
parent
f4505e2ee3
commit
3c91ebdf55
@@ -1,4 +1,4 @@
|
|||||||
from typing import Optional
|
from typing import Dict, Optional, Union
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
@@ -711,6 +711,9 @@ class FusedMoEWithLoRA(BaseLayerWithLoRA):
|
|||||||
# initializes FusedMoE with its own moe_runner for base path
|
# initializes FusedMoE with its own moe_runner for base path
|
||||||
super().__init__(base_layer, lora_backend)
|
super().__init__(base_layer, lora_backend)
|
||||||
|
|
||||||
|
self.experts_shared_outer_loras: bool = False
|
||||||
|
self.quant_method = base_layer.quant_method
|
||||||
|
|
||||||
self.tp_size = getattr(base_layer, "moe_tp_size", 1)
|
self.tp_size = getattr(base_layer, "moe_tp_size", 1)
|
||||||
self.tp_rank = getattr(base_layer, "moe_tp_rank", 0)
|
self.tp_rank = getattr(base_layer, "moe_tp_rank", 0)
|
||||||
self.intermediate_size_per_partition = getattr(
|
self.intermediate_size_per_partition = getattr(
|
||||||
@@ -782,6 +785,7 @@ class FusedMoEWithLoRA(BaseLayerWithLoRA):
|
|||||||
adapter_enabled=adapter_enabled,
|
adapter_enabled=adapter_enabled,
|
||||||
max_lora_rank=max_lora_rank,
|
max_lora_rank=max_lora_rank,
|
||||||
num_experts=self.base_layer.num_experts,
|
num_experts=self.base_layer.num_experts,
|
||||||
|
experts_shared_outer_loras=self.experts_shared_outer_loras,
|
||||||
tp_size=self.tp_size,
|
tp_size=self.tp_size,
|
||||||
tp_rank=self.tp_rank,
|
tp_rank=self.tp_rank,
|
||||||
hidden_size=getattr(self.base_layer, "hidden_size", 0),
|
hidden_size=getattr(self.base_layer, "hidden_size", 0),
|
||||||
@@ -839,34 +843,79 @@ class FusedMoEWithLoRA(BaseLayerWithLoRA):
|
|||||||
return B
|
return B
|
||||||
|
|
||||||
def slice_moe_lora_a_weights(
|
def slice_moe_lora_a_weights(
|
||||||
self, A: torch.Tensor, tp_rank: int, target_module: str
|
self,
|
||||||
) -> torch.Tensor:
|
A: Union[torch.Tensor, Dict[int, torch.Tensor]],
|
||||||
|
tp_rank: int,
|
||||||
|
target_module: str,
|
||||||
|
):
|
||||||
"""Slice LoRA A weights for MoE with TP.
|
"""Slice LoRA A weights for MoE with TP.
|
||||||
|
|
||||||
|
Accepts:
|
||||||
|
- 2D tensor [rank, hidden] (single expert)
|
||||||
|
- 3D tensor [num_experts_or_1, rank, hidden]
|
||||||
|
- dict {expert_id: 2D tensor}
|
||||||
|
|
||||||
Per-expert weight shapes:
|
Per-expert weight shapes:
|
||||||
gate_up_proj_moe A: [rank, hidden_size] — input is full hidden_states, no slice
|
gate_up_proj_moe A: [rank, hidden_size] — input is full hidden_states, no slice
|
||||||
down_proj_moe A: [rank, intermediate_size] — input is sharded intermediate
|
down_proj_moe A: [rank, intermediate_size] — input is sharded intermediate
|
||||||
"""
|
"""
|
||||||
if self.tp_size <= 1:
|
if self.tp_size <= 1:
|
||||||
return A
|
return A
|
||||||
if target_module == "down_proj_moe":
|
if target_module != "down_proj_moe":
|
||||||
shard_size = self.intermediate_size_per_partition
|
return A
|
||||||
start = tp_rank * shard_size
|
if isinstance(A, dict):
|
||||||
end = start + shard_size
|
return {
|
||||||
return A[:, start:end].contiguous()
|
eid: self._slice_moe_a(w, tp_rank, target_module)
|
||||||
return A
|
for eid, w in A.items()
|
||||||
|
}
|
||||||
|
return self._slice_moe_a(A, tp_rank, target_module)
|
||||||
|
|
||||||
|
def _slice_moe_a(
|
||||||
|
self, A: torch.Tensor, tp_rank: int, target_module: str
|
||||||
|
) -> torch.Tensor:
|
||||||
|
shard_size = self.intermediate_size_per_partition
|
||||||
|
start = tp_rank * shard_size
|
||||||
|
end = start + shard_size
|
||||||
|
return A[..., start:end].contiguous()
|
||||||
|
|
||||||
def slice_moe_lora_b_weights(
|
def slice_moe_lora_b_weights(
|
||||||
self, B: torch.Tensor, tp_rank: int, target_module: str
|
self,
|
||||||
) -> torch.Tensor:
|
B: Union[torch.Tensor, Dict[int, torch.Tensor]],
|
||||||
|
tp_rank: int,
|
||||||
|
target_module: str,
|
||||||
|
):
|
||||||
"""Slice LoRA B weights for MoE with TP.
|
"""Slice LoRA B weights for MoE with TP.
|
||||||
|
|
||||||
|
Accepts:
|
||||||
|
- 2D tensor [output_dim, rank] (single expert)
|
||||||
|
- 3D tensor [num_experts_or_1, output_dim, rank]
|
||||||
|
- dict {expert_id: 2D tensor}
|
||||||
|
|
||||||
Per-expert weight shapes:
|
Per-expert weight shapes:
|
||||||
gate_up_proj_moe B: [intermediate_size*2, rank] — output matches sharded base w13
|
gate_up_proj_moe B: [intermediate_size*2, rank] — output matches sharded base w13
|
||||||
down_proj_moe B: [hidden_size, rank] — output is all-reduced, no slice
|
down_proj_moe B: [hidden_size, rank] — output is all-reduced, no slice
|
||||||
"""
|
"""
|
||||||
if self.tp_size <= 1:
|
if self.tp_size <= 1:
|
||||||
return B
|
return B
|
||||||
|
if target_module != "gate_up_proj_moe":
|
||||||
|
return B
|
||||||
|
if isinstance(B, dict):
|
||||||
|
return {
|
||||||
|
eid: self._slice_moe_b_2d(w, tp_rank, target_module)
|
||||||
|
for eid, w in B.items()
|
||||||
|
}
|
||||||
|
if isinstance(B, torch.Tensor) and B.dim() == 3:
|
||||||
|
return torch.stack(
|
||||||
|
[
|
||||||
|
self._slice_moe_b_2d(B[i], tp_rank, target_module)
|
||||||
|
for i in range(B.shape[0])
|
||||||
|
]
|
||||||
|
)
|
||||||
|
return self._slice_moe_b_2d(B, tp_rank, target_module)
|
||||||
|
|
||||||
|
def _slice_moe_b_2d(
|
||||||
|
self, B: torch.Tensor, tp_rank: int, target_module: str
|
||||||
|
) -> torch.Tensor:
|
||||||
if target_module == "gate_up_proj_moe":
|
if target_module == "gate_up_proj_moe":
|
||||||
shard_size = self.intermediate_size_per_partition
|
shard_size = self.intermediate_size_per_partition
|
||||||
start = tp_rank * shard_size
|
start = tp_rank * shard_size
|
||||||
|
|||||||
@@ -137,6 +137,8 @@ class LoRAAdapter(nn.Module):
|
|||||||
for layer in self.layers:
|
for layer in self.layers:
|
||||||
weight_names = list(layer.weights.keys())
|
weight_names = list(layer.weights.keys())
|
||||||
self.normalize_qkv_proj(weight_names, layer.weights)
|
self.normalize_qkv_proj(weight_names, layer.weights)
|
||||||
|
self._rename_expert_w_to_proj(layer.weights)
|
||||||
|
weight_names = list(layer.weights.keys())
|
||||||
self.normalize_gate_up_proj(weight_names, layer.weights)
|
self.normalize_gate_up_proj(weight_names, layer.weights)
|
||||||
|
|
||||||
def normalize_qkv_proj(
|
def normalize_qkv_proj(
|
||||||
@@ -192,6 +194,23 @@ class LoRAAdapter(nn.Module):
|
|||||||
weights[qkv_name] = weights[qkv_name].repeat(3, 1)
|
weights[qkv_name] = weights[qkv_name].repeat(3, 1)
|
||||||
# else: no-op as LoRA B weight is already stacked.
|
# else: no-op as LoRA B weight is already stacked.
|
||||||
|
|
||||||
|
def _rename_expert_w_to_proj(self, weights: Dict[str, torch.Tensor]):
|
||||||
|
"""Rename w1 -> gate_proj, w3 -> up_proj, w2 -> down_proj so that
|
||||||
|
normalize_gate_up_proj can stack them into gate_up_proj."""
|
||||||
|
renames = {}
|
||||||
|
for name in list(weights.keys()):
|
||||||
|
new_name = name
|
||||||
|
if ".w1." in name:
|
||||||
|
new_name = name.replace(".w1.", ".gate_proj.")
|
||||||
|
elif ".w3." in name:
|
||||||
|
new_name = name.replace(".w3.", ".up_proj.")
|
||||||
|
elif ".w2." in name:
|
||||||
|
new_name = name.replace(".w2.", ".down_proj.")
|
||||||
|
if new_name != name:
|
||||||
|
renames[name] = new_name
|
||||||
|
for old_name, new_name in renames.items():
|
||||||
|
weights[new_name] = weights.pop(old_name)
|
||||||
|
|
||||||
def normalize_gate_up_proj(
|
def normalize_gate_up_proj(
|
||||||
self, weight_names: List[str], weights: Dict[str, torch.Tensor]
|
self, weight_names: List[str], weights: Dict[str, torch.Tensor]
|
||||||
):
|
):
|
||||||
@@ -206,8 +225,9 @@ class LoRAAdapter(nn.Module):
|
|||||||
f"Received backend: {self.lora_backend.name}. Please verify your backend configuration "
|
f"Received backend: {self.lora_backend.name}. Please verify your backend configuration "
|
||||||
f"or consider implementing custom initialization logic for other backends."
|
f"or consider implementing custom initialization logic for other backends."
|
||||||
)
|
)
|
||||||
|
cat_dim = weights[weight_name].dim() - 2
|
||||||
weights[gate_up_name] = torch.cat(
|
weights[gate_up_name] = torch.cat(
|
||||||
(weights[weight_name], weights[up_name]), 0
|
(weights[weight_name], weights[up_name]), cat_dim
|
||||||
)
|
)
|
||||||
weights.pop(weight_name)
|
weights.pop(weight_name)
|
||||||
if up_name in weights:
|
if up_name in weights:
|
||||||
@@ -216,7 +236,10 @@ class LoRAAdapter(nn.Module):
|
|||||||
# If gate_up_proj is already stacked, we normalize it following the SGL convention
|
# If gate_up_proj is already stacked, we normalize it following the SGL convention
|
||||||
gate_up_name = weight_name
|
gate_up_name = weight_name
|
||||||
if "lora_A" in weight_name:
|
if "lora_A" in weight_name:
|
||||||
weights[gate_up_name] = weights[gate_up_name].repeat(2, 1)
|
ndim = weights[gate_up_name].dim()
|
||||||
|
repeat_dims = [1] * ndim
|
||||||
|
repeat_dims[ndim - 2] = 2
|
||||||
|
weights[gate_up_name] = weights[gate_up_name].repeat(*repeat_dims)
|
||||||
# else: no-op as LoRA B weight is already stacked.
|
# else: no-op as LoRA B weight is already stacked.
|
||||||
|
|
||||||
def pin_weights_in_cpu(self):
|
def pin_weights_in_cpu(self):
|
||||||
|
|||||||
@@ -78,8 +78,10 @@ class LoRAManager:
|
|||||||
server_args.enable_lora_overlap_loading
|
server_args.enable_lora_overlap_loading
|
||||||
)
|
)
|
||||||
|
|
||||||
# Store eviction policy from server args
|
|
||||||
self.eviction_policy = server_args.lora_eviction_policy
|
self.eviction_policy = server_args.lora_eviction_policy
|
||||||
|
self._experts_shared_outer_override: Optional[bool] = (
|
||||||
|
server_args.experts_shared_outer_loras
|
||||||
|
)
|
||||||
|
|
||||||
# LoRA backend for running sgemm kernels
|
# LoRA backend for running sgemm kernels
|
||||||
logger.info(f"Using {lora_backend} as backend of LoRA kernels.")
|
logger.info(f"Using {lora_backend} as backend of LoRA kernels.")
|
||||||
@@ -303,23 +305,33 @@ class LoRAManager:
|
|||||||
if isinstance(module, FusedMoEWithLoRA) and all(
|
if isinstance(module, FusedMoEWithLoRA) and all(
|
||||||
x in self.target_modules for x in ["gate_up_proj", "down_proj"]
|
x in self.target_modules for x in ["gate_up_proj", "down_proj"]
|
||||||
):
|
):
|
||||||
|
gate_up_key = (
|
||||||
|
"gate_up_proj_moe"
|
||||||
|
if "gate_up_proj_moe" in self.memory_pool.A_buffer
|
||||||
|
else "gate_up_proj"
|
||||||
|
)
|
||||||
|
down_key = (
|
||||||
|
"down_proj_moe"
|
||||||
|
if "down_proj_moe" in self.memory_pool.A_buffer
|
||||||
|
else "down_proj"
|
||||||
|
)
|
||||||
gate_up_a = self.memory_pool.get_tensor(
|
gate_up_a = self.memory_pool.get_tensor(
|
||||||
target_module="gate_up_proj_moe",
|
target_module=gate_up_key,
|
||||||
layer_id=layer_id,
|
layer_id=layer_id,
|
||||||
lora_type=LoRAType.LORA_A,
|
lora_type=LoRAType.LORA_A,
|
||||||
)
|
)
|
||||||
gate_up_b = self.memory_pool.get_tensor(
|
gate_up_b = self.memory_pool.get_tensor(
|
||||||
target_module="gate_up_proj_moe",
|
target_module=gate_up_key,
|
||||||
layer_id=layer_id,
|
layer_id=layer_id,
|
||||||
lora_type=LoRAType.LORA_B,
|
lora_type=LoRAType.LORA_B,
|
||||||
)
|
)
|
||||||
down_a = self.memory_pool.get_tensor(
|
down_a = self.memory_pool.get_tensor(
|
||||||
target_module="down_proj_moe",
|
target_module=down_key,
|
||||||
layer_id=layer_id,
|
layer_id=layer_id,
|
||||||
lora_type=LoRAType.LORA_A,
|
lora_type=LoRAType.LORA_A,
|
||||||
)
|
)
|
||||||
down_b = self.memory_pool.get_tensor(
|
down_b = self.memory_pool.get_tensor(
|
||||||
target_module="down_proj_moe",
|
target_module=down_key,
|
||||||
layer_id=layer_id,
|
layer_id=layer_id,
|
||||||
lora_type=LoRAType.LORA_B,
|
lora_type=LoRAType.LORA_B,
|
||||||
)
|
)
|
||||||
@@ -387,6 +399,16 @@ class LoRAManager:
|
|||||||
target_modules=target_modules,
|
target_modules=target_modules,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if self._experts_shared_outer_override is not None:
|
||||||
|
self.experts_shared_outer_loras = self._experts_shared_outer_override
|
||||||
|
else:
|
||||||
|
self.experts_shared_outer_loras = self._detect_shared_outer_loras()
|
||||||
|
if self.experts_shared_outer_loras:
|
||||||
|
logger.info(
|
||||||
|
"Shared outer LoRA mode enabled: gate_up lora_A and "
|
||||||
|
"down lora_B will be shared across experts (expert_dim=1)."
|
||||||
|
)
|
||||||
|
|
||||||
self.init_lora_modules()
|
self.init_lora_modules()
|
||||||
self.init_memory_pool()
|
self.init_memory_pool()
|
||||||
self.update_lora_info()
|
self.update_lora_info()
|
||||||
@@ -412,6 +434,26 @@ class LoRAManager:
|
|||||||
f"Failed to load LoRA adapter {lora_ref.lora_name}: {result.error_message}"
|
f"Failed to load LoRA adapter {lora_ref.lora_name}: {result.error_message}"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _detect_shared_outer_loras(self) -> bool:
|
||||||
|
"""Auto-detect shared outer LoRA format from loaded adapter weights.
|
||||||
|
|
||||||
|
MoE adapters with shared outer experts store 3D tensors where
|
||||||
|
dim[0]=1 indicates weights shared across all experts, while
|
||||||
|
dim[0]=num_experts indicates per-expert weights.
|
||||||
|
Returns True if gate_up lora_A has expert_dim=1 (shared).
|
||||||
|
"""
|
||||||
|
for adapter in self.loras.values():
|
||||||
|
for layer in adapter.layers:
|
||||||
|
for name, weight in layer.weights.items():
|
||||||
|
if (
|
||||||
|
"gate_up_proj" in name
|
||||||
|
and "lora_A" in name
|
||||||
|
and weight.dim() == 3
|
||||||
|
):
|
||||||
|
return weight.shape[0] == 1
|
||||||
|
break
|
||||||
|
return False
|
||||||
|
|
||||||
def init_lora_shapes(
|
def init_lora_shapes(
|
||||||
self,
|
self,
|
||||||
max_lora_rank: Optional[int] = None,
|
max_lora_rank: Optional[int] = None,
|
||||||
@@ -589,6 +631,7 @@ class LoRAManager:
|
|||||||
base_model=self.base_model,
|
base_model=self.base_model,
|
||||||
eviction_policy=self.eviction_policy,
|
eviction_policy=self.eviction_policy,
|
||||||
lora_added_tokens_size=self.lora_added_tokens_size,
|
lora_added_tokens_size=self.lora_added_tokens_size,
|
||||||
|
experts_shared_outer_loras=self.experts_shared_outer_loras,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Initializing memory pool with base model
|
# Initializing memory pool with base model
|
||||||
@@ -683,11 +726,10 @@ class LoRAManager:
|
|||||||
)
|
)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
# Temporarily workaround for FusedMoE layer
|
|
||||||
if isinstance(module, FusedMoE) and all(
|
if isinstance(module, FusedMoE) and all(
|
||||||
x in self.target_modules for x in ["gate_up_proj", "down_proj"]
|
x in self.target_modules for x in ["gate_up_proj", "down_proj"]
|
||||||
):
|
):
|
||||||
layer_id = get_layer_id(module_name)
|
layer_id = get_layer_id(module_name)
|
||||||
self.lora_modules[layer_id][module_name] = self.set_lora_module(
|
lora_module = self.set_lora_module(module_name, module)
|
||||||
module_name, module
|
lora_module.experts_shared_outer_loras = self.experts_shared_outer_loras
|
||||||
)
|
self.lora_modules[layer_id][module_name] = lora_module
|
||||||
|
|||||||
@@ -71,17 +71,22 @@ if _is_cuda or _is_hip or _is_xpu:
|
|||||||
class LoRAInfo:
|
class LoRAInfo:
|
||||||
"""LoRA weights and dispatch info for MoE computation."""
|
"""LoRA weights and dispatch info for MoE computation."""
|
||||||
|
|
||||||
# LoRA weights: [num_loras, num_experts, dim1, dim2]
|
# LoRA weights: [num_loras, num_experts_or_1, dim1, dim2]
|
||||||
|
# When experts_shared_outer_loras=True:
|
||||||
|
# gate_up_lora_a: [num_loras, 1, max_rank, hidden_dim] (shared)
|
||||||
|
# down_lora_b: [num_loras, 1, hidden_dim, max_rank] (shared)
|
||||||
gate_up_lora_a_weights: (
|
gate_up_lora_a_weights: (
|
||||||
torch.Tensor
|
torch.Tensor
|
||||||
) # [num_loras, num_experts, max_rank, hidden_dim]
|
) # [num_loras, num_experts_or_1, max_rank, hidden_dim]
|
||||||
gate_up_lora_b_weights: (
|
gate_up_lora_b_weights: (
|
||||||
torch.Tensor
|
torch.Tensor
|
||||||
) # [num_loras, num_experts, gate_up_dim, max_rank]
|
) # [num_loras, num_experts, gate_up_dim, max_rank]
|
||||||
down_lora_a_weights: (
|
down_lora_a_weights: (
|
||||||
torch.Tensor
|
torch.Tensor
|
||||||
) # [num_loras, num_experts, max_rank, intermediate_dim]
|
) # [num_loras, num_experts, max_rank, intermediate_dim]
|
||||||
down_lora_b_weights: torch.Tensor # [num_loras, num_experts, hidden_dim, max_rank]
|
down_lora_b_weights: (
|
||||||
|
torch.Tensor
|
||||||
|
) # [num_loras, num_experts_or_1, hidden_dim, max_rank]
|
||||||
|
|
||||||
# Indice pointers of each segment in shape (num_segments + 1, )
|
# Indice pointers of each segment in shape (num_segments + 1, )
|
||||||
seg_indptr: torch.Tensor
|
seg_indptr: torch.Tensor
|
||||||
@@ -95,6 +100,7 @@ class LoRAInfo:
|
|||||||
max_lora_rank: int # Maximum LoRA rank across all adapters
|
max_lora_rank: int # Maximum LoRA rank across all adapters
|
||||||
|
|
||||||
num_experts: int
|
num_experts: int
|
||||||
|
experts_shared_outer_loras: bool = False
|
||||||
|
|
||||||
fully_sharded: bool = False
|
fully_sharded: bool = False
|
||||||
tp_size: int = 1
|
tp_size: int = 1
|
||||||
@@ -469,16 +475,11 @@ class TritonRunnerCoreWithLoRA(TritonRunnerCore):
|
|||||||
|
|
||||||
r = lora_info.max_lora_rank
|
r = lora_info.max_lora_rank
|
||||||
gate_up_a = lora_info.gate_up_lora_a_weights
|
gate_up_a = lora_info.gate_up_lora_a_weights
|
||||||
|
if lora_info.experts_shared_outer_loras:
|
||||||
|
gate_up_a = gate_up_a.expand(-1, lora_info.num_experts, -1, -1)
|
||||||
gate_up_b = lora_info.gate_up_lora_b_weights
|
gate_up_b = lora_info.gate_up_lora_b_weights
|
||||||
inter_size = gate_up_b.shape[2] // 2
|
inter_size = gate_up_b.shape[2] // 2
|
||||||
|
|
||||||
# Split packed gate_up weights into separate gate and up slices.
|
|
||||||
# gate_up_lora_a has shape [max_loras, num_experts, 2*r, hidden_dim]
|
|
||||||
# where the first r rows are gate_lora_a and the next r are up_lora_a.
|
|
||||||
# gate_up_lora_b has shape [max_loras, num_experts, 2*inter_size, r]
|
|
||||||
# where the first inter_size rows are gate_lora_b and the rest up_lora_b.
|
|
||||||
# Using num_slices=2 lets the kernel handle gate and up independently,
|
|
||||||
# keeping the rank dimension at r so shrink and expand both match.
|
|
||||||
lora_a_stacked = [gate_up_a[:, :, :r, :], gate_up_a[:, :, r : 2 * r, :]]
|
lora_a_stacked = [gate_up_a[:, :, :r, :], gate_up_a[:, :, r : 2 * r, :]]
|
||||||
lora_b_stacked = [
|
lora_b_stacked = [
|
||||||
gate_up_b[:, :, :inter_size, :],
|
gate_up_b[:, :, :inter_size, :],
|
||||||
@@ -542,8 +543,12 @@ class TritonRunnerCoreWithLoRA(TritonRunnerCore):
|
|||||||
if lora_info.max_lora_rank == 0:
|
if lora_info.max_lora_rank == 0:
|
||||||
return
|
return
|
||||||
|
|
||||||
|
down_lora_b = lora_info.down_lora_b_weights
|
||||||
|
if lora_info.experts_shared_outer_loras:
|
||||||
|
down_lora_b = down_lora_b.expand(-1, lora_info.num_experts, -1, -1)
|
||||||
|
|
||||||
lora_a_stacked = [lora_info.down_lora_a_weights]
|
lora_a_stacked = [lora_info.down_lora_a_weights]
|
||||||
lora_b_stacked = [lora_info.down_lora_b_weights]
|
lora_b_stacked = [down_lora_b]
|
||||||
|
|
||||||
if lora_info.fully_sharded and lora_info.tp_size > 1:
|
if lora_info.fully_sharded and lora_info.tp_size > 1:
|
||||||
shard_size = lora_info.hidden_size // lora_info.tp_size
|
shard_size = lora_info.hidden_size // lora_info.tp_size
|
||||||
|
|||||||
@@ -60,6 +60,7 @@ class LoRAMemoryPool:
|
|||||||
base_model: torch.nn.Module,
|
base_model: torch.nn.Module,
|
||||||
eviction_policy: str,
|
eviction_policy: str,
|
||||||
lora_added_tokens_size: int,
|
lora_added_tokens_size: int,
|
||||||
|
experts_shared_outer_loras: bool = False,
|
||||||
):
|
):
|
||||||
self.base_hf_config: AutoConfig = base_hf_config
|
self.base_hf_config: AutoConfig = base_hf_config
|
||||||
self.num_layer: int = base_hf_config.num_hidden_layers
|
self.num_layer: int = base_hf_config.num_hidden_layers
|
||||||
@@ -70,6 +71,7 @@ class LoRAMemoryPool:
|
|||||||
self.lora_added_tokens_size: int = lora_added_tokens_size
|
self.lora_added_tokens_size: int = lora_added_tokens_size
|
||||||
self.max_lora_rank: int = max_lora_rank
|
self.max_lora_rank: int = max_lora_rank
|
||||||
self.target_modules: Set[str] = target_modules
|
self.target_modules: Set[str] = target_modules
|
||||||
|
self.experts_shared_outer_loras: bool = experts_shared_outer_loras
|
||||||
|
|
||||||
# Initialize eviction policy
|
# Initialize eviction policy
|
||||||
self.eviction_policy = get_eviction_policy(eviction_policy)
|
self.eviction_policy = get_eviction_policy(eviction_policy)
|
||||||
@@ -140,6 +142,18 @@ class LoRAMemoryPool:
|
|||||||
"""Check if module is part of MoE experts."""
|
"""Check if module is part of MoE experts."""
|
||||||
return "moe" in module_name
|
return "moe" in module_name
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _get_num_experts(base_model: torch.nn.Module) -> int:
|
||||||
|
cfg = base_model.config
|
||||||
|
if hasattr(cfg, "get_text_config"):
|
||||||
|
cfg = cfg.get_text_config()
|
||||||
|
return (
|
||||||
|
getattr(cfg, "num_experts", None)
|
||||||
|
or getattr(cfg, "num_local_experts", None)
|
||||||
|
or getattr(cfg, "n_routed_experts", None)
|
||||||
|
or 1
|
||||||
|
)
|
||||||
|
|
||||||
def _get_standard_shape(
|
def _get_standard_shape(
|
||||||
self,
|
self,
|
||||||
module_name: str,
|
module_name: str,
|
||||||
@@ -178,10 +192,13 @@ class LoRAMemoryPool:
|
|||||||
input_dim = divide(input_dim, self.tp_size)
|
input_dim = divide(input_dim, self.tp_size)
|
||||||
|
|
||||||
if self.is_moe_module(module_name):
|
if self.is_moe_module(module_name):
|
||||||
num_experts = base_model.config.num_experts
|
num_experts = self._get_num_experts(base_model)
|
||||||
|
expert_dim = num_experts
|
||||||
|
if self.experts_shared_outer_loras and module_name == "gate_up_proj_moe":
|
||||||
|
expert_dim = 1
|
||||||
return (
|
return (
|
||||||
self.max_loras_per_batch,
|
self.max_loras_per_batch,
|
||||||
num_experts,
|
expert_dim,
|
||||||
max_lora_dim * c,
|
max_lora_dim * c,
|
||||||
input_dim,
|
input_dim,
|
||||||
)
|
)
|
||||||
@@ -228,8 +245,11 @@ class LoRAMemoryPool:
|
|||||||
|
|
||||||
# Check if MoE module and return appropriate shape
|
# Check if MoE module and return appropriate shape
|
||||||
if self.is_moe_module(module_name):
|
if self.is_moe_module(module_name):
|
||||||
num_experts = base_model.config.num_experts
|
num_experts = self._get_num_experts(base_model)
|
||||||
return (self.max_loras_per_batch, num_experts, output_dim, max_lora_dim)
|
expert_dim = num_experts
|
||||||
|
if self.experts_shared_outer_loras and module_name == "down_proj_moe":
|
||||||
|
expert_dim = 1
|
||||||
|
return (self.max_loras_per_batch, expert_dim, output_dim, max_lora_dim)
|
||||||
else:
|
else:
|
||||||
return (self.max_loras_per_batch, output_dim, max_lora_dim)
|
return (self.max_loras_per_batch, output_dim, max_lora_dim)
|
||||||
|
|
||||||
@@ -264,32 +284,33 @@ class LoRAMemoryPool:
|
|||||||
target_modules: Set[str],
|
target_modules: Set[str],
|
||||||
get_lora_shape_fn: Callable[[str, torch.nn.Module, int, int], Tuple[int]],
|
get_lora_shape_fn: Callable[[str, torch.nn.Module, int, int], Tuple[int]],
|
||||||
):
|
):
|
||||||
# Check if model has both shared experts and MoE experts
|
cfg = base_model.config
|
||||||
|
if hasattr(cfg, "get_text_config"):
|
||||||
|
cfg = cfg.get_text_config()
|
||||||
has_shared_experts = (
|
has_shared_experts = (
|
||||||
hasattr(base_model.config, "shared_expert_intermediate_size")
|
hasattr(cfg, "shared_expert_intermediate_size")
|
||||||
and base_model.config.shared_expert_intermediate_size > 0
|
and cfg.shared_expert_intermediate_size > 0
|
||||||
)
|
) or (getattr(cfg, "n_shared_experts", 0) or 0) > 0
|
||||||
has_moe = getattr(base_model.config, "num_experts", 1) > 1
|
has_moe = self._get_num_experts(base_model) > 1
|
||||||
|
|
||||||
# Shape functions automatically handle both 3D (standard) and 4D (MoE)
|
# Shape functions automatically handle both 3D (standard) and 4D (MoE)
|
||||||
target_modules = target_modules - set(EMBEDDING_NAMES)
|
target_modules = target_modules - set(EMBEDDING_NAMES)
|
||||||
for module_name in target_modules:
|
for module_name in target_modules:
|
||||||
# Special handling for ambiguous target modules that can be in different contexts
|
# Special handling for ambiguous target modules that can be in different contexts
|
||||||
ambiguous_modules = {"gate_up_proj", "down_proj"}
|
ambiguous_modules = {"gate_up_proj", "down_proj"}
|
||||||
if module_name in ambiguous_modules and has_shared_experts and has_moe:
|
if module_name in ambiguous_modules and has_moe:
|
||||||
# Allocate separate buffers for shared and MoE contexts
|
# Allocate shared expert version (3D) only when model has shared experts
|
||||||
# Shared expert version (3D)
|
if has_shared_experts:
|
||||||
shared_key = module_name
|
buffer[module_name] = [
|
||||||
buffer[shared_key] = [
|
torch.zeros(
|
||||||
torch.empty(
|
get_lora_shape_fn(
|
||||||
get_lora_shape_fn(
|
module_name, base_model, self.max_lora_rank, idx
|
||||||
module_name, base_model, self.max_lora_rank, idx
|
),
|
||||||
),
|
dtype=self.dtype,
|
||||||
dtype=self.dtype,
|
device=device,
|
||||||
device=device,
|
)
|
||||||
)
|
for idx in range(self.num_layer)
|
||||||
for idx in range(self.num_layer)
|
]
|
||||||
]
|
|
||||||
|
|
||||||
# MoE expert version (4D)
|
# MoE expert version (4D)
|
||||||
moe_key = f"{module_name}_moe"
|
moe_key = f"{module_name}_moe"
|
||||||
@@ -521,8 +542,8 @@ class LoRAMemoryPool:
|
|||||||
expert_match = re.search(r"experts\.(\d+)\.", name)
|
expert_match = re.search(r"experts\.(\d+)\.", name)
|
||||||
|
|
||||||
if expert_match:
|
if expert_match:
|
||||||
|
# Per-expert MoE weight — 2D tensors, one per expert
|
||||||
target_module = target_module + "_moe"
|
target_module = target_module + "_moe"
|
||||||
# MoE weight - multiple tensors per module (one per expert)
|
|
||||||
if temp_A_buffer[target_module] is None:
|
if temp_A_buffer[target_module] is None:
|
||||||
temp_A_buffer[target_module] = {}
|
temp_A_buffer[target_module] = {}
|
||||||
temp_B_buffer[target_module] = {}
|
temp_B_buffer[target_module] = {}
|
||||||
@@ -532,8 +553,15 @@ class LoRAMemoryPool:
|
|||||||
temp_A_buffer[target_module][expert_id] = weights
|
temp_A_buffer[target_module][expert_id] = weights
|
||||||
else:
|
else:
|
||||||
temp_B_buffer[target_module][expert_id] = weights
|
temp_B_buffer[target_module][expert_id] = weights
|
||||||
|
elif "experts" in name and weights.dim() == 3:
|
||||||
|
# Shared outer MoE weight — 3D tensor [expert_dim, rank, hidden]
|
||||||
|
target_module = target_module + "_moe"
|
||||||
|
if "lora_A" in name:
|
||||||
|
temp_A_buffer[target_module] = weights
|
||||||
|
else:
|
||||||
|
temp_B_buffer[target_module] = weights
|
||||||
else:
|
else:
|
||||||
# Standard weight - single tensor per module
|
# Standard weight — single tensor per module
|
||||||
if "lora_A" in name:
|
if "lora_A" in name:
|
||||||
temp_A_buffer[target_module] = weights
|
temp_A_buffer[target_module] = weights
|
||||||
else:
|
else:
|
||||||
@@ -549,20 +577,18 @@ class LoRAMemoryPool:
|
|||||||
if isinstance(module, FusedMoEWithLoRA):
|
if isinstance(module, FusedMoEWithLoRA):
|
||||||
moe_target_modules = ["gate_up_proj_moe", "down_proj_moe"]
|
moe_target_modules = ["gate_up_proj_moe", "down_proj_moe"]
|
||||||
for target_module in moe_target_modules:
|
for target_module in moe_target_modules:
|
||||||
if temp_A_buffer[target_module] is None:
|
if temp_A_buffer.get(target_module) is not None:
|
||||||
continue
|
temp_A_buffer[target_module] = (
|
||||||
|
|
||||||
for expert_id in temp_A_buffer[target_module].keys():
|
|
||||||
temp_A_buffer[target_module][expert_id] = (
|
|
||||||
module.slice_moe_lora_a_weights(
|
module.slice_moe_lora_a_weights(
|
||||||
temp_A_buffer[target_module][expert_id],
|
temp_A_buffer[target_module],
|
||||||
self.tp_rank,
|
self.tp_rank,
|
||||||
target_module,
|
target_module,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
temp_B_buffer[target_module][expert_id] = (
|
if temp_B_buffer.get(target_module) is not None:
|
||||||
|
temp_B_buffer[target_module] = (
|
||||||
module.slice_moe_lora_b_weights(
|
module.slice_moe_lora_b_weights(
|
||||||
temp_B_buffer[target_module][expert_id],
|
temp_B_buffer[target_module],
|
||||||
self.tp_rank,
|
self.tp_rank,
|
||||||
target_module,
|
target_module,
|
||||||
)
|
)
|
||||||
@@ -587,22 +613,42 @@ class LoRAMemoryPool:
|
|||||||
temp_B_buffer[target_module], self.tp_rank
|
temp_B_buffer[target_module], self.tp_rank
|
||||||
)
|
)
|
||||||
|
|
||||||
# Load weights into buffers (handles both 3D standard and 4D MoE)
|
|
||||||
for name, weights in temp_A_buffer.items():
|
for name, weights in temp_A_buffer.items():
|
||||||
c = get_stacked_multiply(name)
|
c = get_stacked_multiply(name)
|
||||||
target_buffer = self.A_buffer[name][layer_id]
|
target_buffer = self.A_buffer[name][layer_id]
|
||||||
|
|
||||||
if name in ["gate_up_proj_moe", "down_proj_moe"]:
|
if name in ["gate_up_proj_moe", "down_proj_moe"]:
|
||||||
# MoE: multiple tensors per module (one per expert)
|
if self.experts_shared_outer_loras and name == "gate_up_proj_moe":
|
||||||
for expert_id, expert_weight in weights.items():
|
if isinstance(weights, torch.Tensor) and weights.dim() == 3:
|
||||||
# Buffer shape: [num_loras, num_experts, max_rank, hidden_dim]
|
buffer_view = target_buffer[
|
||||||
buffer_view = target_buffer[
|
buffer_id, 0, : lora_rank * c, :
|
||||||
buffer_id, expert_id, : lora_rank * c, :
|
]
|
||||||
]
|
load_lora_weight_tensor(buffer_view, weights[0])
|
||||||
load_lora_weight_tensor(buffer_view, expert_weight)
|
elif isinstance(weights, dict) and len(weights) > 0:
|
||||||
|
rep = next(iter(weights.values()))
|
||||||
|
buffer_view = target_buffer[
|
||||||
|
buffer_id, 0, : lora_rank * c, :
|
||||||
|
]
|
||||||
|
load_lora_weight_tensor(buffer_view, rep)
|
||||||
|
else:
|
||||||
|
raise ValueError(
|
||||||
|
f"Unexpected weight format for shared outer gate_up_proj_moe lora_A: "
|
||||||
|
f"type={type(weights)}, "
|
||||||
|
f"shape={weights.shape if isinstance(weights, torch.Tensor) else 'N/A'}"
|
||||||
|
)
|
||||||
|
elif isinstance(weights, torch.Tensor) and weights.dim() == 3:
|
||||||
|
for eid in range(weights.shape[0]):
|
||||||
|
buffer_view = target_buffer[
|
||||||
|
buffer_id, eid, : lora_rank * c, :
|
||||||
|
]
|
||||||
|
load_lora_weight_tensor(buffer_view, weights[eid])
|
||||||
|
elif isinstance(weights, dict):
|
||||||
|
for expert_id, expert_weight in weights.items():
|
||||||
|
buffer_view = target_buffer[
|
||||||
|
buffer_id, expert_id, : lora_rank * c, :
|
||||||
|
]
|
||||||
|
load_lora_weight_tensor(buffer_view, expert_weight)
|
||||||
else:
|
else:
|
||||||
# Standard: single tensor per module
|
|
||||||
c = get_stacked_multiply(name)
|
|
||||||
buffer_view = target_buffer[buffer_id, : lora_rank * c, :]
|
buffer_view = target_buffer[buffer_id, : lora_rank * c, :]
|
||||||
load_lora_weight_tensor(buffer_view, weights)
|
load_lora_weight_tensor(buffer_view, weights)
|
||||||
|
|
||||||
@@ -610,18 +656,42 @@ class LoRAMemoryPool:
|
|||||||
target_buffer = self.B_buffer[name][layer_id]
|
target_buffer = self.B_buffer[name][layer_id]
|
||||||
|
|
||||||
if name in ["gate_up_proj_moe", "down_proj_moe"]:
|
if name in ["gate_up_proj_moe", "down_proj_moe"]:
|
||||||
# MoE: multiple tensors per module (one per expert)
|
if self.experts_shared_outer_loras and name == "down_proj_moe":
|
||||||
for expert_id, expert_weight in weights.items():
|
if isinstance(weights, torch.Tensor) and weights.dim() == 3:
|
||||||
# Buffer shape: [num_loras, num_experts, intermediate_dim, max_rank]
|
buffer_view = target_buffer[buffer_id, 0, :, :lora_rank]
|
||||||
buffer_view = target_buffer[buffer_id, expert_id, :, :lora_rank]
|
w = weights[0]
|
||||||
|
if w is not None:
|
||||||
weight_to_load = expert_weight
|
w = w * lora_adapter.scaling
|
||||||
if weight_to_load is not None:
|
load_lora_weight_tensor(buffer_view, w)
|
||||||
weight_to_load = weight_to_load * lora_adapter.scaling
|
elif isinstance(weights, dict) and len(weights) > 0:
|
||||||
|
rep = next(iter(weights.values()))
|
||||||
load_lora_weight_tensor(buffer_view, weight_to_load)
|
buffer_view = target_buffer[buffer_id, 0, :, :lora_rank]
|
||||||
|
if rep is not None:
|
||||||
|
rep = rep * lora_adapter.scaling
|
||||||
|
load_lora_weight_tensor(buffer_view, rep)
|
||||||
|
else:
|
||||||
|
raise ValueError(
|
||||||
|
f"Unexpected weight format for shared outer down_proj_moe lora_B: "
|
||||||
|
f"type={type(weights)}, "
|
||||||
|
f"shape={weights.shape if isinstance(weights, torch.Tensor) else 'N/A'}"
|
||||||
|
)
|
||||||
|
elif isinstance(weights, torch.Tensor) and weights.dim() == 3:
|
||||||
|
for eid in range(weights.shape[0]):
|
||||||
|
buffer_view = target_buffer[buffer_id, eid, :, :lora_rank]
|
||||||
|
w = weights[eid]
|
||||||
|
if w is not None:
|
||||||
|
w = w * lora_adapter.scaling
|
||||||
|
load_lora_weight_tensor(buffer_view, w)
|
||||||
|
elif isinstance(weights, dict):
|
||||||
|
for expert_id, expert_weight in weights.items():
|
||||||
|
buffer_view = target_buffer[
|
||||||
|
buffer_id, expert_id, :, :lora_rank
|
||||||
|
]
|
||||||
|
w = expert_weight
|
||||||
|
if w is not None:
|
||||||
|
w = w * lora_adapter.scaling
|
||||||
|
load_lora_weight_tensor(buffer_view, w)
|
||||||
else:
|
else:
|
||||||
# Standard: single tensor per module
|
|
||||||
buffer_view = target_buffer[buffer_id, :, :lora_rank]
|
buffer_view = target_buffer[buffer_id, :, :lora_rank]
|
||||||
load_lora_weight_tensor(buffer_view, weights)
|
load_lora_weight_tensor(buffer_view, weights)
|
||||||
|
|
||||||
|
|||||||
@@ -87,6 +87,7 @@ def _sgemm_lora_b_kernel(
|
|||||||
)
|
)
|
||||||
|
|
||||||
# Iterate to compute the block in output matrix
|
# Iterate to compute the block in output matrix
|
||||||
|
n_mask = n_offset[None, :] < N
|
||||||
partial_sum = tl.zeros((BLOCK_S, BLOCK_N), dtype=tl.float32)
|
partial_sum = tl.zeros((BLOCK_S, BLOCK_N), dtype=tl.float32)
|
||||||
for k in range(0, tl.cdiv(K, BLOCK_K)):
|
for k in range(0, tl.cdiv(K, BLOCK_K)):
|
||||||
x_tile = tl.load(
|
x_tile = tl.load(
|
||||||
@@ -96,7 +97,7 @@ def _sgemm_lora_b_kernel(
|
|||||||
)
|
)
|
||||||
w_tile = tl.load(
|
w_tile = tl.load(
|
||||||
w_ptrs,
|
w_ptrs,
|
||||||
mask=(k_offset[:, None] < K - k * BLOCK_K),
|
mask=(k_offset[:, None] < K - k * BLOCK_K) & n_mask,
|
||||||
other=0.0,
|
other=0.0,
|
||||||
)
|
)
|
||||||
partial_sum += tl.dot(x_tile, w_tile)
|
partial_sum += tl.dot(x_tile, w_tile)
|
||||||
@@ -110,8 +111,8 @@ def _sgemm_lora_b_kernel(
|
|||||||
output_ptr = (output + seg_start * output_stride_0) + (
|
output_ptr = (output + seg_start * output_stride_0) + (
|
||||||
s_offset[:, None] * output_stride_0 + n_offset[None, :] * output_stride_1
|
s_offset[:, None] * output_stride_0 + n_offset[None, :] * output_stride_1
|
||||||
)
|
)
|
||||||
output_mask = s_offset[:, None] < seg_len
|
output_mask = (s_offset[:, None] < seg_len) & n_mask
|
||||||
partial_sum += tl.load(output_ptr, mask=output_mask)
|
partial_sum += tl.load(output_ptr, mask=output_mask, other=0.0)
|
||||||
tl.store(output_ptr, partial_sum, mask=output_mask)
|
tl.store(output_ptr, partial_sum, mask=output_mask)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -464,6 +464,7 @@ class ServerArgs:
|
|||||||
lora_eviction_policy: str = "lru"
|
lora_eviction_policy: str = "lru"
|
||||||
lora_backend: str = "csgmv"
|
lora_backend: str = "csgmv"
|
||||||
max_lora_chunk_size: Optional[int] = 16
|
max_lora_chunk_size: Optional[int] = 16
|
||||||
|
experts_shared_outer_loras: Optional[bool] = None
|
||||||
|
|
||||||
# Kernel backend
|
# Kernel backend
|
||||||
attention_backend: Optional[str] = None
|
attention_backend: Optional[str] = None
|
||||||
@@ -4595,6 +4596,14 @@ class ServerArgs:
|
|||||||
choices=[16, 32, 64, 128],
|
choices=[16, 32, 64, 128],
|
||||||
help="Maximum chunk size for the ChunkedSGMV LoRA backend. Only used when --lora-backend is 'csgmv'. Choosing a larger value might improve performance.",
|
help="Maximum chunk size for the ChunkedSGMV LoRA backend. Only used when --lora-backend is 'csgmv'. Choosing a larger value might improve performance.",
|
||||||
)
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--experts-shared-outer-loras",
|
||||||
|
default=ServerArgs.experts_shared_outer_loras,
|
||||||
|
action="store_true",
|
||||||
|
help="Force shared outer LoRA mode for MoE models. "
|
||||||
|
"When set, w1/w3 lora_A and w2 lora_B are shared across experts "
|
||||||
|
"(expert_dim=1). By default this is auto-detected from adapter weights.",
|
||||||
|
)
|
||||||
|
|
||||||
# Kernel backend
|
# Kernel backend
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
|
|||||||
@@ -0,0 +1,151 @@
|
|||||||
|
# Copyright 2023-2025 SGLang Team
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
# ==============================================================================
|
||||||
|
|
||||||
|
"""
|
||||||
|
Regression test for Qwen3-30B-A3B-Instruct-2507 LoRA logprob accuracy.
|
||||||
|
|
||||||
|
Compares SGLang LoRA logprobs against reference training logprobs from a
|
||||||
|
pre-computed dataset. The LoRA adapter and reference data are downloaded from:
|
||||||
|
https://huggingface.co/datasets/yushengsu/lora-diff-Qwen3-30B-A3B-Instruct-2507
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
python -m unittest test_lora_qwen3_30b_a3b_instruct_2507_logprob_diff
|
||||||
|
"""
|
||||||
|
|
||||||
|
import multiprocessing as mp
|
||||||
|
import os
|
||||||
|
import unittest
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from huggingface_hub import snapshot_download
|
||||||
|
|
||||||
|
import sglang as sgl
|
||||||
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
register_cuda_ci(
|
||||||
|
est_time=300,
|
||||||
|
suite="stage-c-test-4-gpu-b200",
|
||||||
|
)
|
||||||
|
|
||||||
|
BASE_MODEL = "Qwen/Qwen3-30B-A3B-Instruct-2507"
|
||||||
|
LORA_HF_REPO = "yushengsu/lora-diff-Qwen3-30B-A3B-Instruct-2507"
|
||||||
|
LORA_BACKEND = "triton"
|
||||||
|
MAX_LORA_RANK = 32
|
||||||
|
TP_SIZE = 4
|
||||||
|
DISABLE_CUDA_GRAPH = True
|
||||||
|
MOE_RUNNER_BACKEND = "triton"
|
||||||
|
EXPERTS_SHARED_OUTER_LORAS = True
|
||||||
|
PREFILL_ATTENTION_BACKEND = "fa4"
|
||||||
|
DECODE_ATTENTION_BACKEND = "fa4"
|
||||||
|
|
||||||
|
KL_THRESHOLD = 5e-3
|
||||||
|
|
||||||
|
|
||||||
|
def kl_v2(a, b):
|
||||||
|
a = torch.tensor(a) if not torch.is_tensor(a) else a
|
||||||
|
b = torch.tensor(b) if not torch.is_tensor(b) else b
|
||||||
|
return (((a - b) ** 2) * 0.5).mean().item()
|
||||||
|
|
||||||
|
|
||||||
|
def get_prompt_logprobs(engine, input_ids, lora_path):
|
||||||
|
out = engine.generate(
|
||||||
|
input_ids=input_ids,
|
||||||
|
sampling_params={"max_new_tokens": 0, "temperature": 0.0},
|
||||||
|
return_logprob=True,
|
||||||
|
logprob_start_len=0,
|
||||||
|
lora_path=lora_path,
|
||||||
|
)
|
||||||
|
return [logprob for logprob, _, _ in out["meta_info"]["input_token_logprobs"]][1:]
|
||||||
|
|
||||||
|
|
||||||
|
class TestLoRAQwen3_30B_A3B_Instruct_2507_LogprobDiff(CustomTestCase):
|
||||||
|
|
||||||
|
def test_lora_qwen3_30b_a3b_instruct_2507_logprob_accuracy(self):
|
||||||
|
adapter_path = snapshot_download(
|
||||||
|
LORA_HF_REPO,
|
||||||
|
repo_type="dataset",
|
||||||
|
)
|
||||||
|
|
||||||
|
engine = sgl.Engine(
|
||||||
|
model_path=BASE_MODEL,
|
||||||
|
tp_size=TP_SIZE,
|
||||||
|
enable_lora=True,
|
||||||
|
max_lora_rank=MAX_LORA_RANK,
|
||||||
|
lora_paths={"my_lora": adapter_path},
|
||||||
|
lora_backend=LORA_BACKEND,
|
||||||
|
attention_backend="flashinfer",
|
||||||
|
disable_cuda_graph=DISABLE_CUDA_GRAPH,
|
||||||
|
moe_runner_backend=MOE_RUNNER_BACKEND,
|
||||||
|
experts_shared_outer_loras=EXPERTS_SHARED_OUTER_LORAS,
|
||||||
|
prefill_attention_backend=PREFILL_ATTENTION_BACKEND,
|
||||||
|
decode_attention_backend=DECODE_ATTENTION_BACKEND,
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
cdata = torch.load(
|
||||||
|
os.path.join(adapter_path, "compare_sample_train_data.pt"),
|
||||||
|
weights_only=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
base_logprobs = get_prompt_logprobs(engine, cdata["tokens"], lora_path=None)
|
||||||
|
logprobs = get_prompt_logprobs(engine, cdata["tokens"], lora_path="my_lora")
|
||||||
|
|
||||||
|
base_t = torch.tensor(base_logprobs)
|
||||||
|
lora_t = torch.tensor(logprobs)
|
||||||
|
diff = (base_t - lora_t).abs()
|
||||||
|
print(
|
||||||
|
f"[VERIFY] base vs lora: mean_diff={diff.mean().item():.6f}, "
|
||||||
|
f"max_diff={diff.max().item():.6f}, "
|
||||||
|
f"identical={torch.equal(base_t, lora_t)}"
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertFalse(
|
||||||
|
torch.equal(base_t, lora_t),
|
||||||
|
"LoRA logprobs should differ from base model logprobs",
|
||||||
|
)
|
||||||
|
|
||||||
|
kl_sglang_trainer = kl_v2(cdata["training_logprobs"], logprobs)
|
||||||
|
kl_orig_trainer = kl_v2(
|
||||||
|
cdata["training_logprobs"], cdata["sampling_logprobs"]
|
||||||
|
)
|
||||||
|
kl_sglang_orig = kl_v2(logprobs, cdata["sampling_logprobs"])
|
||||||
|
|
||||||
|
print(f"KL(orig_sampler, trainer) = {kl_orig_trainer:.6e}")
|
||||||
|
print(f"KL(sglang, trainer) = {kl_sglang_trainer:.6e}")
|
||||||
|
print(f"KL(sglang, orig_sampler) = {kl_sglang_orig:.6e}")
|
||||||
|
|
||||||
|
self.assertLessEqual(
|
||||||
|
kl_sglang_trainer,
|
||||||
|
KL_THRESHOLD,
|
||||||
|
f"KL(sglang, trainer) = {kl_sglang_trainer:.6e} exceeds "
|
||||||
|
f"threshold {KL_THRESHOLD}",
|
||||||
|
)
|
||||||
|
|
||||||
|
finally:
|
||||||
|
engine.shutdown()
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
try:
|
||||||
|
mp.set_start_method("spawn")
|
||||||
|
except RuntimeError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
try:
|
||||||
|
unittest.main(warnings="ignore", verbosity=2)
|
||||||
|
finally:
|
||||||
|
if torch.cuda.is_available():
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
torch.cuda.synchronize()
|
||||||
Reference in New Issue
Block a user