add dflash gemma4 support (#27471)
Co-authored-by: kpham-sgl <khoa.pham@radixark.ai> Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
co-authored by
kpham-sgl
Claude Opus 4.8
parent
cd60c4edd0
commit
5ea0d1d093
@@ -1127,6 +1127,14 @@ class Gemma4ForCausalLM(PreTrainedModel):
|
|||||||
def dtype(self) -> torch.dtype:
|
def dtype(self) -> torch.dtype:
|
||||||
return next(self.parameters()).dtype
|
return next(self.parameters()).dtype
|
||||||
|
|
||||||
|
def set_dflash_layers_to_capture(self, layer_ids: list[int]):
|
||||||
|
if layer_ids is None:
|
||||||
|
raise ValueError(
|
||||||
|
"DFLASH requires explicit layer_ids for aux hidden capture."
|
||||||
|
)
|
||||||
|
self.capture_aux_hidden_states = True
|
||||||
|
self.model.layers_to_capture = [val + 1 for val in layer_ids]
|
||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
|
|||||||
@@ -301,6 +301,14 @@ class Gemma4ForConditionalGeneration(PreTrainedModel):
|
|||||||
def get_attention_sliding_window_size(self):
|
def get_attention_sliding_window_size(self):
|
||||||
return getattr(self.config.text_config, "sliding_window", -1) - 1
|
return getattr(self.config.text_config, "sliding_window", -1) - 1
|
||||||
|
|
||||||
|
def set_dflash_layers_to_capture(self, layer_ids: List[int]):
|
||||||
|
if layer_ids is None:
|
||||||
|
raise ValueError(
|
||||||
|
"DFLASH requires explicit layer_ids for aux hidden capture."
|
||||||
|
)
|
||||||
|
self.capture_aux_hidden_states = True
|
||||||
|
self.language_model.layers_to_capture = [val + 1 for val in layer_ids]
|
||||||
|
|
||||||
def prepare_attn_masks(
|
def prepare_attn_masks(
|
||||||
self,
|
self,
|
||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
|
|||||||
@@ -620,17 +620,27 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
if hidden_states.numel() == 0:
|
if hidden_states.numel() == 0:
|
||||||
return torch.empty((0,), dtype=torch.long, device=hidden_states.device)
|
return torch.empty((0,), dtype=torch.long, device=hidden_states.device)
|
||||||
|
|
||||||
tp_group = get_tp_group()
|
|
||||||
tp_size = int(tp_group.world_size)
|
|
||||||
|
|
||||||
if not hasattr(lm_head, "weight") or not hasattr(lm_head, "shard_indices"):
|
|
||||||
raise RuntimeError(
|
|
||||||
"DFLASH greedy sampling requires a vocab-parallel head with `weight` and `shard_indices`."
|
|
||||||
)
|
|
||||||
|
|
||||||
shard = lm_head.shard_indices
|
|
||||||
weight = lm_head.weight # [local_vocab_padded, hidden]
|
weight = lm_head.weight # [local_vocab_padded, hidden]
|
||||||
weight_dtype = weight.dtype
|
weight_dtype = weight.dtype
|
||||||
|
num_tokens = int(hidden_states.shape[0])
|
||||||
|
out_tokens = torch.empty(
|
||||||
|
(num_tokens,), dtype=torch.long, device=hidden_states.device
|
||||||
|
)
|
||||||
|
|
||||||
|
def _cast_hs(x: torch.Tensor) -> torch.Tensor:
|
||||||
|
return x if x.dtype == weight_dtype else x.to(weight_dtype)
|
||||||
|
|
||||||
|
if not hasattr(lm_head, "shard_indices"):
|
||||||
|
for start in range(0, num_tokens, int(chunk_size)):
|
||||||
|
end = min(num_tokens, start + int(chunk_size))
|
||||||
|
hs = _cast_hs(hidden_states[start:end])
|
||||||
|
logits = torch.matmul(hs, weight.T)
|
||||||
|
out_tokens[start:end] = torch.argmax(logits, dim=-1).to(torch.long)
|
||||||
|
return out_tokens
|
||||||
|
|
||||||
|
shard = lm_head.shard_indices
|
||||||
|
tp_group = get_tp_group()
|
||||||
|
tp_size = int(tp_group.world_size)
|
||||||
|
|
||||||
# Valid ranges in the local shard (excluding padding):
|
# Valid ranges in the local shard (excluding padding):
|
||||||
# base vocab: [0, num_org)
|
# base vocab: [0, num_org)
|
||||||
@@ -641,14 +651,6 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
org_vocab_start = int(shard.org_vocab_start_index)
|
org_vocab_start = int(shard.org_vocab_start_index)
|
||||||
added_vocab_start = int(shard.added_vocab_start_index)
|
added_vocab_start = int(shard.added_vocab_start_index)
|
||||||
|
|
||||||
num_tokens = int(hidden_states.shape[0])
|
|
||||||
out_tokens = torch.empty(
|
|
||||||
(num_tokens,), dtype=torch.long, device=hidden_states.device
|
|
||||||
)
|
|
||||||
|
|
||||||
def _cast_hs(x: torch.Tensor) -> torch.Tensor:
|
|
||||||
return x if x.dtype == weight_dtype else x.to(weight_dtype)
|
|
||||||
|
|
||||||
def _ensure_local_reduce_buffers(
|
def _ensure_local_reduce_buffers(
|
||||||
chunk_len: int,
|
chunk_len: int,
|
||||||
value_dtype: torch.dtype,
|
value_dtype: torch.dtype,
|
||||||
|
|||||||
@@ -0,0 +1,177 @@
|
|||||||
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import requests
|
||||||
|
|
||||||
|
from sglang.srt.utils import kill_process_tree
|
||||||
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
from sglang.test.run_eval import run_eval
|
||||||
|
from sglang.test.test_utils import (
|
||||||
|
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
DEFAULT_URL_FOR_TEST,
|
||||||
|
CustomTestCase,
|
||||||
|
is_in_ci,
|
||||||
|
popen_launch_server,
|
||||||
|
write_github_step_summary,
|
||||||
|
)
|
||||||
|
|
||||||
|
register_cuda_ci(est_time=720, stage="extra-a", runner_config="2-gpu-large")
|
||||||
|
|
||||||
|
MODEL_NAME = "31B"
|
||||||
|
TARGET_PATH = "google/gemma-4-31B-it"
|
||||||
|
DRAFT_PATH = "z-lab/gemma-4-31B-it-DFlash"
|
||||||
|
TENSOR_PARALLEL_SIZE = 2
|
||||||
|
|
||||||
|
DRAFT_ATTENTION_BACKEND = "flashinfer"
|
||||||
|
SPECULATIVE_NUM_DRAFT_TOKENS = 16
|
||||||
|
GSM8K_NUM_EXAMPLES = 200
|
||||||
|
GSM8K_NUM_THREADS = 128
|
||||||
|
SERVER_LAUNCH_TIMEOUT = DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH * 3
|
||||||
|
|
||||||
|
# Match the existing Gemma4 31B MTP accuracy floor.
|
||||||
|
GSM8K_SCORE_THRESHOLD = 0.75
|
||||||
|
ACCEPT_LENGTH_THRESHOLD = 5.4
|
||||||
|
|
||||||
|
|
||||||
|
def get_server_info(base_url: str) -> dict:
|
||||||
|
response = requests.get(base_url + "/server_info", timeout=10)
|
||||||
|
response.raise_for_status()
|
||||||
|
return response.json()
|
||||||
|
|
||||||
|
|
||||||
|
def get_avg_spec_accept_length(base_url: str) -> Optional[float]:
|
||||||
|
try:
|
||||||
|
info = get_server_info(base_url)
|
||||||
|
except Exception:
|
||||||
|
return None
|
||||||
|
internal_states = info.get("internal_states") or []
|
||||||
|
if not internal_states:
|
||||||
|
return None
|
||||||
|
value = internal_states[0].get("avg_spec_accept_length")
|
||||||
|
if value is None:
|
||||||
|
return None
|
||||||
|
return float(value)
|
||||||
|
|
||||||
|
|
||||||
|
class TestGemma4DFlash31B(CustomTestCase):
|
||||||
|
base_url = DEFAULT_URL_FOR_TEST
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _common_server_args(cls) -> list[str]:
|
||||||
|
args = [
|
||||||
|
"--trust-remote-code",
|
||||||
|
"--attention-backend",
|
||||||
|
"triton",
|
||||||
|
"--dtype",
|
||||||
|
"bfloat16",
|
||||||
|
"--mem-fraction-static",
|
||||||
|
"0.55",
|
||||||
|
"--max-running-requests",
|
||||||
|
"16",
|
||||||
|
"--context-length",
|
||||||
|
"2048",
|
||||||
|
"--max-total-tokens",
|
||||||
|
"32768",
|
||||||
|
"--skip-server-warmup",
|
||||||
|
]
|
||||||
|
if TENSOR_PARALLEL_SIZE > 1:
|
||||||
|
args += ["--tp-size", str(TENSOR_PARALLEL_SIZE)]
|
||||||
|
return args
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _server_args(cls) -> list[str]:
|
||||||
|
return [
|
||||||
|
"--speculative-algorithm",
|
||||||
|
"DFLASH",
|
||||||
|
"--speculative-draft-model-path",
|
||||||
|
DRAFT_PATH,
|
||||||
|
"--speculative-num-draft-tokens",
|
||||||
|
str(SPECULATIVE_NUM_DRAFT_TOKENS),
|
||||||
|
"--speculative-draft-attention-backend",
|
||||||
|
DRAFT_ATTENTION_BACKEND,
|
||||||
|
] + cls._common_server_args()
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _gsm8k_args(cls) -> SimpleNamespace:
|
||||||
|
return SimpleNamespace(
|
||||||
|
base_url=cls.base_url,
|
||||||
|
model=TARGET_PATH,
|
||||||
|
eval_name="gsm8k",
|
||||||
|
api="completion",
|
||||||
|
max_tokens=512,
|
||||||
|
num_examples=GSM8K_NUM_EXAMPLES,
|
||||||
|
num_threads=GSM8K_NUM_THREADS,
|
||||||
|
num_shots=5,
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _stop_process(process) -> None:
|
||||||
|
try:
|
||||||
|
kill_process_tree(process.pid)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
def test_gsm8k_dflash(self) -> None:
|
||||||
|
process = None
|
||||||
|
try:
|
||||||
|
process = popen_launch_server(
|
||||||
|
TARGET_PATH,
|
||||||
|
self.base_url,
|
||||||
|
timeout=SERVER_LAUNCH_TIMEOUT,
|
||||||
|
other_args=self._server_args(),
|
||||||
|
)
|
||||||
|
requests.get(self.base_url + "/flush_cache", timeout=30)
|
||||||
|
|
||||||
|
server_info = get_server_info(self.base_url)
|
||||||
|
self.assertEqual(
|
||||||
|
server_info.get("speculative_algorithm"),
|
||||||
|
"DFLASH",
|
||||||
|
f"{MODEL_NAME}: server did not start with DFLASH",
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
server_info.get("speculative_draft_attention_backend"),
|
||||||
|
DRAFT_ATTENTION_BACKEND,
|
||||||
|
f"{MODEL_NAME}: unexpected DFLASH draft attention backend",
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
server_info.get("speculative_num_draft_tokens"),
|
||||||
|
SPECULATIVE_NUM_DRAFT_TOKENS,
|
||||||
|
f"{MODEL_NAME}: unexpected DFLASH block size",
|
||||||
|
)
|
||||||
|
self.assertFalse(
|
||||||
|
bool(server_info.get("disable_cuda_graph")),
|
||||||
|
f"{MODEL_NAME}: CUDA graph is disabled",
|
||||||
|
)
|
||||||
|
|
||||||
|
metrics = run_eval(self._gsm8k_args())
|
||||||
|
dflash_score = float(metrics["score"])
|
||||||
|
avg_accept = get_avg_spec_accept_length(self.base_url)
|
||||||
|
finally:
|
||||||
|
if process is not None:
|
||||||
|
self._stop_process(process)
|
||||||
|
|
||||||
|
print(
|
||||||
|
f"[Gemma4 {MODEL_NAME} DFlash] "
|
||||||
|
f"score={dflash_score:.4f} threshold={GSM8K_SCORE_THRESHOLD:.4f} "
|
||||||
|
f"avg_spec_accept_length={avg_accept}"
|
||||||
|
)
|
||||||
|
if is_in_ci():
|
||||||
|
write_github_step_summary(
|
||||||
|
f"### Gemma4 {MODEL_NAME} DFlash\n"
|
||||||
|
f"score={dflash_score:.4f}\n"
|
||||||
|
f"threshold={GSM8K_SCORE_THRESHOLD:.4f}\n"
|
||||||
|
f"avg_spec_accept_length={avg_accept}\n"
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertGreaterEqual(dflash_score, GSM8K_SCORE_THRESHOLD)
|
||||||
|
self.assertIsNotNone(avg_accept)
|
||||||
|
self.assertGreaterEqual(
|
||||||
|
avg_accept,
|
||||||
|
ACCEPT_LENGTH_THRESHOLD,
|
||||||
|
f"{MODEL_NAME}: accept length too low",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user