Skip to content

[Model] Add BerryLM - #39972

Draft
DanilSmorchkov wants to merge 5 commits into
sgl-project:mainfrom
DanilSmorchkov:add-berrylm
Draft

DanilSmorchkov wants to merge 5 commits into
sgl-project:mainfrom
DanilSmorchkov:add-berrylm

Conversation

@DanilSmorchkov

@DanilSmorchkov DanilSmorchkov commented Sep 17, 2026 •

Copy link
Copy Markdown

Motivation

Add BerryLM (BerryLMForCausalLM, model_type: berrylm), a hybrid linear/full-attention MoE decoder with a
thinking mode, released as rwb-ai/BerryLM-OS (~3.2B active / ~18.3B total, 40 layers, 180k text-only vocabulary,
256k context; transformers: huggingface/transformers#48901, vLLM: vllm-project/vllm#57384).

Architecture: three linear-attention layers per full-attention layer, the linear layers being a gated delta rule with
a per-channel forget gate (rank-128 projection of the layer input → one log-decay per key channel per head, the
Kimi Delta Attention recurrence, GVA 16 key / 32 value heads); GQA full attention with a sigmoid output gate and
partial rotary; Gated Block AttnRes (every layer reads a softmax mixture of the residual streams committed at
block boundaries, gated toward the identity by a per-layer scalar; token-local, transparent to the caches); 128-expert
top-8 MoE plus a shared expert.

Modifications

  • python/sglang/srt/models/berrylm.py — self-contained model on the layer-boundary stages (make_stages; the AttnRes
    mixer sits between layers as residual_batch.fold → mix → set_written): attention block, KDA block with a merged
    in_proj_qkvzb projection and the low-rank gate head, MoE + shared expert, the mixer as a fused Triton kernel (torch
    fallback off-GPU). Linear attention goes through the existing HybridLinearAttnBackend / KDAAttnBackend (Kimi conv
    layout, KimiLinearStateShape).
  • python/sglang/srt/configs/berrylm.py + registration in configs/__init__.py, hf_transformers/common.py
    (the registry plumbing this model needs — the matched config in attn_backend_wrapper, the
    support_mamba_cache_extra_buffer flag — is on main since Let predicate-registered linear-attention models carry the mamba radix-cache leaves #41165).
  • KDA backend: GVA support (num_k_heads != num_v_heads, q/k repeated to the value heads in the backend rather than
    in the model), and a per-layer opt-out of the fused intra-chunk prefill path (layer.kda_fused_intra = False →
    chunk_kda(..., fused_intra=False)): the fused-diagonal kernel clamps the exp2(g_i - g_n) factors to ±126, which
    collapses decays beyond 2^-126 inside a 16-token block; BerryLM's per-channel gates reach that range (issue
    [KDA] Fused intra-chunk prefill path (chunk_kda_fwd_intra(fuse_diagonal=True)) collapses for strong per-channel decays because of the ±126 clamp in the exp2 factorization #39971). Default behaviour for every other model is unchanged.
  • python/sglang/srt/function_call/berrylm_detector.py: --reasoning-parser berrylm (<think> blocks) and
    --tool-call-parser berrylm (XML <tool_call> with typed JSON arguments), streaming and non-streaming; both names
    added to the CLI name lists (parser/reasoning_parser_names.py, function_call/parser_names.py).
  • test/registered/unit/function_call/test_berrylm_detector.py: tool-call detector (typed arguments, parallel calls, multiline
    values, streaming) and reasoning detector (think/content, template-opened think, <tool_call> without </think>).
  • test/registered/e2e/models/test_berrylm.py: GSM8K through the server with the reasoning parser, on the kit's chat backend
    (gsm8k_backend = "sgl_eval", thinking, 16k tokens) like the other reasoning-model tests (stage="extra-a",
    runner_config="1-gpu-large"; registered disabled until the release checkpoint is public).
  • Docs: rows in supported-models/generative_models.mdx, advanced_features/separate_reasoning.mdx, advanced_features/tool_parser.mdx.

Accuracy Tests

Serving-path parity (our harness, H200, release checkpoint): a real launch_server scores fixed token ids over /generate
(input logprobs + greedy decode) and transformers teacher-forces the same ids; default CUDA graphs, radix + mamba prefix
cache on, short prompts and 6k–8k-token documents, cache-hit and multi-turn continuations.

6k–8k prompts, mean|Δlogprob| greedy decode = transformers argmax greedy-decode max|Δlogprob| cache pass
TP1 0.079 87.5–100 % 0.32 identical ids
TP2 0.082 87.5–100 % 0.26 identical ids

0.08 is the bf16 floor of this model: vLLM's port sits at the same distance from transformers, the two engines differ
from each other by 0.078, and two transformers builds (5.18.0 and main of 17.09) differ by 0.080 on the same ids. The worst
single prompt position is 2.7 at TP1 and 3.4 at TP2; the same positions deviate in vLLM and between the two transformers
builds (up to 2.8). The model card's serving commands were also replayed from a clean environment on H200 and H100; parser
check over the chat API: 17/17.

Measured on the 17.09 head of this branch (before the port to the layer-boundary stages, earlier checkpoint) and not
repeated on the current head:

  • layer-wise vs transformers on the same inputs: the fused intra-chunk path is 13.5 % off (rel-L2 of the final hidden
    state, top-1 95.3 %), the unfused path 8.9 % / top-1 97.3 % — the same profile as vLLM's port of the kernels (6.0 % /
    98.8 %) and as transformers' fla vs its own torch reference (5.6 % / 97.3 %);
  • long context (32k / 64k-token documents): per-8k-window mean |Δlogprob| 0.06–0.08 from 6k to 64k, radix + mamba cache
    passes identical. On the current head the longest prompt in the parity run is 8k tokens.

test_berrylm.py run as is on the release checkpoint (our cluster, model path and port swapped, offline data): GSM8K
0.95 over 200 problems (3 of 200 hit the 16k-token cap), so the threshold is 0.90; about 10 minutes including the
server start. It stays registered disabled until the checkpoint is public.

Speed Tests and Profiling

vllm bench serve client against sglang.launch_server on H200 (random prompts, ignore_eos, default settings). TP1
measured on this branch's tree with the release checkpoint; TP2 is from the 17.09 head of the branch (before the port to
the layer-boundary stages), not re-measured:

1k→512 c=1 TPOT ms 1k→512 c=16 tok/s 1k→512 c=64 tok/s 1k→2k c=64 tok/s 8k→512 c=64 tok/s
TP1 5.13 1612 3957 4448 1689
TP2 (17.09) 4.0 2321 6035 6888 2891

No change to any other model's path (the fused-intra switch defaults to the upstream behaviour).

Checklist

  • Format your code according to the Format code with pre-commit guide.
  • Add unit tests according to the Run and add unit tests guide.
  • Update documentation according to Write documentations (supported-models/generative_models.mdx, separate_reasoning.mdx, tool_parser.mdx).
  • Provide accuracy and speed benchmark results.
  • Follow the SGLang code style guidance.

CI States

Latest PR Test (Base): ❌ Run #37432786991
Latest PR Test (Extra): ❌ Run #37432786905
Latest PR Test (AMD ROCm 10): ❌ Run #37432787080

@DanilSmorchkov DanilSmorchkov changed the title [Model] Support BerryLM (hybrid KDA MoE with Gated Block AttnRes) [Model] Add BerryLM Oct 6, 2026

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

documentation Improvements or additions to documentation jit-kernel

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant