Instructions to use Taykhoom/Evo2-40B-8K with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use Taykhoom/Evo2-40B-8K with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-generation", model="Taykhoom/Evo2-40B-8K", trust_remote_code=True)# pip install -U transformers accelerate # Load model directly from transformers import AutoModelForCausalLM model = AutoModelForCausalLM.from_pretrained("Taykhoom/Evo2-40B-8K", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- vLLM
How to use Taykhoom/Evo2-40B-8K with vLLM:
Install from pip and serve model
# Install vLLM from pip: pip install vllm # Start the vLLM server: vllm serve "Taykhoom/Evo2-40B-8K" # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:8000/v1/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "Taykhoom/Evo2-40B-8K", "prompt": "Once upon a time,", "max_tokens": 512, "temperature": 0.5 }'Use Docker
docker model run hf.co/Taykhoom/Evo2-40B-8K
- SGLang
How to use Taykhoom/Evo2-40B-8K with SGLang:
Install from pip and serve model
# Install SGLang from pip: pip install sglang # Start the SGLang server: python3 -m sglang.launch_server \ --model-path "Taykhoom/Evo2-40B-8K" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "Taykhoom/Evo2-40B-8K", "prompt": "Once upon a time,", "max_tokens": 512, "temperature": 0.5 }'Use Docker images
docker run --gpus all \ --shm-size 32g \ -p 30000:30000 \ -v ~/.cache/huggingface:/root/.cache/huggingface \ --env "HF_TOKEN=<secret>" \ --ipc=host \ lmsysorg/sglang:latest \ python3 -m sglang.launch_server \ --model-path "Taykhoom/Evo2-40B-8K" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "Taykhoom/Evo2-40B-8K", "prompt": "Once upon a time,", "max_tokens": 512, "temperature": 0.5 }' - Docker Model Runner
How to use Taykhoom/Evo2-40B-8K with Docker Model Runner:
docker model run hf.co/Taykhoom/Evo2-40B-8K
Download kernels.py from Taykhoom/Evo2-40B-8K: direct link, hf CLI and curl.
- Browser
- Download file 24.2 kB
-
https://huggingface.co/Taykhoom/Evo2-40B-8K/resolve/main/kernels.py
- Command line
-
hf download hf://Taykhoom/Evo2-40B-8K/kernels.py
-
curl -L -o kernels.py https://huggingface.co/Taykhoom/Evo2-40B-8K/resolve/main/kernels.py
24.2 kB
| """Optional Triton kernels for the Hyena short / medium / long blocks (inference only). | |
| These kernels were written by @AlphaKhaw for vortex in | |
| https://github.com/Zymrael/vortex/pull/77. They are vendored unchanged from | |
| Zymrael/vortex@9c54bb31c7 (vortex/ops/{triton_common,hcs_interface, | |
| hcm_interface,hcl_interface}.py, Apache-2.0), merged into one file. Enabled | |
| per block type with the use_hc{s,m,l}_kernel config flags; off by default. | |
| """ | |
| from collections.abc import Callable | |
| import torch | |
| import triton | |
| import triton.language as tl | |
| # triton_common.py | |
| # 2-D (BLOCK_D, BLOCK_L) tile sweep for memory-bound elementwise kernels over | |
| # a (D, L) plane. Re-benchmarked per shape key by each decorated kernel. | |
| BDL_TILE_CONFIGS: list[triton.Config] = [ | |
| triton.Config({"BLOCK_D": 32, "BLOCK_L": 64}, num_warps=2), | |
| triton.Config({"BLOCK_D": 64, "BLOCK_L": 64}, num_warps=4), | |
| triton.Config({"BLOCK_D": 64, "BLOCK_L": 128}, num_warps=4), | |
| triton.Config({"BLOCK_D": 128, "BLOCK_L": 64}, num_warps=4), | |
| triton.Config({"BLOCK_D": 64, "BLOCK_L": 256}, num_warps=8), | |
| triton.Config({"BLOCK_D": 128, "BLOCK_L": 128}, num_warps=8), | |
| ] | |
| def bdl_grid_3d(B: int, D: int, L: int) -> Callable[[dict], tuple[int, int, int]]: | |
| """ | |
| Standard (B, cdiv(D, BLOCK_D), cdiv(L, BLOCK_L)) grid for the HC bias-residual | |
| and HCS conv kernels. Reads tile sizes from the autotune meta-dict. | |
| """ | |
| return lambda meta: ( | |
| B, | |
| triton.cdiv(D, meta["BLOCK_D"]), | |
| triton.cdiv(L, meta["BLOCK_L"]), | |
| ) | |
| def bdl_grid_2d(D: int, L: int) -> Callable[[dict], tuple[int, int]]: | |
| """ | |
| 2-D (cdiv(D, BLOCK_D), cdiv(L, BLOCK_L)) grid for kernels without a batch | |
| axis (the HCL filter build). | |
| """ | |
| return lambda meta: ( | |
| triton.cdiv(D, meta["BLOCK_D"]), | |
| triton.cdiv(L, meta["BLOCK_L"]), | |
| ) | |
| # hcs_interface.py | |
| def _hcs_depthwise_conv_kernel( | |
| u_ptr, | |
| w_ptr, | |
| z_ptr, | |
| D, | |
| L, | |
| stride_ub, | |
| stride_ud, | |
| stride_ul, | |
| stride_wd, | |
| stride_wk, | |
| FIR_LEN: tl.constexpr, | |
| BLOCK_D: tl.constexpr, | |
| BLOCK_L: tl.constexpr, | |
| ): | |
| """ | |
| Depthwise causal conv: z[b, d, t] = sum_k w[d, k] * u[b, d, t - FIR_LEN + 1 + k]. | |
| One program covers a (BLOCK_D, BLOCK_L) tile of one batch element. The | |
| FIR_LEN tap loop is unrolled at compile time. Input positions before 0 | |
| are masked to zero, giving a causal (left-padded) convolution. | |
| """ | |
| pid_b = tl.program_id(0) | |
| pid_d = tl.program_id(1) | |
| pid_l = tl.program_id(2) | |
| offs_d = pid_d * BLOCK_D + tl.arange(0, BLOCK_D) | |
| offs_l = pid_l * BLOCK_L + tl.arange(0, BLOCK_L) | |
| mask_d = offs_d < D | |
| mask_l = offs_l < L | |
| tile_mask = mask_d[:, None] & mask_l[None, :] | |
| u_base = u_ptr + pid_b * stride_ub + offs_d[:, None] * stride_ud | |
| acc = tl.zeros((BLOCK_D, BLOCK_L), dtype=tl.float32) | |
| # hcs_conv forces fp32 before launch, so loaded tiles are already fp32 -- | |
| # the .to(tl.float32) calls are no-ops at runtime. | |
| for k in tl.static_range(FIR_LEN): | |
| w_k = tl.load( | |
| w_ptr + offs_d * stride_wd + k * stride_wk, mask=mask_d, other=0.0 | |
| ) | |
| pos = offs_l - (FIR_LEN - 1) + k | |
| mask_pos = tile_mask & (pos[None, :] >= 0) & (pos[None, :] < L) | |
| u_tile = tl.load(u_base + pos[None, :] * stride_ul, mask=mask_pos, other=0.0) | |
| acc += w_k[:, None] * u_tile | |
| z_ptrs = ( | |
| z_ptr | |
| + pid_b * stride_ub | |
| + offs_d[:, None] * stride_ud | |
| + offs_l[None, :] * stride_ul | |
| ) | |
| tl.store(z_ptrs, acc, mask=tile_mask) | |
| def hcs_depthwise_conv(u: torch.Tensor, weight: torch.Tensor) -> torch.Tensor: | |
| """ | |
| Depthwise causal 1D convolution, the HCS short-filter time-mixing op. | |
| Equivalent to F.conv1d(u, weight, padding=fir_length - 1, groups=D) | |
| trimmed to length L, but in a single fused Triton launch. | |
| Args: | |
| u (torch.Tensor): Input activations, shape (B, D, L), contiguous. | |
| weight (torch.Tensor): Depthwise filter, shape (D, 1, fir_length), | |
| contiguous. Every channel has its own filter. | |
| Returns: | |
| torch.Tensor: Convolved output, shape (B, D, L), same dtype as u. | |
| """ | |
| u = u.contiguous() | |
| weight = weight.contiguous() | |
| if u.dim() != 3 or weight.dim() != 3: | |
| raise ValueError(f"expected 3-D u and weight, got {u.shape} and {weight.shape}") | |
| B, D, L = u.shape | |
| Dw, in_per_group, fir_length = weight.shape | |
| if Dw != D or in_per_group != 1: | |
| raise ValueError(f"weight {tuple(weight.shape)} is not depthwise for D={D}") | |
| z: torch.Tensor = torch.empty_like(u) | |
| _hcs_depthwise_conv_kernel[bdl_grid_3d(B, D, L)]( | |
| u, | |
| weight, | |
| z, | |
| D, | |
| L, | |
| u.stride(0), | |
| u.stride(1), | |
| u.stride(2), | |
| weight.stride(0), | |
| weight.stride(2), | |
| FIR_LEN=fir_length, | |
| ) | |
| return z | |
| def hcs_conv( | |
| x1: torch.Tensor, | |
| x2: torch.Tensor, | |
| v: torch.Tensor, | |
| weight: torch.Tensor, | |
| bias: torch.Tensor | None = None, | |
| *, | |
| gated_bias: bool = False, | |
| padding_mask: torch.Tensor | None = None, | |
| ) -> torch.Tensor: | |
| """ | |
| Fully-gated HCS short conv: z = x2 * (conv(x1 * v, weight) + bias). | |
| Drop-in for the gated fir_length < 128 branch of | |
| HyenaInferenceEngine.parallel_fir. Conv runs in fp32 for parity with | |
| F.conv1d, then casts back to x1.dtype. | |
| Args: | |
| x1 (torch.Tensor): Pre-gate "key" stream, shape (B, D, L). | |
| x2 (torch.Tensor): Post-gate stream, shape (B, D, L). | |
| v (torch.Tensor): "Value" stream, shape (B, D, L). | |
| weight (torch.Tensor): Depthwise filter, shape (D, 1, fir_length). | |
| bias (torch.Tensor | None): Per-channel skip-gain, shape (D,). | |
| gated_bias (bool): If True, bias is applied multiplicatively | |
| (bias * x1 * v); HCS uses additive (False). | |
| padding_mask (torch.Tensor | None): If set, zeros masked positions | |
| after the conv, shape (B, L). | |
| Returns: | |
| torch.Tensor: Gated HCS output, shape (B, D, L), x1's dtype. | |
| """ | |
| u: torch.Tensor = x1 * v | |
| z: torch.Tensor = hcs_depthwise_conv(u=u.float(), weight=weight.float()) | |
| z = z.to(u.dtype) | |
| if bias is not None: | |
| if gated_bias: | |
| z = z + bias[None, :, None] * u | |
| else: | |
| z = z + bias[None, :, None] | |
| if isinstance(padding_mask, torch.Tensor): | |
| z = z * padding_mask[:, None] | |
| return x2 * z | |
| # hcm_interface.py | |
| # 1-D BLOCK sweep for the complex multiply -- the only HC kernel that flattens | |
| # (D, F) to a single axis, so it can't share BDL_TILE_CONFIGS. | |
| _COMPLEX_MUL_CONFIGS: list[triton.Config] = [ | |
| triton.Config({"BLOCK": 256}, num_warps=2), | |
| triton.Config({"BLOCK": 512}, num_warps=4), | |
| triton.Config({"BLOCK": 1024}, num_warps=4), | |
| triton.Config({"BLOCK": 2048}, num_warps=8), | |
| ] | |
| def _hcm_complex_mul_kernel( | |
| u_ptr, | |
| k_ptr, | |
| y_ptr, | |
| DF, | |
| inv_fft_size, | |
| stride_batch, | |
| BLOCK: tl.constexpr, | |
| ): | |
| """ | |
| Broadcast complex multiply: y[b] = u_f[b] * k_f[0] * inv_fft_size. | |
| One program covers a BLOCK-element slice of the flattened (D, F) plane | |
| of one batch element. Complex values are stored interleaved (real, imag) | |
| -- view_as_real's layout -- so element n's real part is at offset 2n and | |
| its imag part at 2n + 1. The filter k_f carries no batch stride: every | |
| batch element multiplies against the same spectrum. | |
| The flat (D, F) tile is grid axis 0: cdiv(DF, BLOCK) overruns the 65535 | |
| cap on axes 1 and 2 at long context, so the small batch sits on axis 1. | |
| """ | |
| pid = tl.program_id(0) | |
| pid_b = tl.program_id(1) | |
| offs = pid * BLOCK + tl.arange(0, BLOCK) | |
| mask = offs < DF | |
| u_base = u_ptr + pid_b * stride_batch | |
| y_base = y_ptr + pid_b * stride_batch | |
| u_re = tl.load(u_base + 2 * offs, mask=mask, other=0.0) | |
| u_im = tl.load(u_base + 2 * offs + 1, mask=mask, other=0.0) | |
| k_re = tl.load(k_ptr + 2 * offs, mask=mask, other=0.0) | |
| k_im = tl.load(k_ptr + 2 * offs + 1, mask=mask, other=0.0) | |
| y_re = (u_re * k_re - u_im * k_im) * inv_fft_size | |
| y_im = (u_re * k_im + u_im * k_re) * inv_fft_size | |
| tl.store(y_base + 2 * offs, y_re, mask=mask) | |
| tl.store(y_base + 2 * offs + 1, y_im, mask=mask) | |
| def _hcm_complex_mul( | |
| u_f: torch.Tensor, k_f: torch.Tensor, fft_size: int | |
| ) -> torch.Tensor: | |
| """ | |
| Fused broadcast complex multiply with the 1/fft_size filter scale. | |
| Computes u_f * k_f / fft_size -- stage 3 of fftconv_func, with stage 1's | |
| filter normalisation folded in. u_f is the activation spectrum; k_f is | |
| the *unscaled* filter spectrum, already shaped for broadcast over the | |
| batch by adjust_filter_shape_for_broadcast. | |
| Args: | |
| u_f (torch.Tensor): Activation spectrum, complex, shape (B, D, F). | |
| k_f (torch.Tensor): Filter spectrum, complex, shape (1, D, F), shared | |
| across the batch and not yet scaled by 1/fft_size. | |
| fft_size (int): The FFT length n = 2 * seqlen; its reciprocal folds | |
| in as the filter normalisation. | |
| Returns: | |
| torch.Tensor: The scaled product u_f * k_f / fft_size, complex, shape | |
| (B, D, F), u_f's dtype. | |
| Raises: | |
| ValueError: If the tensors are not 3-D complex, or k_f is not | |
| broadcastable over the batch of u_f. | |
| """ | |
| if u_f.dim() != 3 or k_f.dim() != 3: | |
| raise ValueError( | |
| f"expected 3-D u_f and k_f, got {tuple(u_f.shape)} and {tuple(k_f.shape)}" | |
| ) | |
| if not u_f.is_complex() or not k_f.is_complex(): | |
| raise ValueError("u_f and k_f must be complex tensors") | |
| B, D, F = u_f.shape | |
| if tuple(k_f.shape) != (1, D, F): | |
| raise ValueError( | |
| f"k_f {tuple(k_f.shape)} is not broadcastable over u_f {tuple(u_f.shape)}" | |
| ) | |
| u_f = u_f.contiguous() | |
| k_f = k_f.contiguous() | |
| y_f: torch.Tensor = torch.empty_like(u_f) | |
| # Triton has no complex dtype: operate on the (..., 2) real/imag view. | |
| u_r = torch.view_as_real(u_f) | |
| k_r = torch.view_as_real(k_f) | |
| y_r = torch.view_as_real(y_f) | |
| DF = D * F | |
| grid: Callable[[dict], tuple[int, int]] = lambda meta: ( | |
| triton.cdiv(DF, meta["BLOCK"]), | |
| B, | |
| ) | |
| _hcm_complex_mul_kernel[grid]( | |
| u_r, | |
| k_r, | |
| y_r, | |
| DF, | |
| 1.0 / fft_size, | |
| u_r.stride(0), | |
| ) | |
| return y_f | |
| def _hcm_bias_residual_kernel( | |
| y_ptr, | |
| u_ptr, | |
| bias_ptr, | |
| out_ptr, | |
| D, | |
| L, | |
| stride_b, | |
| stride_d, | |
| stride_l, | |
| BLOCK_D: tl.constexpr, | |
| BLOCK_L: tl.constexpr, | |
| ): | |
| """ | |
| Skip-residual add: out[b, d, l] = y[b, d, l] + u[b, d, l] * bias[d]. | |
| One program covers a (BLOCK_D, BLOCK_L) tile of one batch element. y, u | |
| and out share a contiguous (B, D, L) layout; bias is per-channel, shape | |
| (D,), broadcast over batch and length. The fp32 accumulator is cast to | |
| out's dtype on the store -- fftconv_func's stage-6 .to(u.dtype) cast. | |
| """ | |
| pid_b = tl.program_id(0) | |
| pid_d = tl.program_id(1) | |
| pid_l = tl.program_id(2) | |
| offs_d = pid_d * BLOCK_D + tl.arange(0, BLOCK_D) | |
| offs_l = pid_l * BLOCK_L + tl.arange(0, BLOCK_L) | |
| mask_d = offs_d < D | |
| mask_l = offs_l < L | |
| mask = mask_d[:, None] & mask_l[None, :] | |
| offs = pid_b * stride_b + offs_d[:, None] * stride_d + offs_l[None, :] * stride_l | |
| y = tl.load(y_ptr + offs, mask=mask, other=0.0).to(tl.float32) | |
| u = tl.load(u_ptr + offs, mask=mask, other=0.0).to(tl.float32) | |
| bias = tl.load(bias_ptr + offs_d, mask=mask_d, other=0.0).to(tl.float32) | |
| acc = y + u * bias[:, None] | |
| tl.store(out_ptr + offs, acc, mask=mask) | |
| def _hcm_bias_residual( | |
| y: torch.Tensor, u: torch.Tensor, bias: torch.Tensor | |
| ) -> torch.Tensor: | |
| """ | |
| Fused skip-residual add -- stage 5 of fftconv_func. | |
| Computes y + u * bias[:, None] and writes it at u's dtype, fusing | |
| fftconv_func's broadcast multiply and residual add (and its stage-6 | |
| dtype cast) into a single Triton launch. | |
| Args: | |
| y (torch.Tensor): The irfft output, shape (B, D, L). | |
| u (torch.Tensor): The activations, shape (B, D, L); its dtype is the | |
| output dtype -- fftconv_func's stage-6 cast target. | |
| bias (torch.Tensor): Per-channel skip gain, shape (D,), broadcast | |
| over batch and length. | |
| Returns: | |
| torch.Tensor: y + u * bias[:, None], shape (B, D, L), u's dtype. | |
| Raises: | |
| ValueError: If y and u are not matching 3-D tensors, or bias is not | |
| 1-D of length D. | |
| """ | |
| if y.dim() != 3 or u.dim() != 3: | |
| raise ValueError( | |
| f"expected 3-D y and u, got {tuple(y.shape)} and {tuple(u.shape)}" | |
| ) | |
| if y.shape != u.shape: | |
| raise ValueError(f"y {tuple(y.shape)} and u {tuple(u.shape)} must match") | |
| B, D, L = u.shape | |
| if bias.dim() != 1 or bias.shape[0] != D: | |
| raise ValueError(f"bias {tuple(bias.shape)} must be 1-D of length D={D}") | |
| y = y.contiguous() | |
| u = u.contiguous() | |
| bias = bias.contiguous() | |
| out: torch.Tensor = torch.empty_like(u) | |
| _hcm_bias_residual_kernel[bdl_grid_3d(B, D, L)]( | |
| y, | |
| u, | |
| bias, | |
| out, | |
| D, | |
| L, | |
| u.stride(0), | |
| u.stride(1), | |
| u.stride(2), | |
| ) | |
| return out | |
| def hcm_fft_conv( | |
| u: torch.Tensor, | |
| k: torch.Tensor, | |
| D: torch.Tensor, | |
| dropout_mask: torch.Tensor | None, | |
| gelu: bool = False, | |
| k_rev: torch.Tensor | None = None, | |
| bidirectional: bool = False, | |
| print_activations: bool = False, | |
| layer_idx: int | None = None, | |
| **kwargs, | |
| ) -> torch.Tensor: | |
| """ | |
| Fused HCM FFT-convolution -- drop-in for fftconv_func. | |
| cuFFT keeps the three transforms; _hcm_complex_mul does stage 3 (scaled | |
| spectral product), _hcm_bias_residual does stage 5 (skip-residual add). | |
| Trailing kwargs exist only for signature parity with fftconv_func so the | |
| engine dispatch is a one-line swap. gelu and dropout_mask are part of that | |
| parity surface but unsupported -- the kernel has no activation or dropout | |
| stage, so a set value raises instead of being silently dropped. | |
| Args: | |
| u (torch.Tensor): Input activations, shape (B, D, L). | |
| k (torch.Tensor): Filter, shape (D, 1, K). | |
| D (torch.Tensor): Per-channel skip-connection bias, shape (D,). | |
| dropout_mask (torch.Tensor | None): Unsupported; must be None (parity only). | |
| gelu (bool): Unsupported; must be False; the kernel has no activation stage. | |
| Returns: | |
| torch.Tensor: y + u * D[:, None], shape (B, D, L), u's dtype. | |
| Raises: | |
| NotImplementedError: If bidirectional is True, k_rev is set, gelu is | |
| True, or dropout_mask is not None -- the HCM | |
| kernel implements none of these. | |
| """ | |
| if bidirectional or k_rev is not None: | |
| raise NotImplementedError( | |
| "hcm_fft_conv handles only the causal, non-reverse path" | |
| ) | |
| if gelu or dropout_mask is not None: | |
| raise NotImplementedError( | |
| "hcm_fft_conv implements only the gelu=False, dropout_mask=None " | |
| "path used by evo2; the kernel has no activation or dropout stage." | |
| ) | |
| seqlen = u.shape[-1] | |
| fft_size = 2 * seqlen | |
| # rfft(k) reshaped to (1, D, F) for the batch broadcast -- inlined to avoid | |
| # the adjust_filter_shape_for_broadcast import cycle. squeeze(1) drops only | |
| # the channel-group axis; .squeeze() with no arg would also collapse a D=1 | |
| # case and break the (1, D, F) broadcast contract. | |
| k_f = torch.fft.rfft(k, n=fft_size).squeeze(1).unsqueeze(0) | |
| u_f = torch.fft.rfft(u.to(dtype=k.dtype), n=fft_size) | |
| prod = _hcm_complex_mul(u_f, k_f, fft_size) | |
| y = torch.fft.irfft(prod, n=fft_size, norm="forward")[..., :seqlen] | |
| return _hcm_bias_residual(y, u, D) | |
| # hcl_interface.py | |
| def _hcl_compute_filter_kernel( | |
| residues_ptr, | |
| log_poles_ptr, | |
| t_ptr, | |
| h_ptr, | |
| D, | |
| L, | |
| stride_rs_d, | |
| stride_rs_s, | |
| stride_h_d, | |
| stride_h_l, | |
| S: tl.constexpr, | |
| BLOCK_D: tl.constexpr, | |
| BLOCK_L: tl.constexpr, | |
| ): | |
| """ | |
| Modal filter: h[d, l] = sum_s residues[d, s] * exp(log_poles[d, s] * t[l]). | |
| One program covers a (BLOCK_D, BLOCK_L) tile of h; the state-size sum (S | |
| terms) runs in a fp32 register accumulator so the (D, S, L) intermediate | |
| that OOMs the stock compute_filter at L=131k never exists. | |
| """ | |
| pid_d = tl.program_id(0) | |
| pid_l = tl.program_id(1) | |
| offs_d = pid_d * BLOCK_D + tl.arange(0, BLOCK_D) | |
| offs_l = pid_l * BLOCK_L + tl.arange(0, BLOCK_L) | |
| mask_d = offs_d < D | |
| mask_l = offs_l < L | |
| t_tile = tl.load(t_ptr + offs_l, mask=mask_l, other=0.0).to(tl.float32) | |
| acc = tl.zeros((BLOCK_D, BLOCK_L), dtype=tl.float32) | |
| for s in tl.static_range(S): | |
| rs_offs = offs_d * stride_rs_d + s * stride_rs_s | |
| r_s = tl.load(residues_ptr + rs_offs, mask=mask_d, other=0.0).to(tl.float32) | |
| lp_s = tl.load(log_poles_ptr + rs_offs, mask=mask_d, other=0.0).to(tl.float32) | |
| acc += r_s[:, None] * tl.exp(lp_s[:, None] * t_tile[None, :]) | |
| h_ptrs = h_ptr + offs_d[:, None] * stride_h_d + offs_l[None, :] * stride_h_l | |
| tl.store(h_ptrs, acc, mask=mask_d[:, None] & mask_l[None, :]) | |
| def _hcl_compute_filter( | |
| residues: torch.Tensor, log_poles: torch.Tensor, t: torch.Tensor | |
| ) -> torch.Tensor: | |
| """ | |
| Tiled modal-filter build -- the HCL compute_filter without the OOM. | |
| Computes h[d, l] = sum_s residues[d, s] * exp(log_poles[d, s] * t[l]), the | |
| (D, L) filter compute_filter builds, with the state-size sum done | |
| in-register so the (D, state_size, L) intermediate never exists. | |
| Args: | |
| residues (torch.Tensor): Modal residues, shape (D, S). | |
| log_poles (torch.Tensor): Modal log-poles, shape (D, S); negative for | |
| a stable (decaying) filter. | |
| t (torch.Tensor): Time index [0, 1, ..., L-1], shape (L,). | |
| Returns: | |
| torch.Tensor: The modal filter h, shape (D, L), fp32. | |
| Raises: | |
| ValueError: If residues and log_poles are not matching 2-D tensors, | |
| or t is not 1-D. | |
| """ | |
| if residues.dim() != 2 or log_poles.dim() != 2: | |
| raise ValueError( | |
| f"expected 2-D residues and log_poles, got {tuple(residues.shape)} " | |
| f"and {tuple(log_poles.shape)}" | |
| ) | |
| if residues.shape != log_poles.shape: | |
| raise ValueError( | |
| f"residues {tuple(residues.shape)} and log_poles " | |
| f"{tuple(log_poles.shape)} must match" | |
| ) | |
| if t.dim() != 1: | |
| raise ValueError(f"expected 1-D t, got {tuple(t.shape)}") | |
| D, S = residues.shape | |
| L: int = t.shape[0] | |
| residues = residues.contiguous().float() | |
| log_poles = log_poles.contiguous().float() | |
| t = t.contiguous().float() | |
| h: torch.Tensor = torch.empty(D, L, dtype=torch.float32, device=residues.device) | |
| _hcl_compute_filter_kernel[bdl_grid_2d(D, L)]( | |
| residues, | |
| log_poles, | |
| t, | |
| h, | |
| D, | |
| L, | |
| residues.stride(0), | |
| residues.stride(1), | |
| h.stride(0), | |
| h.stride(1), | |
| S, | |
| ) | |
| return h | |
| def _hcl_bias_residual_gate_kernel( | |
| y_ptr, | |
| x1v_ptr, | |
| bias_ptr, | |
| x2_ptr, | |
| out_ptr, | |
| D, | |
| L, | |
| stride_b, | |
| stride_d, | |
| stride_l, | |
| BLOCK_D: tl.constexpr, | |
| BLOCK_L: tl.constexpr, | |
| ): | |
| """ | |
| Bias-residual + gate: out[b,d,l] = (y[b,d,l] + x1v[b,d,l]*bias[d]) * x2[b,d,l]. | |
| One program covers a (BLOCK_D, BLOCK_L) tile of one batch element. y, x1v, | |
| x2 and out share a contiguous (B, D, L) layout; bias is per-channel, shape | |
| (D,), broadcast over batch and length. The fp32 accumulator is cast to | |
| out's dtype on the store -- parallel_iir's y.to(x1v.dtype) cast. | |
| """ | |
| pid_b = tl.program_id(0) | |
| pid_d = tl.program_id(1) | |
| pid_l = tl.program_id(2) | |
| offs_d = pid_d * BLOCK_D + tl.arange(0, BLOCK_D) | |
| offs_l = pid_l * BLOCK_L + tl.arange(0, BLOCK_L) | |
| mask_d = offs_d < D | |
| mask_l = offs_l < L | |
| mask = mask_d[:, None] & mask_l[None, :] | |
| offs = pid_b * stride_b + offs_d[:, None] * stride_d + offs_l[None, :] * stride_l | |
| y = tl.load(y_ptr + offs, mask=mask, other=0.0).to(tl.float32) | |
| x1v = tl.load(x1v_ptr + offs, mask=mask, other=0.0).to(tl.float32) | |
| x2 = tl.load(x2_ptr + offs, mask=mask, other=0.0).to(tl.float32) | |
| bias = tl.load(bias_ptr + offs_d, mask=mask_d, other=0.0).to(tl.float32) | |
| acc = (y + x1v * bias[:, None]) * x2 | |
| tl.store(out_ptr + offs, acc, mask=mask) | |
| def _hcl_bias_residual_gate( | |
| y: torch.Tensor, x1v: torch.Tensor, bias: torch.Tensor, x2: torch.Tensor | |
| ) -> torch.Tensor: | |
| """ | |
| Fused bias-residual + gate -- the HCL FFT-conv epilogue. | |
| Computes (y + x1v * bias[:, None]) * x2 and writes it at x1v's dtype -- | |
| parallel_iir's post-conv `y = (y + x1v * D.unsqueeze(-1)) * x2`, fusing the | |
| broadcast multiply, the residual add, the gate and the dtype cast into one | |
| Triton launch. | |
| Args: | |
| y (torch.Tensor): The irfft output, shape (B, D, L). | |
| x1v (torch.Tensor): The conv input, shape (B, D, L); its dtype is the | |
| output dtype. | |
| bias (torch.Tensor): Per-channel skip gain, shape (D,), broadcast over | |
| batch and length. | |
| x2 (torch.Tensor): The post-gate stream, shape (B, D, L). | |
| Returns: | |
| torch.Tensor: (y + x1v * bias[:, None]) * x2, shape (B, D, L), x1v's | |
| dtype. | |
| Raises: | |
| ValueError: If y, x1v and x2 are not matching 3-D tensors, or bias is | |
| not 1-D of length D. | |
| """ | |
| if y.dim() != 3 or x1v.dim() != 3 or x2.dim() != 3: | |
| raise ValueError( | |
| f"expected 3-D y, x1v, x2, got {tuple(y.shape)}, " | |
| f"{tuple(x1v.shape)}, {tuple(x2.shape)}" | |
| ) | |
| if not (y.shape == x1v.shape == x2.shape): | |
| raise ValueError( | |
| f"y {tuple(y.shape)}, x1v {tuple(x1v.shape)}, x2 {tuple(x2.shape)} " | |
| f"must all match" | |
| ) | |
| B, D, L = x1v.shape | |
| if bias.dim() != 1 or bias.shape[0] != D: | |
| raise ValueError(f"bias {tuple(bias.shape)} must be 1-D of length D={D}") | |
| y = y.contiguous() | |
| x1v = x1v.contiguous() | |
| x2 = x2.contiguous() | |
| bias = bias.contiguous() | |
| out: torch.Tensor = torch.empty_like(x1v) | |
| _hcl_bias_residual_gate_kernel[bdl_grid_3d(B, D, L)]( | |
| y, | |
| x1v, | |
| bias, | |
| x2, | |
| out, | |
| D, | |
| L, | |
| x1v.stride(0), | |
| x1v.stride(1), | |
| x1v.stride(2), | |
| ) | |
| return out | |
| def hcl_fft_conv( | |
| h: torch.Tensor, | |
| x1v: torch.Tensor, | |
| x2: torch.Tensor, | |
| D: torch.Tensor, | |
| L: int, | |
| fft_size: int, | |
| ) -> torch.Tensor: | |
| """ | |
| Fused HCL FFT-convolution epilogue. | |
| Reproduces parallel_iir's long_fir_threshold-is-None branch. cuFFT keeps | |
| the three transforms; _hcm_complex_mul does the spectral product X*H | |
| scaled by 1/fft_size; _hcl_bias_residual_gate does the post-conv | |
| (y + x1v*D[:, None]) * x2. Uses rfft (half the bins of fft on real input). | |
| Args: | |
| h (torch.Tensor): The modal filter, shape (1, D, L). | |
| x1v (torch.Tensor): The conv input, shape (1, D, L). | |
| x2 (torch.Tensor): The post-gate stream, shape (1, D, L). | |
| D (torch.Tensor): Per-channel skip-connection bias, shape (D,). | |
| L (int): Sequence length. | |
| fft_size (int): The FFT length, n = 2 * L. | |
| Returns: | |
| torch.Tensor: (y + x1v*D[:, None]) * x2, shape (1, D, L), x1v's dtype. | |
| """ | |
| H = torch.fft.rfft(h.to(dtype=torch.float32), n=fft_size) | |
| X = torch.fft.rfft(x1v.to(dtype=torch.float32), n=fft_size) | |
| prod = _hcm_complex_mul(X, H, fft_size) | |
| y = torch.fft.irfft(prod, n=fft_size, norm="forward")[..., :L] | |
| return _hcl_bias_residual_gate(y, x1v, D, x2) | |