Prhub

#29198 Convert SamplingParams to msgspec Struct

原始 PR 作者 merrymercy 合并时间 2026-06-25 04:42 文件变更 3 提交数 3 评论 4 代码增减 +158 / -58

执行摘要

将 SamplingParams 转换为 msgspec Struct 以支持 msgpack 序列化

从更大的 msgpack 迁移中拆分出 SamplingParams 独立转换,以便单独审查和逐步迁移。该 PR 不改变模型计算逻辑,仅重构序列化结构,为后续全面使用 msgspec/msgpack 做准备。

值得精读。重点关注:

1) msgspec.Struct 的继承设计和 post_init 的使用方式;
2) is_normalized 标志如何回避 msgpack 反序列化中的重复处理;
3) 通过 omit_defaults 和清除别名实现紧凑序列化的技巧。建议在合并后密切关注运行中的行为是否符合预期,特别是 top_k 默认值的变化。

讨论亮点

(无审核评论。PR 提交者已在 PR 描述中详细解释了设计动机:从更大 msgpack 迁移中拆分、保留构造行为、添加 is_normalized 守卫。设计决策包括:使用 post_init 而非自定义 init 以兼容 msgspec 反序列化;通过清除 stop 别名避免冗余序列化;默认字段使用 omit_defaults 以压缩序列化体积。)

实现拆解

  1. 核心类重构:在 python/sglang/srt/sampling/sampling_params.py 中将 SamplingParams 改为继承 msgspec.Struct(kw_only=True, omit_defaults=True),并将所有参数声明为带类型注解的类属性。补充了 stop_strsstop_regex_strsstop_str_max_lenstop_regex_max_lenis_normalized 等内部字段。
  2. 构造逻辑迁移:用 __post_init__ 替换原有的 __init__ 方法,保留 None 默认值转换(如 temperature: None → 1.0)、stop 字段别名复制、贪婪采样归一化。同时添加 is_normalized 判断:如果已归一化(msgpack 解码后)则直接返回,避免重复处理。
  3. 序列化优化:在 normalize() 方法中将 stopstop_regex 字段清空(设为 None),避免在序列化时被 omit_defaults 机制保留(因为它们已复制到 stop_strs/stop_regex_strs 中)。
  4. 专项测试:在 test/registered/unit/sampling/test_sampling_params.py 中新增 TestSamplingParamsMsgspecStruct 测试类,覆盖 copy.copy 后独立性、所有 None 字段默认值替换、默认构造实例 msgpack 编码后解码为空字典(默认字段被省略)、归一化后 msgpack round-trip 保留所有字段。
  5. CI 加速:在 scripts/ci/utils/compute_partitions.py 中将 base-a-test-cpu 的并行分片数从 4 调整为 8。
文件 模块 状态 重要度
python/sglang/srt/sampling/sampling_params.py 调度器 modified 8.4
test/registered/unit/sampling/test_sampling_params.py 抽样参数 modified 7.14
scripts/ci/utils/compute_partitions.py CI 脚本 modified 2.55

关键符号

SamplingParams.__post_init__ SamplingParams.normalize TestSamplingParamsMsgspecStruct.test_copy_remains_mutable_and_independent TestSamplingParamsMsgspecStruct.test_none_values_still_use_constructor_defaults TestSamplingParamsMsgspecStruct.test_msgpack_omits_default_fields TestSamplingParamsMsgspecStruct.test_msgpack_round_trip_preserves_normalized_state

关键源码片段

python/sglang/srt/sampling/sampling_params.py core-logic

核心变更文件,将 SamplingParams 从普通类转为 msgspec Struct,并重写构造逻辑。

# python/sglang/srt/sampling/sampling_params.py ( 重构后核心片段 )import msgspec
from typing import Any, Dict, List, Optional, Set, Union# 定义采样参数类为 msgspec.Struct,以便利用 msgpack 序列化
class SamplingParams(msgspec.Struct, kw_only=True, omit_defaults=True):
    """采样参数,用于控制文本生成时的采样行为。"""
    # --- API 参数(由调用者设置) ---
    max_new_tokens: Optional[int] = 128
    stop: Optional[Union[str, List[str]]] = None # API 别名,在 normalize 中复制到 stop_strs 并清空
    stop_token_ids: Optional[Set[int]] = None
    stop_regex: Optional[Union[str, List[str]]] = None # 同 stop
    temperature: float = 1.0
    top_p: float = 1.0
    top_k: int = TOP_K_ALL # 默认值改为 TOP_K_ALL(1<<30),原为 -1 后经 normalize 转换
    min_p: float = 0.0
    frequency_penalty: float = 0.0
    presence_penalty: float = 0.0
    repetition_penalty: float = 1.0
    min_new_tokens: int = 0
    n: int = 1
    json_schema: Optional[str] = None
    regex: Optional[str] = None
    ebnf: Optional[str] = None
    structural_tag: Optional[str] = None
    ignore_eos: bool = False
    skip_special_tokens: bool = True
    spaces_between_special_tokens: bool = True
    no_stop_trim: bool = False
    custom_params: Optional[Dict[str, Any]] = None
    stream_interval: Optional[int] = None
    logit_bias: Optional[Dict[str, float]] = None
    sampling_seed: Optional[int] = None
