config: delete the dead get_server_args() bindings across the repo (#33888)
This commit is contained in:
@@ -115,7 +115,6 @@ class InklingShortConvAttnBackend(ShortConvAttnBackend):
|
|||||||
"""Sized from the CONFIGURED capture shapes, once, never reallocated:
|
"""Sized from the CONFIGURED capture shapes, once, never reallocated:
|
||||||
growing a buffer after a graph captured it moves the address that graph
|
growing a buffer after a graph captured it moves the address that graph
|
||||||
reads, and prefill captures before the decode runner reports its bounds."""
|
reads, and prefill captures before the decode runner reports its bounds."""
|
||||||
server_args = get_server_args()
|
|
||||||
cuda_graph_config = get_exec().graph.cuda_graph_config
|
cuda_graph_config = get_exec().graph.cuda_graph_config
|
||||||
decode_bs: list[int] = []
|
decode_bs: list[int] = []
|
||||||
prefill_tokens: list[int] = []
|
prefill_tokens: list[int] = []
|
||||||
|
|||||||
@@ -1374,7 +1374,7 @@ def _eplb_remap_enabled() -> bool:
|
|||||||
from sglang.srt.runtime_context import get_server_args
|
from sglang.srt.runtime_context import get_server_args
|
||||||
|
|
||||||
try:
|
try:
|
||||||
server_args = get_server_args()
|
get_server_args() # probes that a config is published
|
||||||
except ValueError:
|
except ValueError:
|
||||||
# Global server args are not initialized outside the server runtime
|
# Global server args are not initialized outside the server runtime
|
||||||
# (e.g. in unit tests that call select_experts directly). In that case
|
# (e.g. in unit tests that call select_experts directly). In that case
|
||||||
|
|||||||
@@ -3244,7 +3244,6 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
def maybe_evict_swa(self):
|
def maybe_evict_swa(self):
|
||||||
if self.tree_cache.supports_swa():
|
if self.tree_cache.supports_swa():
|
||||||
sliding_window_size = self.tree_cache.sliding_window_size
|
sliding_window_size = self.tree_cache.sliding_window_size
|
||||||
server_args = get_server_args()
|
|
||||||
|
|
||||||
release_leaf_lock = (
|
release_leaf_lock = (
|
||||||
envs.SGLANG_OPT_SWA_RELEASE_LEAF_LOCK_AFTER_WINDOW.get()
|
envs.SGLANG_OPT_SWA_RELEASE_LEAF_LOCK_AFTER_WINDOW.get()
|
||||||
|
|||||||
@@ -21,7 +21,7 @@ from sglang.srt.environ import envs
|
|||||||
from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool
|
from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool
|
||||||
from sglang.srt.mem_cache.deepseek_v4_compress_state import CompressStatePool
|
from sglang.srt.mem_cache.deepseek_v4_compress_state import CompressStatePool
|
||||||
from sglang.srt.mem_cache.memory_pool import KVCache
|
from sglang.srt.mem_cache.memory_pool import KVCache
|
||||||
from sglang.srt.runtime_context import get_exec, get_server_args, get_spec
|
from sglang.srt.runtime_context import get_exec, get_spec
|
||||||
from sglang.srt.utils import ceil_div, is_hip
|
from sglang.srt.utils import ceil_div, is_hip
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -575,7 +575,6 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
|||||||
self.swa_kv_pool = None
|
self.swa_kv_pool = None
|
||||||
self.c4_kv_pool = None
|
self.c4_kv_pool = None
|
||||||
self.c128_kv_pool = None
|
self.c128_kv_pool = None
|
||||||
server_args = get_server_args()
|
|
||||||
spec_extra = (
|
spec_extra = (
|
||||||
(get_spec().speculative_num_draft_tokens - 1)
|
(get_spec().speculative_num_draft_tokens - 1)
|
||||||
if get_spec().speculative_algorithm is not None
|
if get_spec().speculative_algorithm is not None
|
||||||
@@ -658,7 +657,6 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
|||||||
self.full_to_swa_index_mapping = full_to_swa_index_mapping
|
self.full_to_swa_index_mapping = full_to_swa_index_mapping
|
||||||
|
|
||||||
def get_ring_size(self, compress_ratio: int) -> int:
|
def get_ring_size(self, compress_ratio: int) -> int:
|
||||||
server_args = get_server_args()
|
|
||||||
is_speculative = get_spec().speculative_algorithm is not None
|
is_speculative = get_spec().speculative_algorithm is not None
|
||||||
return get_compress_state_ring_size(compress_ratio, is_speculative)
|
return get_compress_state_ring_size(compress_ratio, is_speculative)
|
||||||
|
|
||||||
|
|||||||
@@ -391,9 +391,6 @@ class FlexKVRadixCache(RadixCache):
|
|||||||
return
|
return
|
||||||
|
|
||||||
# Compute the committed prefix mirroring LMCRadixCache's logic.
|
# Compute the committed prefix mirroring LMCRadixCache's logic.
|
||||||
from sglang.srt.runtime_context import get_server_args
|
|
||||||
|
|
||||||
global_server_args = get_server_args()
|
|
||||||
topk = get_spec().speculative_eagle_topk
|
topk = get_spec().speculative_eagle_topk
|
||||||
enable_kv_committed_len = topk is None or topk == 1
|
enable_kv_committed_len = topk is None or topk == 1
|
||||||
if enable_kv_committed_len:
|
if enable_kv_committed_len:
|
||||||
|
|||||||
@@ -16,7 +16,7 @@ from sglang.srt.mem_cache.base_prefix_cache import (
|
|||||||
MatchResult,
|
MatchResult,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode
|
from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode
|
||||||
from sglang.srt.runtime_context import get_memory, get_server_args, get_spec
|
from sglang.srt.runtime_context import get_memory, get_spec
|
||||||
from sglang.srt.utils import create_device_stream, device_stream_context
|
from sglang.srt.utils import create_device_stream, device_stream_context
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -447,7 +447,6 @@ class LMCRadixCache(RadixCache):
|
|||||||
self.lmcache_connector.end_session(req.rid)
|
self.lmcache_connector.end_session(req.rid)
|
||||||
return
|
return
|
||||||
|
|
||||||
global_server_args = get_server_args()
|
|
||||||
topk = get_spec().speculative_eagle_topk
|
topk = get_spec().speculative_eagle_topk
|
||||||
enable_kv_committed_len = topk is None or topk == 1
|
enable_kv_committed_len = topk is None or topk == 1
|
||||||
if enable_kv_committed_len:
|
if enable_kv_committed_len:
|
||||||
|
|||||||
@@ -586,7 +586,6 @@ class DefaultModelLoader(BaseModelLoader):
|
|||||||
hf_weights_files,
|
hf_weights_files,
|
||||||
)
|
)
|
||||||
elif use_safetensors:
|
elif use_safetensors:
|
||||||
server_args = get_server_args()
|
|
||||||
weight_loader_disable_mmap = get_model().weight_loader_disable_mmap
|
weight_loader_disable_mmap = get_model().weight_loader_disable_mmap
|
||||||
weight_loader_prefetch = get_model().weight_loader_prefetch_checkpoints
|
weight_loader_prefetch = get_model().weight_loader_prefetch_checkpoints
|
||||||
prefetch_num_threads = get_model().weight_loader_prefetch_num_threads
|
prefetch_num_threads = get_model().weight_loader_prefetch_num_threads
|
||||||
|
|||||||
@@ -69,7 +69,7 @@ from sglang.srt.models.deepseek_common.utils import (
|
|||||||
_use_aiter_bpreshuffle_gfx95,
|
_use_aiter_bpreshuffle_gfx95,
|
||||||
_use_aiter_gfx95,
|
_use_aiter_gfx95,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args
|
from sglang.srt.runtime_context import get_exec, get_parallel
|
||||||
from sglang.srt.state_capturer.indexer_topk import (
|
from sglang.srt.state_capturer.indexer_topk import (
|
||||||
maybe_capture_indexer_topk,
|
maybe_capture_indexer_topk,
|
||||||
)
|
)
|
||||||
@@ -1280,7 +1280,6 @@ class DeepseekMLAForwardMixin:
|
|||||||
"""
|
"""
|
||||||
Check if we should skip rope and use fused rope+cache path for TileLang DSA on gfx95.
|
Check if we should skip rope and use fused rope+cache path for TileLang DSA on gfx95.
|
||||||
"""
|
"""
|
||||||
server_args = get_server_args()
|
|
||||||
return (
|
return (
|
||||||
_use_aiter_gfx95
|
_use_aiter_gfx95
|
||||||
and self.current_attention_backend in ("dsa", "nsa")
|
and self.current_attention_backend in ("dsa", "nsa")
|
||||||
|
|||||||
@@ -954,7 +954,6 @@ class DeepseekV2MoE(nn.Module):
|
|||||||
has_shared_output = (
|
has_shared_output = (
|
||||||
hidden_states.shape[0] > 0 and self.num_fused_shared_experts == 0
|
hidden_states.shape[0] > 0 and self.num_fused_shared_experts == 0
|
||||||
)
|
)
|
||||||
server_args = get_server_args()
|
|
||||||
dispatch_info = (
|
dispatch_info = (
|
||||||
ExpertLocationDispatchInfo.init_new(layer_id=self.layer_id)
|
ExpertLocationDispatchInfo.init_new(layer_id=self.layer_id)
|
||||||
if get_exec().moe.enable_eplb and not self.is_nextn
|
if get_exec().moe.enable_eplb and not self.is_nextn
|
||||||
@@ -1055,7 +1054,6 @@ class DeepseekV2MoE(nn.Module):
|
|||||||
self.shared_experts.gate_up_proj
|
self.shared_experts.gate_up_proj
|
||||||
):
|
):
|
||||||
return self.forward_cpu(hidden_states)
|
return self.forward_cpu(hidden_states)
|
||||||
server_args = get_server_args()
|
|
||||||
dispatch_info = (
|
dispatch_info = (
|
||||||
ExpertLocationDispatchInfo.init_new(layer_id=self.layer_id)
|
ExpertLocationDispatchInfo.init_new(layer_id=self.layer_id)
|
||||||
if get_exec().moe.enable_eplb and not self.is_nextn
|
if get_exec().moe.enable_eplb and not self.is_nextn
|
||||||
@@ -2972,7 +2970,6 @@ class DeepseekV2ForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin):
|
|||||||
self, architecture: str = "DeepseekV3ForCausalLM"
|
self, architecture: str = "DeepseekV3ForCausalLM"
|
||||||
):
|
):
|
||||||
self.num_fused_shared_experts = 0
|
self.num_fused_shared_experts = 0
|
||||||
server_args = get_server_args()
|
|
||||||
|
|
||||||
if get_exec().moe.disable_shared_experts_fusion:
|
if get_exec().moe.disable_shared_experts_fusion:
|
||||||
return
|
return
|
||||||
|
|||||||
@@ -43,7 +43,7 @@ from sglang.srt.models.minimax_vl_common import (
|
|||||||
merge_vit_qkv_weights,
|
merge_vit_qkv_weights,
|
||||||
)
|
)
|
||||||
from sglang.srt.models.utils import WeightsMapper
|
from sglang.srt.models.utils import WeightsMapper
|
||||||
from sglang.srt.runtime_context import get_exec, get_mm, get_parallel, get_server_args
|
from sglang.srt.runtime_context import get_exec, get_mm, get_parallel
|
||||||
from sglang.srt.utils import add_prefix, get_device_sm, is_cuda, log_info_on_rank0
|
from sglang.srt.utils import add_prefix, get_device_sm, is_cuda, log_info_on_rank0
|
||||||
from sglang.srt.utils.hf_transformers_utils import get_rope_config
|
from sglang.srt.utils.hf_transformers_utils import get_rope_config
|
||||||
|
|
||||||
@@ -134,7 +134,6 @@ class MiniMaxM3SparseForConditionalGeneration(nn.Module):
|
|||||||
|
|
||||||
def _determine_num_fused_shared_experts(self) -> None:
|
def _determine_num_fused_shared_experts(self) -> None:
|
||||||
text_config = self.config.text_config
|
text_config = self.config.text_config
|
||||||
server_args = get_server_args()
|
|
||||||
if get_exec().moe.disable_shared_experts_fusion:
|
if get_exec().moe.disable_shared_experts_fusion:
|
||||||
return
|
return
|
||||||
|
|
||||||
|
|||||||
@@ -12,7 +12,7 @@ from sglang.srt.constrained.base_grammar_backend import (
|
|||||||
GrammarMask,
|
GrammarMask,
|
||||||
GrammarRow,
|
GrammarRow,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_exec, get_server_args
|
from sglang.srt.runtime_context import get_exec
|
||||||
from sglang.srt.sampling.custom_logit_processor import CustomLogitProcessor
|
from sglang.srt.sampling.custom_logit_processor import CustomLogitProcessor
|
||||||
from sglang.srt.sampling.penaltylib.repetition_penalty import apply_scaling_penalties
|
from sglang.srt.sampling.penaltylib.repetition_penalty import apply_scaling_penalties
|
||||||
from sglang.srt.sampling.sampling_params import TOP_K_ALL
|
from sglang.srt.sampling.sampling_params import TOP_K_ALL
|
||||||
@@ -85,7 +85,6 @@ class SamplingBatchInfo:
|
|||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_schedule_batch(cls, batch: ScheduleBatch, vocab_size: int):
|
def from_schedule_batch(cls, batch: ScheduleBatch, vocab_size: int):
|
||||||
global_server_args = get_server_args()
|
|
||||||
enable_deterministic = get_exec().deterministic.enable_deterministic_inference
|
enable_deterministic = get_exec().deterministic.enable_deterministic_inference
|
||||||
|
|
||||||
reqs = batch.reqs
|
reqs = batch.reqs
|
||||||
|
|||||||
@@ -21,8 +21,6 @@ class IndexerTopkCapturer(BaseTopkCapturer):
|
|||||||
max_running_requests: int,
|
max_running_requests: int,
|
||||||
device: str,
|
device: str,
|
||||||
):
|
):
|
||||||
from sglang.srt.runtime_context import get_server_args
|
|
||||||
|
|
||||||
self.num_indexer_layers = num_indexer_layers
|
self.num_indexer_layers = num_indexer_layers
|
||||||
self.index_topk = index_topk
|
self.index_topk = index_topk
|
||||||
|
|
||||||
@@ -31,7 +29,6 @@ class IndexerTopkCapturer(BaseTopkCapturer):
|
|||||||
|
|
||||||
# DP-attention capture is per-rank-local: each rank writes [:local_batch, ...]
|
# DP-attention capture is per-rank-local: each rank writes [:local_batch, ...]
|
||||||
# to its own device_cache, so the buffer only needs to fit one rank's batch.
|
# to its own device_cache, so the buffer only needs to fit one rank's batch.
|
||||||
server_args = get_server_args()
|
|
||||||
max_batch_size = max(get_schedule().chunked_prefill_size, max_running_requests)
|
max_batch_size = max(get_schedule().chunked_prefill_size, max_running_requests)
|
||||||
|
|
||||||
super().__init__(
|
super().__init__(
|
||||||
|
|||||||
@@ -16,7 +16,6 @@ from sglang.srt.runtime_context import (
|
|||||||
get_exec,
|
get_exec,
|
||||||
get_parallel,
|
get_parallel,
|
||||||
get_schedule,
|
get_schedule,
|
||||||
get_server_args,
|
|
||||||
)
|
)
|
||||||
from sglang.srt.state_capturer.base import BaseTopkCapturer
|
from sglang.srt.state_capturer.base import BaseTopkCapturer
|
||||||
|
|
||||||
@@ -40,7 +39,6 @@ class RoutedExpertsCapturer(BaseTopkCapturer):
|
|||||||
max_running_requests: int,
|
max_running_requests: int,
|
||||||
device: str,
|
device: str,
|
||||||
) -> Optional["RoutedExpertsCapturer"]:
|
) -> Optional["RoutedExpertsCapturer"]:
|
||||||
server_args = get_server_args()
|
|
||||||
if not get_exec().features.enable_return_routed_experts:
|
if not get_exec().features.enable_return_routed_experts:
|
||||||
return None
|
return None
|
||||||
if not get_exec().moe.disable_shared_experts_fusion and hasattr(
|
if not get_exec().moe.disable_shared_experts_fusion and hasattr(
|
||||||
|
|||||||
@@ -503,11 +503,8 @@ class TestFromScheduleBatch(CustomTestCase):
|
|||||||
req.tokenizer.eos_token_id = eos_id
|
req.tokenizer.eos_token_id = eos_id
|
||||||
return req
|
return req
|
||||||
|
|
||||||
@patch("sglang.srt.sampling.sampling_batch_info.get_server_args")
|
def test_basic_construction(self):
|
||||||
def test_basic_construction(self, mock_server_args):
|
|
||||||
"""Test that from_schedule_batch correctly extracts sampling params from requests."""
|
"""Test that from_schedule_batch correctly extracts sampling params from requests."""
|
||||||
mock_server_args.return_value.enable_deterministic_inference = False
|
|
||||||
mock_server_args.return_value.enable_custom_logit_processor = False
|
|
||||||
|
|
||||||
reqs = [self._make_req(temp=0.8, top_p=0.9, top_k=50, min_p=0.1)]
|
reqs = [self._make_req(temp=0.8, top_p=0.9, top_k=50, min_p=0.1)]
|
||||||
batch = MagicMock()
|
batch = MagicMock()
|
||||||
@@ -520,11 +517,8 @@ class TestFromScheduleBatch(CustomTestCase):
|
|||||||
self.assertAlmostEqual(info.top_ps[0].item(), 0.9, places=5)
|
self.assertAlmostEqual(info.top_ps[0].item(), 0.9, places=5)
|
||||||
self.assertEqual(info.top_ks[0].item(), 50)
|
self.assertEqual(info.top_ks[0].item(), 50)
|
||||||
|
|
||||||
@patch("sglang.srt.sampling.sampling_batch_info.get_server_args")
|
def test_greedy_detection(self):
|
||||||
def test_greedy_detection(self, mock_server_args):
|
|
||||||
"""Test that top_k=1 sets is_all_greedy=True."""
|
"""Test that top_k=1 sets is_all_greedy=True."""
|
||||||
mock_server_args.return_value.enable_deterministic_inference = False
|
|
||||||
mock_server_args.return_value.enable_custom_logit_processor = False
|
|
||||||
|
|
||||||
reqs = [self._make_req(top_k=1)]
|
reqs = [self._make_req(top_k=1)]
|
||||||
batch = MagicMock()
|
batch = MagicMock()
|
||||||
@@ -533,11 +527,8 @@ class TestFromScheduleBatch(CustomTestCase):
|
|||||||
info = SamplingBatchInfo.from_schedule_batch(batch, VOCAB_SIZE)
|
info = SamplingBatchInfo.from_schedule_batch(batch, VOCAB_SIZE)
|
||||||
self.assertTrue(info.is_all_greedy)
|
self.assertTrue(info.is_all_greedy)
|
||||||
|
|
||||||
@patch("sglang.srt.sampling.sampling_batch_info.get_server_args")
|
def test_logit_bias_construction(self):
|
||||||
def test_logit_bias_construction(self, mock_server_args):
|
|
||||||
"""Test that logit_bias dict is converted to a tensor with correct values."""
|
"""Test that logit_bias dict is converted to a tensor with correct values."""
|
||||||
mock_server_args.return_value.enable_deterministic_inference = False
|
|
||||||
mock_server_args.return_value.enable_custom_logit_processor = False
|
|
||||||
|
|
||||||
reqs = [self._make_req(logit_bias={"5": 2.0, "10": -1.0})]
|
reqs = [self._make_req(logit_bias={"5": 2.0, "10": -1.0})]
|
||||||
batch = MagicMock()
|
batch = MagicMock()
|
||||||
@@ -549,11 +540,8 @@ class TestFromScheduleBatch(CustomTestCase):
|
|||||||
self.assertAlmostEqual(info.logit_bias[0, 10].item(), -1.0)
|
self.assertAlmostEqual(info.logit_bias[0, 10].item(), -1.0)
|
||||||
self.assertAlmostEqual(info.logit_bias[0, 0].item(), 0.0)
|
self.assertAlmostEqual(info.logit_bias[0, 0].item(), 0.0)
|
||||||
|
|
||||||
@patch("sglang.srt.sampling.sampling_batch_info.get_server_args")
|
def test_deterministic_seed(self):
|
||||||
def test_deterministic_seed(self, mock_server_args):
|
|
||||||
"""Test that explicit seed=123 is kept and missing seed defaults to 42."""
|
"""Test that explicit seed=123 is kept and missing seed defaults to 42."""
|
||||||
mock_server_args.return_value.enable_deterministic_inference = True
|
|
||||||
mock_server_args.return_value.enable_custom_logit_processor = False
|
|
||||||
self._exec_ns.deterministic.enable_deterministic_inference = True
|
self._exec_ns.deterministic.enable_deterministic_inference = True
|
||||||
|
|
||||||
reqs = [self._make_req(seed=123), self._make_req(seed=None)]
|
reqs = [self._make_req(seed=123), self._make_req(seed=None)]
|
||||||
@@ -565,11 +553,8 @@ class TestFromScheduleBatch(CustomTestCase):
|
|||||||
self.assertEqual(info.sampling_seed[0].item(), 123)
|
self.assertEqual(info.sampling_seed[0].item(), 123)
|
||||||
self.assertEqual(info.sampling_seed[1].item(), 42) # default
|
self.assertEqual(info.sampling_seed[1].item(), 42) # default
|
||||||
|
|
||||||
@patch("sglang.srt.sampling.sampling_batch_info.get_server_args")
|
def test_from_schedule_batch_sampling_flags(self):
|
||||||
def test_from_schedule_batch_sampling_flags(self, mock_server_args):
|
|
||||||
"""Test that sampling flags (need_top_p/top_k/min_p) are set correctly."""
|
"""Test that sampling flags (need_top_p/top_k/min_p) are set correctly."""
|
||||||
mock_server_args.return_value.enable_deterministic_inference = False
|
|
||||||
mock_server_args.return_value.enable_custom_logit_processor = False
|
|
||||||
|
|
||||||
reqs = [self._make_req(top_p=0.9, top_k=50, min_p=0.1)]
|
reqs = [self._make_req(top_p=0.9, top_k=50, min_p=0.1)]
|
||||||
batch = MagicMock()
|
batch = MagicMock()
|
||||||
@@ -581,11 +566,8 @@ class TestFromScheduleBatch(CustomTestCase):
|
|||||||
self.assertTrue(info.need_min_p_sampling) # 0.1 > 0
|
self.assertTrue(info.need_min_p_sampling) # 0.1 > 0
|
||||||
self.assertFalse(info.is_all_greedy) # top_k=50 > 1
|
self.assertFalse(info.is_all_greedy) # top_k=50 > 1
|
||||||
|
|
||||||
@patch("sglang.srt.sampling.sampling_batch_info.get_server_args")
|
def test_no_logit_bias_when_all_none(self):
|
||||||
def test_no_logit_bias_when_all_none(self, mock_server_args):
|
|
||||||
"""Test that logit_bias stays None when no request has logit_bias set."""
|
"""Test that logit_bias stays None when no request has logit_bias set."""
|
||||||
mock_server_args.return_value.enable_deterministic_inference = False
|
|
||||||
mock_server_args.return_value.enable_custom_logit_processor = False
|
|
||||||
|
|
||||||
reqs = [self._make_req(), self._make_req()]
|
reqs = [self._make_req(), self._make_req()]
|
||||||
batch = MagicMock()
|
batch = MagicMock()
|
||||||
@@ -594,15 +576,12 @@ class TestFromScheduleBatch(CustomTestCase):
|
|||||||
info = SamplingBatchInfo.from_schedule_batch(batch, VOCAB_SIZE)
|
info = SamplingBatchInfo.from_schedule_batch(batch, VOCAB_SIZE)
|
||||||
self.assertIsNone(info.logit_bias)
|
self.assertIsNone(info.logit_bias)
|
||||||
|
|
||||||
@patch("sglang.srt.sampling.sampling_batch_info.get_server_args")
|
def test_custom_logit_processor_merging(self):
|
||||||
def test_custom_logit_processor_merging(self, mock_server_args):
|
|
||||||
"""Test deserialization and merging of custom logit processors."""
|
"""Test deserialization and merging of custom logit processors."""
|
||||||
from sglang.srt.sampling.custom_logit_processor import (
|
from sglang.srt.sampling.custom_logit_processor import (
|
||||||
DisallowedTokensLogitsProcessor,
|
DisallowedTokensLogitsProcessor,
|
||||||
)
|
)
|
||||||
|
|
||||||
mock_server_args.return_value.enable_deterministic_inference = False
|
|
||||||
mock_server_args.return_value.enable_custom_logit_processor = True
|
|
||||||
self._exec_ns.features.enable_custom_logit_processor = True
|
self._exec_ns.features.enable_custom_logit_processor = True
|
||||||
|
|
||||||
proc_str = DisallowedTokensLogitsProcessor.to_str()
|
proc_str = DisallowedTokensLogitsProcessor.to_str()
|
||||||
|
|||||||
Reference in New Issue
Block a user