Download decision1_rocm_conv.py from vllm-sr/Decision-1.0-Route-0.6B: direct link, hf CLI and curl.
- Browser
- Download file 3.02 kB
-
https://huggingface.co/vllm-sr/Decision-1.0-Route-0.6B/resolve/a5b21dffbf09f7926b6f20e9ec352574c3aa5860/decision1_rocm_conv.py
- Command line
-
hf download hf://vllm-sr/Decision-1.0-Route-0.6B@a5b21dffbf09f7926b6f20e9ec352574c3aa5860/decision1_rocm_conv.py
-
curl -L -o decision1_rocm_conv.py https://huggingface.co/vllm-sr/Decision-1.0-Route-0.6B/resolve/a5b21dffbf09f7926b6f20e9ec352574c3aa5860/decision1_rocm_conv.py
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") | |
| 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) | |