[Piecewise CUDA Graph] Fix recompile issue for Mixtral and Grok2 (#13667)
Co-authored-by: Minglei Zhu <mingleizhu1122@gmail.com> Co-authored-by: Ke Bao <ISPObaoke@163.com> Co-authored-by: Oasis-Git <ayw.sirius19@gmail.com>
This commit is contained in:
co-authored by
Minglei Zhu
Ke Bao
Oasis-Git
parent
6b262ac839
commit
b5344b31b8
@@ -179,6 +179,9 @@ benchmark/llava_bench/mme_pack
|
|||||||
*.jsonl
|
*.jsonl
|
||||||
tmp*.txt
|
tmp*.txt
|
||||||
|
|
||||||
|
# Torch Compile logs
|
||||||
|
tl_out/
|
||||||
|
|
||||||
# Plots
|
# Plots
|
||||||
*.png
|
*.png
|
||||||
*.pdf
|
*.pdf
|
||||||
|
|||||||
@@ -12,17 +12,11 @@
|
|||||||
# limitations under the License.
|
# limitations under the License.
|
||||||
# ==============================================================================
|
# ==============================================================================
|
||||||
|
|
||||||
# Adapted from
|
|
||||||
# https://github.com/vllm-project/vllm/blob/c7f2cf2b7f67bce5842fedfdba508440fe257375/vllm/model_executor/models/mixtral.py#L1
|
|
||||||
"""Inference-only Grok1 model."""
|
|
||||||
import functools
|
import functools
|
||||||
import logging
|
import logging
|
||||||
import math
|
import math
|
||||||
import os
|
|
||||||
import warnings
|
|
||||||
from typing import Iterable, Optional, Tuple
|
from typing import Iterable, Optional, Tuple
|
||||||
|
|
||||||
import numpy as np
|
|
||||||
import torch
|
import torch
|
||||||
from torch import nn
|
from torch import nn
|
||||||
from transformers import PretrainedConfig
|
from transformers import PretrainedConfig
|
||||||
@@ -30,7 +24,6 @@ from transformers import PretrainedConfig
|
|||||||
from sglang.srt.distributed import (
|
from sglang.srt.distributed import (
|
||||||
get_tensor_model_parallel_rank,
|
get_tensor_model_parallel_rank,
|
||||||
get_tensor_model_parallel_world_size,
|
get_tensor_model_parallel_world_size,
|
||||||
tensor_model_parallel_all_gather,
|
|
||||||
tensor_model_parallel_all_reduce,
|
tensor_model_parallel_all_reduce,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.activation import GeluAndMul
|
from sglang.srt.layers.activation import GeluAndMul
|
||||||
@@ -62,22 +55,15 @@ from sglang.srt.layers.vocab_parallel_embedding import (
|
|||||||
ParallelLMHead,
|
ParallelLMHead,
|
||||||
VocabParallelEmbedding,
|
VocabParallelEmbedding,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.model_executor.cuda_graph_runner import get_is_capture_mode
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
from sglang.srt.model_loader.loader import DefaultModelLoader
|
from sglang.srt.model_loader.loader import DefaultModelLoader
|
||||||
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
||||||
from sglang.srt.server_args import get_global_server_args
|
from sglang.srt.utils import add_prefix
|
||||||
from sglang.srt.utils import add_prefix, dispose_tensor, dump_to_file
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
# Dump tensors for debugging
|
|
||||||
debug_tensor_dump_output_folder = None
|
|
||||||
debug_tensor_dump_inject = False
|
|
||||||
debug_tensor_dump_layers = None
|
|
||||||
debug_tensor_dump_test = False
|
|
||||||
|
|
||||||
|
|
||||||
class Grok1MLP(nn.Module):
|
class Grok1MLP(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -120,14 +106,6 @@ class Grok1MLP(nn.Module):
|
|||||||
|
|
||||||
|
|
||||||
class Grok1MoE(nn.Module):
|
class Grok1MoE(nn.Module):
|
||||||
"""A tensor-parallel MoE implementation for Grok1 that shards each expert
|
|
||||||
across all ranks.
|
|
||||||
|
|
||||||
Each expert's weights are sharded across all ranks and a fused MoE
|
|
||||||
kernel is used for the forward pass, and finally we reduce the outputs
|
|
||||||
across ranks.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
config: PretrainedConfig,
|
config: PretrainedConfig,
|
||||||
@@ -148,7 +126,6 @@ class Grok1MoE(nn.Module):
|
|||||||
super().__init__()
|
super().__init__()
|
||||||
self.hidden_size = hidden_size
|
self.hidden_size = hidden_size
|
||||||
|
|
||||||
# Gate always runs at full precision for stability (see https://arxiv.org/pdf/2101.03961)
|
|
||||||
self.gate = ReplicatedLinear(
|
self.gate = ReplicatedLinear(
|
||||||
hidden_size,
|
hidden_size,
|
||||||
num_experts,
|
num_experts,
|
||||||
@@ -157,9 +134,7 @@ class Grok1MoE(nn.Module):
|
|||||||
quant_config=None,
|
quant_config=None,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.router_logit_softcapping = getattr(
|
self.router_logit_softcapping = 30.0
|
||||||
config, "router_logit_softcapping", 30.0
|
|
||||||
)
|
|
||||||
custom_routing_function = functools.partial(
|
custom_routing_function = functools.partial(
|
||||||
fused_moe_router_shim, self.router_logit_softcapping
|
fused_moe_router_shim, self.router_logit_softcapping
|
||||||
)
|
)
|
||||||
@@ -186,7 +161,6 @@ class Grok1MoE(nn.Module):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||||
# need to assert self.gate.quant_method is unquantized
|
|
||||||
topk_output = self.topk(hidden_states, self.gate.weight)
|
topk_output = self.topk(hidden_states, self.gate.weight)
|
||||||
return self.experts(hidden_states, topk_output)
|
return self.experts(hidden_states, topk_output)
|
||||||
|
|
||||||
@@ -447,75 +421,12 @@ class Grok1Attention(nn.Module):
|
|||||||
hidden_states: torch.Tensor,
|
hidden_states: torch.Tensor,
|
||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
if hidden_states.shape[0] == 0:
|
|
||||||
assert (
|
|
||||||
not self.o_proj.reduce_results
|
|
||||||
), "short-circuiting allreduce will lead to hangs"
|
|
||||||
return hidden_states
|
|
||||||
if debug_tensor_dump_output_folder:
|
|
||||||
dump_to_file(
|
|
||||||
debug_tensor_dump_output_folder,
|
|
||||||
f"attn_input_{self.layer_id}",
|
|
||||||
hidden_states,
|
|
||||||
)
|
|
||||||
|
|
||||||
if debug_tensor_dump_inject:
|
|
||||||
name = os.path.join(
|
|
||||||
debug_tensor_dump_output_folder,
|
|
||||||
f"jax_dump_attn_input_{self.layer_id}.npy",
|
|
||||||
)
|
|
||||||
logger.info(f"Load {name} from jax.")
|
|
||||||
x = np.load(name)
|
|
||||||
hidden_states = torch.tensor(x[0, : hidden_states.shape[0]]).to(
|
|
||||||
hidden_states
|
|
||||||
)
|
|
||||||
|
|
||||||
qkv, _ = self.qkv_proj(hidden_states)
|
qkv, _ = self.qkv_proj(hidden_states)
|
||||||
dispose_tensor(hidden_states)
|
|
||||||
|
|
||||||
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
|
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
|
||||||
q, k = self.rotary_emb(positions, q, k)
|
q, k = self.rotary_emb(positions, q, k)
|
||||||
|
|
||||||
if debug_tensor_dump_output_folder:
|
|
||||||
num_tokens = q.shape[0]
|
|
||||||
num_heads_q = self.num_heads
|
|
||||||
head_dim = self.head_dim
|
|
||||||
num_heads_kv = k.numel() // (num_tokens * head_dim)
|
|
||||||
|
|
||||||
dump_to_file(
|
|
||||||
debug_tensor_dump_output_folder,
|
|
||||||
f"q_{self.layer_id}",
|
|
||||||
tensor_model_parallel_all_gather(
|
|
||||||
q.reshape(num_tokens, num_heads_q, head_dim).contiguous(), dim=1
|
|
||||||
).contiguous(),
|
|
||||||
)
|
|
||||||
dump_to_file(
|
|
||||||
debug_tensor_dump_output_folder,
|
|
||||||
f"k_{self.layer_id}",
|
|
||||||
tensor_model_parallel_all_gather(
|
|
||||||
k.reshape(num_tokens, num_heads_kv, head_dim).contiguous(), dim=1
|
|
||||||
).contiguous(),
|
|
||||||
)
|
|
||||||
dump_to_file(
|
|
||||||
debug_tensor_dump_output_folder,
|
|
||||||
f"v_{self.layer_id}",
|
|
||||||
tensor_model_parallel_all_gather(
|
|
||||||
v.reshape(num_tokens, num_heads_kv, head_dim).contiguous(), dim=1
|
|
||||||
).contiguous(),
|
|
||||||
)
|
|
||||||
|
|
||||||
attn_output = self.attn(q, k, v, forward_batch)
|
attn_output = self.attn(q, k, v, forward_batch)
|
||||||
del q, k, v, qkv
|
|
||||||
|
|
||||||
if debug_tensor_dump_output_folder:
|
|
||||||
dump_to_file(
|
|
||||||
debug_tensor_dump_output_folder,
|
|
||||||
f"attn_output_{self.layer_id}",
|
|
||||||
tensor_model_parallel_all_gather(
|
|
||||||
attn_output.reshape(num_tokens, num_heads_q, head_dim).contiguous(),
|
|
||||||
dim=1,
|
|
||||||
).contiguous(),
|
|
||||||
)
|
|
||||||
|
|
||||||
output, _ = self.o_proj(attn_output)
|
output, _ = self.o_proj(attn_output)
|
||||||
return output
|
return output
|
||||||
@@ -648,14 +559,6 @@ class Grok1DecoderLayer(nn.Module):
|
|||||||
hidden_states,
|
hidden_states,
|
||||||
)
|
)
|
||||||
|
|
||||||
if residual_original is not None:
|
|
||||||
dispose_tensor(residual_original)
|
|
||||||
|
|
||||||
dispose_flag = False
|
|
||||||
if residual is not hidden_states_original:
|
|
||||||
dispose_flag = True
|
|
||||||
dispose_tensor(hidden_states_original)
|
|
||||||
|
|
||||||
hidden_states = self.self_attn(
|
hidden_states = self.self_attn(
|
||||||
positions=positions,
|
positions=positions,
|
||||||
hidden_states=hidden_states,
|
hidden_states=hidden_states,
|
||||||
@@ -673,21 +576,21 @@ class Grok1DecoderLayer(nn.Module):
|
|||||||
self.post_attn_norm.variance_epsilon,
|
self.post_attn_norm.variance_epsilon,
|
||||||
)
|
)
|
||||||
|
|
||||||
if not dispose_flag:
|
|
||||||
dispose_tensor(hidden_states_original)
|
|
||||||
|
|
||||||
# Fully Connected
|
# Fully Connected
|
||||||
hidden_states = self.ffn(hidden_states)
|
hidden_states = self.ffn(hidden_states)
|
||||||
return hidden_states, residual, self.post_moe_norm # defer layernorm
|
return hidden_states, residual, self.post_moe_norm # defer layernorm
|
||||||
|
|
||||||
def moe_with_rmoe(self, x):
|
def moe_with_rmoe(self, x):
|
||||||
current_stream = torch.cuda.current_stream()
|
if self.alt_stream is not None and get_is_capture_mode():
|
||||||
self.alt_stream.wait_stream(current_stream)
|
current_stream = torch.cuda.current_stream()
|
||||||
mlp_result = self.mlp(x)
|
self.alt_stream.wait_stream(current_stream)
|
||||||
with torch.cuda.stream(self.alt_stream):
|
mlp_result = self.mlp(x)
|
||||||
# moe should not be inplace because of stream race condition
|
with torch.cuda.stream(self.alt_stream):
|
||||||
|
moe_result = self.block_sparse_moe(x)
|
||||||
|
current_stream.wait_stream(self.alt_stream)
|
||||||
|
else:
|
||||||
|
mlp_result = self.mlp(x)
|
||||||
moe_result = self.block_sparse_moe(x)
|
moe_result = self.block_sparse_moe(x)
|
||||||
current_stream.wait_stream(self.alt_stream)
|
|
||||||
return (mlp_result + moe_result) / 1.4142135623730951
|
return (mlp_result + moe_result) / 1.4142135623730951
|
||||||
|
|
||||||
|
|
||||||
@@ -752,41 +655,13 @@ class Grok1Model(nn.Module):
|
|||||||
positions, hidden_states, forward_batch, residual, deferred_norm
|
positions, hidden_states, forward_batch, residual, deferred_norm
|
||||||
)
|
)
|
||||||
|
|
||||||
if debug_tensor_dump_output_folder:
|
hidden_states, _ = fused_dual_residual_rmsnorm(
|
||||||
hidden_states = (
|
hidden_states,
|
||||||
fused_rmsnorm(
|
residual,
|
||||||
hidden_states,
|
deferred_norm.weight,
|
||||||
deferred_norm.weight,
|
self.norm.weight,
|
||||||
deferred_norm.variance_epsilon,
|
deferred_norm.variance_epsilon,
|
||||||
)
|
)
|
||||||
+ residual
|
|
||||||
)
|
|
||||||
|
|
||||||
dump_to_file(
|
|
||||||
debug_tensor_dump_output_folder,
|
|
||||||
"last_hidden_before_norm",
|
|
||||||
hidden_states,
|
|
||||||
)
|
|
||||||
|
|
||||||
hidden_states = fused_rmsnorm(
|
|
||||||
hidden_states,
|
|
||||||
self.norm.weight,
|
|
||||||
self.norm.variance_epsilon,
|
|
||||||
)
|
|
||||||
|
|
||||||
dump_to_file(
|
|
||||||
debug_tensor_dump_output_folder,
|
|
||||||
"last_hidden_after_norm",
|
|
||||||
hidden_states,
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
hidden_states, _ = fused_dual_residual_rmsnorm(
|
|
||||||
hidden_states,
|
|
||||||
residual,
|
|
||||||
deferred_norm.weight,
|
|
||||||
self.norm.weight,
|
|
||||||
deferred_norm.variance_epsilon,
|
|
||||||
)
|
|
||||||
|
|
||||||
return hidden_states
|
return hidden_states
|
||||||
|
|
||||||
@@ -862,21 +737,9 @@ class Grok1ForCausalLM(nn.Module):
|
|||||||
)
|
)
|
||||||
self.logits_processor = LogitsProcessor(config)
|
self.logits_processor = LogitsProcessor(config)
|
||||||
|
|
||||||
# Dump tensors for debugging
|
|
||||||
global debug_tensor_dump_output_folder, debug_tensor_dump_inject
|
|
||||||
debug_tensor_dump_output_folder = (
|
|
||||||
get_global_server_args().debug_tensor_dump_output_folder
|
|
||||||
)
|
|
||||||
debug_tensor_dump_inject = get_global_server_args().debug_tensor_dump_inject
|
|
||||||
warnings.filterwarnings("ignore", category=FutureWarning)
|
|
||||||
|
|
||||||
if get_tensor_model_parallel_rank() == 0:
|
|
||||||
logger.info(
|
|
||||||
f"#parameters (analytical): {self.get_num_params_analytical() / 1e9:.2f} B, "
|
|
||||||
f"#parameters (actual): {self.get_num_params_torch() / 1e9:.2f} B"
|
|
||||||
)
|
|
||||||
self.loaded_param_names = set()
|
self.loaded_param_names = set()
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
input_ids: torch.Tensor,
|
input_ids: torch.Tensor,
|
||||||
@@ -884,9 +747,6 @@ class Grok1ForCausalLM(nn.Module):
|
|||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
input_embeds: torch.Tensor = None,
|
input_embeds: torch.Tensor = None,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
if debug_tensor_dump_output_folder:
|
|
||||||
dump_to_file(debug_tensor_dump_output_folder, "input_ids", input_ids)
|
|
||||||
|
|
||||||
hidden_states = self.model(input_ids, positions, forward_batch, input_embeds)
|
hidden_states = self.model(input_ids, positions, forward_batch, input_embeds)
|
||||||
return self.logits_processor(
|
return self.logits_processor(
|
||||||
input_ids, hidden_states, self.lm_head, forward_batch
|
input_ids, hidden_states, self.lm_head, forward_batch
|
||||||
|
|||||||
@@ -353,6 +353,7 @@ class MixtralForCausalLM(nn.Module):
|
|||||||
)
|
)
|
||||||
self.logits_processor = LogitsProcessor(config)
|
self.logits_processor = LogitsProcessor(config)
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
input_ids: torch.Tensor,
|
input_ids: torch.Tensor,
|
||||||
|
|||||||
Reference in New Issue
Block a user