diff --git a/python/sglang/srt/layers/attention/triton_backend.py b/python/sglang/srt/layers/attention/triton_backend.py index 0f6c80051..cd1492ad9 100644 --- a/python/sglang/srt/layers/attention/triton_backend.py +++ b/python/sglang/srt/layers/attention/triton_backend.py @@ -1019,8 +1019,23 @@ class TritonAttnBackend(AttentionBackend): prefix_kv_indices = self.forward_metadata.kv_indices window_start_pos = None - # Build unified kv_indices using fused Triton kernel + # For SWA layers, mirror SWAKVPool.set_kv_buffer: read from the + # precomputed pool.swa_loc. Translate out_cache_loc to SWA-pool index space + # as a fallback when pool.swa_loc is not pre-populated. extend_kv_indices = forward_batch.out_cache_loc + pool = forward_batch.token_to_kv_pool + if ( + layer.sliding_window_size is not None + and layer.sliding_window_size > -1 + and isinstance(pool, SWAKVPool) + and pool.layers_mapping[layer.layer_id][1] + ): + if pool.swa_loc is not None: + extend_kv_indices = pool.swa_loc + else: + extend_kv_indices = pool.translate_loc_from_full_to_swa( + extend_kv_indices + ) # Handle cases where extend_seq_lens or extend_start_loc might not be set # In speculative decoding, we can infer these from spec_info or compute them diff --git a/test/registered/core/test_gemma4_moe_deterministic.py b/test/registered/core/test_gemma4_moe_deterministic.py new file mode 100644 index 000000000..89675c8e6 --- /dev/null +++ b/test/registered/core/test_gemma4_moe_deterministic.py @@ -0,0 +1,127 @@ +"""Regression test for issue #24394. + +`--enable-deterministic-inference` with `--attention-backend triton` on a +hybrid `SWAKVPool` model (Gemma4 family) used to crash with +`CUDA error: an illegal memory access` inside `_fwd_kernel_unified`: the +unified extend kernel read the new tokens at `out_cache_loc` (full-pool +index space) while `SWAKVPool.set_kv_buffer` had written them at the +SWA-translated indices. With diverse prompts the OOB never materialises; +the repro is same-prompt × high-concurrency, which is what this test fires. + +Adapted from the repro script in the bug report (200 identical completions +at concurrency 128, `--max-running-requests 16`). Pre-fix this loses +~40-50% of requests within ~30-40s; post-fix all 200 succeed. +""" + +import concurrent.futures +import unittest + +import requests + +from sglang.srt.utils import kill_process_tree +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.test_utils import ( + DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + DEFAULT_URL_FOR_TEST, + CustomTestCase, + popen_launch_server, +) + +register_cuda_ci(est_time=420, suite="stage-b-test-2-gpu-large") + + +PROMPT = ( + "Question: Janet's ducks lay 16 eggs per day. She eats three for breakfast " + "every morning and bakes muffins for her friends every day with four. She " + "sells the remainder at the farmers' market daily for $2 per fresh duck " + "egg. How much in dollars does she make every day at the farmers' market?\n" + "Answer:" +) +NUM_REQUESTS = 200 +CONCURRENCY = 128 +MAX_TOKENS = 256 + + +class TestGemma4MoeDeterministic(CustomTestCase): + @classmethod + def setUpClass(cls): + cls.model = "google/gemma-4-26B-A4B-it" + cls.base_url = DEFAULT_URL_FOR_TEST + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=[ + "--tp-size", + "2", + "--attention-backend", + "triton", + "--enable-deterministic-inference", + "--dtype", + "bfloat16", + "--mem-fraction-static", + "0.55", + "--max-running-requests", + "16", + "--context-length", + "2048", + "--max-total-tokens", + "32768", + "--skip-server-warmup", + "--random-seed", + "0", + ], + ) + + @classmethod + def tearDownClass(cls): + if hasattr(cls, "process") and cls.process: + kill_process_tree(cls.process.pid) + + def _fire_one(self): + try: + r = requests.post( + self.base_url + "/v1/completions", + json={ + "model": self.model, + "prompt": PROMPT, + "max_tokens": MAX_TOKENS, + "temperature": 0.0, + "top_k": 1, + }, + timeout=300, + ) + r.raise_for_status() + return True, "" + except Exception as e: + return False, repr(e) + + def test_no_ima_under_concurrent_load(self): + try: + requests.get(self.base_url + "/flush_cache", timeout=30) + except Exception: + pass + + n_ok = n_fail = 0 + first_fail = "" + with concurrent.futures.ThreadPoolExecutor(max_workers=CONCURRENCY) as ex: + futs = [ex.submit(self._fire_one) for _ in range(NUM_REQUESTS)] + for f in concurrent.futures.as_completed(futs): + ok, msg = f.result() + if ok: + n_ok += 1 + else: + if n_fail == 0: + first_fail = msg + n_fail += 1 + + print(f"n_ok={n_ok} n_fail={n_fail} first_fail={first_fail!r}") + self.assertEqual( + n_fail, + 0, + f"{n_fail}/{NUM_REQUESTS} requests failed; first error: {first_fail}", + ) + + +if __name__ == "__main__": + unittest.main()