"""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)