Evo2-40B-8K / kernels.py
Taykhoom's picture
Add opt-in Triton Hyena kernels from vortex PR #77 (@AlphaKhaw)
a873801 verified
Raw History Blame Contribute Delete
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
@triton.autotune(configs=BDL_TILE_CONFIGS, key=["D", "L", "FIR_LEN"])
@triton.jit
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),
]
@triton.autotune(configs=_COMPLEX_MUL_CONFIGS, key=["DF"])
@triton.jit
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
@triton.autotune(configs=BDL_TILE_CONFIGS, key=["D", "L"])
@triton.jit
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
@triton.autotune(configs=BDL_TILE_CONFIGS, key=["D", "L"])
@triton.jit
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
@triton.autotune(configs=BDL_TILE_CONFIGS, key=["D", "L"])
@triton.jit
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)