Prhub

#34045 Add registered short-conv tests and backend extensions

原始 PR 作者 aurickq 合并时间 2026-08-08 14:38 文件变更 5 提交数 2 评论 1 代码增减 +234 / -35

执行摘要

新增短卷积后端扩展点与注册表测试,重构默认路径

PR body 指出:线性注意力模型注册表允许模型集成选择后端,但注册模型无法定制后端如何被 full attention 包装;短卷积层直接调用计算内核,难以在复用现有缓存和状态管理生命周期的情况下定制数值实现。因此需要提供可覆盖的扩展点,并补充数值测试来保护这些路径。

值得精读。核心设计是「用可覆盖方法 + 注册表谓词」提供扩展点而不改变默认行为,是插件化集成的良好范式。建议重点关注 sconv.py 中的方法抽取方式(参数收敛为 precomputed 上下文)以及 LinearAttnModelSpec 的兼容性设计(可选字段默认 None)。测试文件可作为数值一致性测试的参考模板。

讨论亮点

本 PR 仅有一条维护性评论(作者触发 /tag-and-rerun-ci),实质 review 讨论为空。审核者 ispobock 直接 APPROVED,无否决或疑问。

实现拆解

  1. 扩展 LinearAttnModelSpec 数据契约:在 python/sglang/srt/configs/linear_attn_model_registry.py 中新增 hybrid_backend_class_name(可选混合后端类名)和 config_predicate(可选配置谓词)字段;get_linear_attn_config 在类型匹配基础上额外校验谓词,谓词默认为 None 时行为不变。
  2. 更新后端构造逻辑:在 attention_registry.pyattn_backend_wrapper 中,当 spec.hybrid_backend_class_name 存在时动态导入并覆盖默认的 ShortConvHybridAttnBackend,同时提前将 cfg 赋值为 runner.model_config 供后续 full_attention_layer_ids 使用。
  3. 增加混合后端转发能力:在 inkling_sconv_backend.py 中为 forward_metadata 属性补充 setter,将赋值操作转发到 full_attn_backend.forward_metadata,与只读 getter 对称,确保外部写入能正确路由。
  4. 重构短卷积 forward 路径:在 python/sglang/srt/models/inkling_common/sconv.py 中将原先内联的 causal_conv1d 调用抽取为 _apply_causal_sconv_kernel,将 fused_causal_conv1d_update_decode 调用抽取为 _apply_decode_sconv_kernel,两者默认可覆盖;_apply_training_sconv_kernel 转而调用新方法,并移除冗余的 weightis_decode 参数(由方法内部从 self 读取)。
  5. 新增数值测试:新建 test/registered/kernels/ops/mamba/test_sconv_cache.py,包含 test_update_sconv_cache_matches_reference(验证不同 query 长度、初始状态模式和填充槽位下的缓存更新)和 test_cached_continuations_match_full_prefill(验证逐 token decode 与缓存前缀扩展均与全量 prefill 数值对齐)。测试注册到 CI 的 base-b-kernel-unit 阶段,使用 1 卡大 runner。
文件 模块 状态 重要度
python/sglang/srt/models/inkling_common/sconv.py 短卷积 modified 7.85
test/registered/kernels/ops/mamba/test_sconv_cache.py 短卷积测试 added 7.41
python/sglang/srt/configs/linear_attn_model_registry.py 模型注册 modified 6.18
python/sglang/srt/layers/attention/attention_registry.py 后端注册 modified 5.3
python/sglang/srt/layers/attention/linear/inkling_sconv_backend.py 混合后端 modified 5.46

关键符号

_apply_causal_sconv_kernel _apply_decode_sconv_kernel _apply_training_sconv_kernel forward get_linear_attn_config forward_metadata test_update_sconv_cache_matches_reference test_cached_continuations_match_full_prefill

关键源码片段

python/sglang/srt/configs/linear_attn_model_registry.py data-contract

注册表数据契约扩展,新增可选字段,是支持定制混合后端和配置谓词的入口。

# python/sglang/srt/configs/linear_attn_model_registry.py
from collections.abc import Callable
from dataclasses import dataclass, field
from typing import Any, Optional
​
​
@dataclass
class LinearAttnModelSpec:
    """Specification for a hybrid (softmax + linear attention) model."""
​
    config_class: type
    backend_class_name: str # fully-qualified class name, lazily imported
    arch_names: list[str] = field(default_factory=list)
    uses_mamba_radix_cache: bool = True
    support_mamba_cache: bool = True
    support_mamba_cache_extra_buffer: bool = False
    unwrap_text_config: bool = False # call get_text_config() before isinstance check
    # 新增:自定义混合后端类的全限定名,None 表示使用默认 ShortConvHybridAttnBackend
    hybrid_backend_class_name: str | None = None
    # 新增:可选的配置谓词,用于依据具体配置细节(如 layer 分布)选择注册条目
    config_predicate: Callable[[Any], bool] | None = None
​
​
def get_linear_attn_config(hf_config: Any) -> Optional[tuple[LinearAttnModelSpec, Any]]:
    for spec in _LINEAR_ATTN_MODEL_REGISTRY:
        config = hf_config.get_text_config() if spec.unwrap_text_config else hf_config
        # 类型匹配且(若无谓词或谓词通过)才返回该注册项
        if isinstance(config, spec.config_class) and (
            spec.config_predicate is None or spec.config_predicate(config)
        ):
            return spec, config
    return None
python/sglang/srt/layers/attention/attention_registry.py core-logic

后端构造逻辑使用 spec.hybrid_backend_class_name 动态替换混合后端类,是注册表扩展生效的关键环节。

# python/sglang/srt/layers/attention/attention_registry.py
# 在 attn_backend_wrapper 的注册模型分支中
spec_result = get_linear_attn_config(runner.model_config.hf_config)
if spec_result is not None:
    spec, _ = spec_result
    cfg = runner.model_config # 提前取出,供下方 full_attention_layer_ids 使用
    BackendClass = import_backend_class(spec.backend_class_name)
    linear_attn_backend = BackendClass(runner)
    # 若注册时指定了自定义混合后端,则动态导入并替换默认类
    if spec.hybrid_backend_class_name is not None:
        hybrid_backend_cls = import_backend_class(spec.hybrid_backend_class_name)
else:
    raise ValueError(
        "Expected hybrid GDN or NemotronH models, but got unknown model. "
        "If this is a custom hybrid model, use register_linear_attn_model() "
        "from sglang.srt.configs.linear_attn_model_registry."
    )
# 后续 full_attn_layers 从 cfg 读取,并用 hybrid_backend_cls 构造最终混合后端
if runner.is_draft_worker:
    full_attn_layers = [0] # FIXME: we assume that MTP/NEXTN always use full-attention
else:
    full_attn_layers = cfg.full_attention_layer_ids
return hybrid_backend_cls(full_attn_backend, linear_attn_backend, full_attn_layers)
python/sglang/srt/layers/attention/linear/inkling_sconv_backend.py core-logic

为 forward_metadata 属性增加 setter,使外部赋值能正确转发到 full-attention 后端,修复混合后端读写不对称问题。

# python/sglang/srt/layers/attention/linear/inkling_sconv_backend.py
# 在 InklingShortConvHybridAttnBackend 中
@property
def forward_metadata(self):
    # 侧车(short-conv)的元数据通过 conv_state_metadata 获取,
    # 因此这里返回 full-attention 后端的 KV 写位置与 SWA 位置翻译信息
    return self.full_attn_backend.forward_metadata
​
​
@forward_metadata.setter
def forward_metadata(self, value):
    # 新增 setter:将外部赋值转发到 full-attention 后端,
    # 确保初始化或更新前向元数据时路径对称,不会静默丢失
    self.full_attn_backend.forward_metadata = value

评论区精华

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

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

风险与影响

  1. 核心 forward 路径重构:sconv.pyforward 方法从内联内核调用改为方法调用,虽默认参数一致,但若子类覆盖不当可能引入行为偏差;新增测试覆盖了主要数值路径,但仅限 CUDA。
  2. 注册表行为变化:config_predicate 可能使之前匹配的模型不再匹配(若谓词返回 False),但默认 None 时行为不变;hybrid_backend_class_name 动态导入可能因类名错误而在运行时失败。
  3. forward_metadata setter 依赖 full_attn_backend.forward_metadata 可写,若某个 full-attention 后端只读属性未实现 setter,赋值会抛 AttributeError。
  4. 测试仅覆盖 CUDA 且依赖 sglang.srt.models.inkling_common.kernels.sconv 内部接口,未来接口调整需同步维护。

对现有用户:默认行为完全不变,所有新增字段和方法均为可选或默认同义,已注册模型无需修改。对外部模型集成者:首次获得定制短卷积数值实现(如替换为低精度或不同内核)和自定义混合后端包装的能力,且能复用现有缓存与推测解码生命周期。对团队维护:引入了两个新扩展点(可覆盖方法和注册表字段),需在文档或代码注释中明确契约。测试增强了短卷积缓存路径的回归防护。

核心 forward 路径重构 测试仅限 CUDA 新增扩展点需文档化 动态导入有运行时失败风险

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论