Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions benchmark/kernels/fused_moe_triton/common_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -78,6 +78,7 @@ def get_model_config(
topk = config.num_experts_per_tok
intermediate_size = config.intermediate_size
elif architecture in [
"BerryLMForCausalLM",
"Qwen2MoeForCausalLM",
"Qwen3MoeForCausalLM",
"Qwen3NextForCausalLM",
Expand Down
6 changes: 6 additions & 0 deletions docs/docs/advanced_features/separate_reasoning.mdx
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,12 @@ SGLang supports parsing reasoning content out from "normal" content for reasonin
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>`deepseek-v3`</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Including [DeepSeek‑V3.2](https://huggingface.co/deepseek-ai/DeepSeek-V3.2-Exp). Supports `thinking` parameter</td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>BerryLM-OS</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>`<think>` … `</think>`</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>`berrylm`</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Supports `enable_thinking`; a `<tool_call>` may follow the reasoning without `</think>`</td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>[Standard Qwen3 models](https://huggingface.co/collections/Qwen/qwen3-67dd247413f0e2e4f653967f)</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>`<think>` … `</think>`</td>
Expand Down
5 changes: 5 additions & 0 deletions docs/docs/advanced_features/tool_parser.mdx
Original file line number Diff line number Diff line change
Expand Up @@ -82,6 +82,11 @@ This guide demonstrates how to use SGLang’s [Function calling](https://platfor
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Qwen series (e.g. `Qwen/Qwen3-Next-80B-A3B-Instruct`, `Qwen/Qwen3-VL-30B-A3B-Thinking`) except Qwen3-Coder</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}></td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>`berrylm`</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>BerryLM-OS (e.g. `rwb-ai/BerryLM-OS`), XML `<tool_call>` format</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}></td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>`qwen3_coder`</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Qwen3-Coder (e.g. `Qwen/Qwen3-Coder-30B-A3B-Instruct`)</td>
Expand Down
5 changes: 5 additions & 0 deletions docs/docs/supported-models/generative_models.mdx
Original file line number Diff line number Diff line change
Expand Up @@ -48,6 +48,11 @@ in the GitHub search bar.
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>`moonshotai/Kimi-K2-Instruct`</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>Moonshot AI's 1 trillion parameter MoE model (32B active) with 128K–256K context; state-of-the-art agentic intelligence with stable long-horizon agency across 200–300 sequential tool calls. Features MLA attention and native INT4 quantization. <a href="../advanced_features/separate_reasoning">See Reasoning Parser docs</a></td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>**BerryLM-OS** (18B-A3B)</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>`rwb-ai/BerryLM-OS`</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>Hybrid linear/full-attention MoE (18.3B total, 3.2B active) with a thinking mode: per-channel KDA forget gates on the linear-attention layers and Gated Block AttnRes depth mixing; 180k vocabulary, 256k context.</td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>**Kimi Linear** (48B-A3B)</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>`moonshotai/Kimi-Linear-48B-A3B-Instruct`</td>
Expand Down
10 changes: 8 additions & 2 deletions python/sglang/kernels/ops/attention/fla/kda.py
Original file line number Diff line number Diff line change
Expand Up @@ -1139,6 +1139,7 @@ def chunk_kda_fwd(
dt_bias: Optional[torch.Tensor] = None,
lower_bound: Optional[float] = None,
output_intermediate_states: bool = False,
fused_intra: Optional[bool] = None,
track_state: Optional[torch.Tensor] = None,
track_chunk_idx: Optional[torch.Tensor] = None,
beta_is_raw: bool = False,
Expand Down Expand Up @@ -1197,6 +1198,9 @@ def chunk_kda_fwd(
_H_pr = q.shape[-2]
_B = q.shape[0]
_small_grid = _B * _NT_pr * _H_pr <= 256
# Layers whose per-channel decays exceed the +-126 (log2) clamp of the fused diagonal
# factorization opt out (RadixLinearAttention.kda_fused_intra = False).
_fused_intra = _small_grid if fused_intra is None else bool(fused_intra)
w, u, _, kg, Aqk, _ = chunk_kda_fwd_intra(
q=q,
k=k,
Expand All @@ -1208,8 +1212,8 @@ def chunk_kda_fwd(
chunk_size=chunk_size,
chunk_indices=chunk_indices,
safe_gate=lower_bound is not None,
fuse_diagonal=_small_grid,
fuse_recompute=_small_grid,
fuse_diagonal=_fused_intra,
fuse_recompute=_fused_intra,
)

h, v_new = chunk_gated_delta_rule_fwd_h(
Expand Down Expand Up @@ -1265,6 +1269,7 @@ def chunk_kda(
dt_bias: Optional[torch.Tensor] = None,
lower_bound: Optional[float] = None,
output_intermediate_states: bool = False,
fused_intra: Optional[bool] = None,
track_state: Optional[torch.Tensor] = None,
track_chunk_idx: Optional[torch.Tensor] = None,
beta_is_raw: bool = False,
Expand Down Expand Up @@ -1292,6 +1297,7 @@ def chunk_kda(
dt_bias=dt_bias,
lower_bound=lower_bound,
output_intermediate_states=output_intermediate_states,
fused_intra=fused_intra,
track_state=track_state,
track_chunk_idx=track_chunk_idx,
beta_is_raw=beta_is_raw,
Expand Down
2 changes: 2 additions & 0 deletions python/sglang/srt/configs/__init__.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
from sglang.srt.configs.afmoe import AfmoeConfig
from sglang.srt.configs.bailing_hybrid import BailingHybridConfig, BailingMoeV3VLConfig
from sglang.srt.configs.bailing_moe_v2 import BailingMM2Config
from sglang.srt.configs.berrylm import BerryLMConfig
from sglang.srt.configs.chatglm import ChatGLMConfig
from sglang.srt.configs.cohere2_moe import Cohere2MoeConfig
from sglang.srt.configs.cosmos3 import (
Expand Down Expand Up @@ -95,6 +96,7 @@
"BailingMM2Config",
"BailingMoeV3VLConfig",
"ExaoneConfig",
"BerryLMConfig",
"ChatGLMConfig",
"Cosmos3Config",
"Cosmos3EdgeConfig",
Expand Down
191 changes: 191 additions & 0 deletions python/sglang/srt/configs/berrylm.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,191 @@
# Copyright 2025-2026 SGLang Team
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
"""BerryLM-OS model configuration (text-only hybrid MoE decoder: gated delta-net
linear attention with a per-channel KDA forget gate, full attention every
``full_attention_interval`` layers, sparse MoE + shared expert, and a Gated
Block AttnRes mixer on the residual stream)."""

from transformers import PretrainedConfig

from sglang.srt.configs.linear_attn_model_registry import (
LinearAttnModelSpec,
register_linear_attn_model,
)
from sglang.srt.configs.mamba_utils import (
KimiLinearCacheParams,
KimiLinearStateShape,
mamba2_state_dtype,
)


class BerryLMConfig(PretrainedConfig):
model_type = "berrylm"
keys_to_ignore_at_inference = ["past_key_values"]

def __init__(
self,
vocab_size=180224,
hidden_size=2048,
num_hidden_layers=40,
num_attention_heads=16,
num_key_value_heads=2,
hidden_act="silu",
max_position_embeddings=262144,
initializer_range=0.02,
rms_norm_eps=1e-6,
use_cache=True,
tie_word_embeddings=False,
rope_parameters=None,
rope_scaling=None,
partial_rotary_factor=0.25,
attention_bias=False,
attention_dropout=0.0,
attn_output_gate=True,
head_dim=256,
linear_conv_kernel_dim=4,
linear_key_head_dim=128,
linear_value_head_dim=128,
linear_num_key_heads=16,
linear_num_value_heads=32,
moe_intermediate_size=512,
shared_expert_intermediate_size=512,
num_experts_per_tok=8,
num_experts=128,
norm_topk_prob=True,
output_router_logits=False,
router_aux_loss_coef=0.001,
layer_types=None,
full_attention_interval=4,
attn_res_block_size=8,
attn_res_gated=True,
attn_res_eps=1e-6,
kda_gate_bottleneck=128,
kda_safe_gate=False,
kda_gate_lower_bound=None,
**kwargs,
):
super().__init__(tie_word_embeddings=tie_word_embeddings, **kwargs)
self.vocab_size = vocab_size
self.max_position_embeddings = max_position_embeddings
self.hidden_size = hidden_size
self.num_hidden_layers = num_hidden_layers
self.num_attention_heads = num_attention_heads
self.num_key_value_heads = num_key_value_heads
self.hidden_act = hidden_act
self.initializer_range = initializer_range
self.rms_norm_eps = rms_norm_eps
self.use_cache = use_cache
self.attention_bias = attention_bias
self.attention_dropout = attention_dropout
self.attn_output_gate = attn_output_gate
self.head_dim = head_dim
self.partial_rotary_factor = partial_rotary_factor
# transformers v5 stores RoPE settings in ``rope_parameters``; keep both names populated.
rope = dict(rope_parameters or rope_scaling or {})
rope.setdefault("rope_type", "default")
rope.setdefault("rope_theta", 10000000.0)
rope.setdefault("partial_rotary_factor", partial_rotary_factor)
self.rope_parameters = rope
self.rope_scaling = rope
self.rope_theta = rope["rope_theta"]

self.full_attention_interval = full_attention_interval
self.layer_types = layer_types
if self.layer_types is None:
self.layer_types = [
(
"linear_attention"
if bool((i + 1) % full_attention_interval)
else "full_attention"
)
for i in range(self.num_hidden_layers)
]
if len(self.layer_types) != self.num_hidden_layers:
raise ValueError("layer_types must have num_hidden_layers entries")

self.linear_conv_kernel_dim = linear_conv_kernel_dim
self.linear_key_head_dim = linear_key_head_dim
self.linear_value_head_dim = linear_value_head_dim
self.linear_num_key_heads = linear_num_key_heads
self.linear_num_value_heads = linear_num_value_heads
self.moe_intermediate_size = moe_intermediate_size
self.shared_expert_intermediate_size = shared_expert_intermediate_size
self.num_experts_per_tok = num_experts_per_tok
self.num_experts = num_experts
self.norm_topk_prob = norm_topk_prob
self.output_router_logits = output_router_logits
self.router_aux_loss_coef = router_aux_loss_coef

self.attn_res_block_size = int(attn_res_block_size)
self.attn_res_gated = bool(attn_res_gated)
self.attn_res_eps = float(attn_res_eps)
self.kda_gate_bottleneck = int(kda_gate_bottleneck)
self.kda_safe_gate = bool(kda_safe_gate)
self.kda_gate_lower_bound = kda_gate_lower_bound

@property
def layers_block_type(self):
# SGLang HybridLayerType values: "attention" / "linear_attention".
return [
"attention" if t == "full_attention" else "linear_attention"
for t in self.layer_types
]

@property
def linear_layer_ids(self):
return [
i for i, t in enumerate(self.layers_block_type) if t == "linear_attention"
]

@property
def full_attention_layer_ids(self):
return [i for i, t in enumerate(self.layers_block_type) if t == "attention"]

@property
def mamba2_cache_params(self) -> KimiLinearCacheParams:
# KDA backend state layout (same as Kimi Linear): one fused q|k|v conv window
# [kernel-1, 2*H*K + HV*V] and a [HV, V, K] recurrent state per linear layer.
from sglang.srt.runtime_context import get_parallel

if self.linear_key_head_dim != self.linear_value_head_dim:
raise ValueError(
"BerryLM KDA cache needs linear_key_head_dim == linear_value_head_dim"
)
shape = KimiLinearStateShape.create(
tp_world_size=get_parallel().attn_tp_size,
num_heads=self.linear_num_value_heads,
head_dim=self.linear_value_head_dim,
num_k_heads=self.linear_num_key_heads,
head_k_dim=self.linear_key_head_dim,
conv_kernel_size=self.linear_conv_kernel_dim,
)
return KimiLinearCacheParams(
shape=shape, layers=self.linear_layer_ids, dtype=mamba2_state_dtype(self)
)


# Hybrid plumbing (mamba pool sizing, linear-attention backend, radix-cache args) is
# driven by the linear-attention model registry: BerryLM shares the KDA backend with
# Kimi Linear (same fused q|k|v conv state layout; grouped value heads are handled in
# KDAAttnBackend.forward_extend).
register_linear_attn_model(
LinearAttnModelSpec(
config_class=BerryLMConfig,
backend_class_name="sglang.srt.layers.attention.linear.kda_backend.KDAAttnBackend",
arch_names=["BerryLMForCausalLM"],
uses_mamba_radix_cache=True,
support_mamba_cache=True,
support_mamba_cache_extra_buffer=True,
)
)
Loading
Loading