​
    # --- 内部字段(由 __post_init__ 或 normalize 填充,不在 API 中出现) ---
    stop_strs: Optional[Union[str, List[str]]] = None # 由 stop 复制
    stop_regex_strs: Optional[Union[str, List[str]]] = None # 由 stop_regex 复制
    stop_str_max_len: int = 0 # normalize 时计算
    stop_regex_max_len: int = 0 # normalize 时计算
    is_normalized: bool = False # 标记是否已调用过 normalize
​
    def __post_init__(self):
        # msgspec 在反序列化后也会调用 __post_init__,需要避免重置已归一化的状态
        if self.is_normalized:
            return
​
        # 复制 stop 和 stop_regex 到内部字段
        self.stop_strs = self.stop
        self.stop_regex_strs = self.stop_regex
​
        # 处理 stop_token_ids:过滤 None 并转 set
        if self.stop_token_ids:
            filtered = {int(t) for t in self.stop_token_ids if t is not None}
            self.stop_token_ids = filtered or None
        else:
            self.stop_token_ids = None
​
        # 以下处理:若用户传入 None,则替换为 msgspec 默认值(类属性中的默认值)
        # 但 msgspec 在 __init__ 时已经为 None 字段指定了默认值,不过通过 __post_init__
        # 我们再次确保,以便兼容旧代码直接传 None 的行为(区别于显式传默认值)
        self.temperature = self.temperature if self.temperature is not None else 1.0
        self.top_p = self.top_p if self.top_p is not None else 1.0
        self.top_k = self.top_k if self.top_k is not None else TOP_K_ALL
        self.min_p = self.min_p if self.min_p is not None else 0.0
        self.frequency_penalty = self.frequency_penalty if self.frequency_penalty is not None else 0.0
        self.presence_penalty = self.presence_penalty if self.presence_penalty is not None else 0.0
        self.repetition_penalty = self.repetition_penalty if self.repetition_penalty is not None else 1.0
        self.min_new_tokens = self.min_new_tokens if self.min_new_tokens is not None else 0
        self.n = self.n if self.n is not None else 1
        self.ignore_eos = self.ignore_eos if self.ignore_eos is not None else False
        self.skip_special_tokens = self.skip_special_tokens if self.skip_special_tokens is not None else True
        self.spaces_between_special_tokens = self.spaces_between_special_tokens if self.spaces_between_special_tokens is not None else True
        self.no_stop_trim = self.no_stop_trim if self.no_stop_trim is not None else False
​
        # 贪婪采样归一化:温度低于阈值则视为 greedy 采样
        if self.temperature < _SAMPLING_EPS:
            self.top_k = 1
            self.temperature = 1.0
​
    def normalize(self, tokenizer=None):
        # 清空 API 别名,避免被序列化(omit_defaults 可自动跳过 None)
        self.stop = None
        self.stop_regex = None
        # 设置 is_normalized 标记
        self.is_normalized = True
        # ... 继续其他归一化逻辑
test/registered/unit/sampling/test_sampling_params.py test-coverage

新增 msgspec 结构化行为的专项测试,确保 copy、None 默认值、msgpack 省略与 round-trip 正确。

# test/registered/unit/sampling/test_sampling_params.py ( 新增测试类 )import copy
import msgspec
from unittest.mock import MagicMock
from sglang.srt.sampling.sampling_params import SamplingParams, TOP_K_ALL
from sglang.test.test_utils import CustomTestCaseclass TestSamplingParamsMsgspecStruct(CustomTestCase):
    # 验证 copy.copy 产生的副本独立于原始对象
    def test_copy_remains_mutable_and_independent(self):
        sp = SamplingParams(max_new_tokens=8, custom_params={"a": 1})
        copied = copy.copy(sp)
        copied.max_new_tokens = 16
        copied.custom_params = {"b": 2}
        self.assertEqual(sp.max_new_tokens, 8)
        self.assertEqual(sp.custom_params, {"a": 1})
        self.assertEqual(copied.max_new_tokens, 16)
        self.assertEqual(copied.custom_params, {"b": 2})
​
    # 验证当所有可选字段传入 None 时,行为与构造默认值一致
    def test_none_values_still_use_constructor_defaults(self):
        sp = SamplingParams(
            temperature=None, top_p=None, top_k=None, min_p=None,
            frequency_penalty=None, presence_penalty=None,
            repetition_penalty=None, min_new_tokens=None, n=None,
            ignore_eos=None, skip_special_tokens=None,
            spaces_between_special_tokens=None, no_stop_trim=None,
        )
        self.assertEqual(sp.temperature, 1.0)
        self.assertEqual(sp.top_p, 1.0)
        self.assertEqual(sp.top_k, TOP_K_ALL)
        self.assertEqual(sp.min_p, 0.0)
        self.assertEqual(sp.frequency_penalty, 0.0)
        self.assertEqual(sp.presence_penalty, 0.0)
        self.assertEqual(sp.repetition_penalty, 1.0)
        self.assertEqual(sp.min_new_tokens, 0)
        self.assertEqual(sp.n, 1)
        self.assertFalse(sp.ignore_eos)
        self.assertTrue(sp.skip_special_tokens)
        self.assertTrue(sp.spaces_between_special_tokens)
        self.assertFalse(sp.no_stop_trim)
​
    # 验证默认构造的实例经过 msgpack 编解码后只有空字段(所有值为默认,omit_defaults=True 省略)
    def test_msgpack_omits_default_fields(self):
        encoded = msgspec.msgpack.encode(SamplingParams())
        self.assertEqual(msgspec.msgpack.decode(encoded), {})
​
    # 验证归一化后的 SamplingParams 经过 msgpack round-trip 保留所有状态
    def test_msgpack_round_trip_preserves_normalized_state(self):
        tokenizer = MagicMock()
        tokenizer.encode.side_effect = lambda s, add_special_tokens=False: {
            "hello": [101, 102],
            "world": [201],
        }[s]
        sp = SamplingParams(
            stop=["hello", "world"], stop_regex=r"[a-z]{3}",
            stop_token_ids=[1, 2], temperature=0.5,
        )
        sp.normalize(tokenizer)
​
        encoder = msgspec.msgpack.Encoder()
        decoder = msgspec.msgpack.Decoder(SamplingParams)
        rebuilt = decoder.decode(encoder.encode(sp))
​
        self.assertIsInstance(rebuilt, SamplingParams)
        self.assertTrue(rebuilt.is_normalized)
        self.assertEqual(rebuilt.stop_strs, ["hello", "world"])
        self.assertEqual(rebuilt.stop_str_max_len, 2)
        self.assertEqual(rebuilt.stop_regex_strs, [r"[a-z]{3}"])
        self.assertEqual(rebuilt.stop_regex_max_len, 3)
        self.assertEqual(rebuilt.stop_token_ids, {1, 2})
        self.assertEqual(rebuilt.temperature, 0.5)

评论区精华

没有提炼出高价值讨论线程

当前评论区没有形成足够清晰的争议点或结论,后续有更多讨论时会体现在这里。

风险与影响

  1. 默认值语义变化:原 top_k 默认构造为 -1(内部通过 normalize 转为 TOP_K_ALL),现直接设为 TOP_K_ALL(1<<30)。如果用户依赖 top_k=-1 作为“未设置”的判别,需要调整。
  2. 类型变化stop_token_idsOptional[List[int]] 变为 Optional[Set[int]],去重但失去顺序性,且反序列化时如果传入列表可能隐式转换(测试中已验证 set 类型)。
  3. init 移除:所有内部通过 SamplingParams(...) 构造的代码均通过 msgspec 的元类和 __post_init__ 处理,如果某些代码依赖自定义字段或构造后逻辑,可能被 is_normalized 阻碍。
  4. 测试覆盖:现有测试未覆盖 top_k=-1 等边界值,也未覆盖旧 __init__ 中的 skip = "max_new_tokens" 等处理。
  5. 新外部依赖:直接 import msgspec,如果环境缺失会在 import 阶段报错。

对用户:默认 API 行为理论上不变,但如果用户直接传递 top_k=-1,现在将被视为有效整数而非“未指定”,可能需要调整为使用 TOP_K_ALL 或移除该参数。对系统:序列化方式从 pickle/JSON 切换到 msgpack,效率提升但需要全链路 msgspec 化(此 PR 仅为第一步)。对团队:此 PR 为后续大规模迁移奠定了基础,设计模式(post_init + is_normalized 保护)可复用于其他 Struct。

默认值语义变化 (top_k) 类型变化 (stop_token_ids) 新外部依赖 (msgspec) 缺少边界值测试

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论