Decision-1.0-Route-0.6B / decision1_rocm_conv.py
Xunzhuo's picture
Load with 🤗 Transformers (trust_remote_code) (#3)
a5b21df
Raw History Blame
3.02 kB
# Copyright 2026 The vLLM Semantic Router Authors.
# SPDX-License-Identifier: Apache-2.0
"""Eos's ROCm depthwise causal convolution + SiLU, as its native runtime runs it.
The native Decision-1.0-Eos runtime replaced the Qwen3.5 gated-delta layers'
``causal_conv1d_fn`` with this Triton kernel on gfx942 GPUs for full-forward
inference calls whose batch times length is at least 2048: the four-tap
convolution accumulates in FP64, is rounded to BF16 like the reference
convolution's output, and SiLU is computed in FP32. Every other call keeps the
reference function. Only this model's layers are rebound; nothing global changes.
"""
from __future__ import annotations
import torch
try:
import triton
import triton.language as tl
from triton.language.extra import libdevice
except ImportError:
triton = None
MINIMUM_BATCH_TIMES_LENGTH = 2048
if triton is None:
raise ImportError("Eos's native ROCm convolution needs Triton")
@triton.jit
def _causal_silu(
X,
W,
Y,
L: tl.constexpr,
D: tl.constexpr,
N: tl.constexpr,
S0: tl.constexpr,
S1: tl.constexpr,
S2: tl.constexpr,
K: tl.constexpr,
BLOCK: tl.constexpr,
):
i = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK)
valid = i < N
d = i % D
t = (i // D) % L
b = i // (D * L)
acc = tl.full((BLOCK,), 0, tl.float64)
for k in tl.static_range(K):
pos = t - (K - 1) + k
x = tl.load(
X + b * S0 + d * S1 + pos * S2, mask=valid & (pos >= 0), other=0
).to(tl.float64)
w = tl.load(W + d * K + k, mask=valid, other=0).to(tl.float64)
acc = tl.fma(x, w, acc)
v = acc.to(tl.float32).to(X.dtype.element_ty).to(tl.float32)
y = tl.div_rn(v, 1.0 + libdevice.exp(-v))
tl.store(Y + i, y, mask=valid)
def causal_silu(x, weight, bias=None, activation=None, **kwargs):
b, d, length = x.shape
y = torch.empty((b, length, d), device=x.device, dtype=x.dtype)
_causal_silu[(triton.cdiv(b * length * d, 256),)](
x,
weight,
y,
length,
d,
b * length * d,
*x.stride(),
4,
256,
num_warps=4,
enable_fp_fusion=False,
)
return y.transpose(1, 2)
class ConvController:
def __init__(self, reference):
self.reference = reference
def __call__(self, x, weight, bias=None, activation=None, **kwargs):
if (
not torch.is_grad_enabled()
and x.device.type == "cuda"
and x.ndim == 3
and x.dtype == torch.bfloat16
and weight.dtype == torch.bfloat16
and bias is None
and activation in ("silu", "swish")
and tuple(weight.shape) == (x.shape[1], 4)
and weight.is_contiguous()
and x.shape[0] * x.shape[2] >= MINIMUM_BATCH_TIMES_LENGTH
):
return causal_silu(x, weight, bias, activation=activation, **kwargs)
return self.reference(x, weight, bias, activation=activation, **kwargs)