causal-conv1d / torch-rocm /functional.py
Ashiedu's picture
Add files using upload-large-folder tool
45eef7f verified
Raw
History Blame Contribute Delete
7.86 kB
"""Triton-backed implementations of the Mamba-2 causal-conv1d ABI.
Layout (channel-first, Mamba-2 convention):
x : (B, D, L)
weight : (D, WIDTH)
bias : (D,)
y : (B, D, L)
All supported activations (silu, swish, relu, identity) run through Triton.
No PyTorch ``F.conv1d`` fallback remains in the production path.
"""
from __future__ import annotations
from typing import Optional
import torch
import triton
from .triton_kernels import (
causal_conv1d_fwd_kernel,
causal_conv1d_update_kernel,
causal_conv1d_bwd_dx_kernel,
causal_conv1d_bwd_dw_kernel,
causal_conv1d_bwd_db_kernel,
)
BLOCK_L = 256
BLOCK_D = 64
def _resolve_activation(activation: Optional[str]):
if activation in ("silu", "swish"):
return True, False
if activation == "relu":
return False, True
return False, False
class _CausalConv1dFunction(torch.autograd.Function):
@staticmethod
def forward(ctx, x, weight, bias, activation, orig_len):
b, d, l_full = x.shape
width = weight.shape[-1]
apply_silu, apply_relu = _resolve_activation(activation)
y = torch.empty_like(x)
grid = (b, triton.cdiv(l_full, BLOCK_L), triton.cdiv(d, BLOCK_D))
causal_conv1d_fwd_kernel[grid](
x,
weight,
bias,
y,
b,
l_full,
d,
width,
x.stride(0),
x.stride(1),
x.stride(2),
weight.stride(0),
weight.stride(1),
y.stride(0),
y.stride(1),
y.stride(2),
apply_silu,
apply_relu,
BLOCK_L,
BLOCK_D,
)
needs_grad = ctx.needs_input_grad[0] or ctx.needs_input_grad[1] or ctx.needs_input_grad[2]
pre_act = None
if needs_grad and (apply_silu or apply_relu):
pre_act = torch.empty_like(x)
causal_conv1d_fwd_kernel[grid](
x, weight, bias, pre_act,
b, l_full, d, width,
x.stride(0), x.stride(1), x.stride(2),
weight.stride(0), weight.stride(1),
pre_act.stride(0), pre_act.stride(1), pre_act.stride(2),
False, apply_relu, BLOCK_L, BLOCK_D,
)
ctx.save_for_backward(x, weight, bias, pre_act)
ctx.activation = activation
ctx.orig_len = orig_len
return y
@staticmethod
def backward(ctx, dy):
x, weight, bias, pre_act = ctx.saved_tensors
activation = ctx.activation
l_full = x.shape[-1]
if pre_act is not None:
if activation in ("silu", "swish"):
sig = torch.sigmoid(pre_act)
d_pre = dy * sig * (1.0 + pre_act * (1.0 - sig))
else:
d_pre = dy * (pre_act > 0).to(dy.dtype)
else:
d_pre = dy
dx = torch.empty_like(x)
b, d, l = d_pre.shape
width = weight.shape[-1]
grid_dx = (b, triton.cdiv(l, BLOCK_L), triton.cdiv(d, BLOCK_D))
causal_conv1d_bwd_dx_kernel[grid_dx](
d_pre,
weight,
dx,
b,
l,
d,
width,
d_pre.stride(0),
d_pre.stride(1),
d_pre.stride(2),
weight.stride(0),
weight.stride(1),
dx.stride(0),
dx.stride(1),
dx.stride(2),
BLOCK_L,
BLOCK_D,
)
dw = None
if ctx.needs_input_grad[1]:
dw = torch.empty(d, width, dtype=x.dtype, device=x.device)
grid_dw = (triton.cdiv(d, BLOCK_D), width)
causal_conv1d_bwd_dw_kernel[grid_dw](
d_pre,
x,
dw,
b,
l,
d,
width,
d_pre.stride(0),
d_pre.stride(1),
d_pre.stride(2),
x.stride(0),
x.stride(1),
x.stride(2),
dw.stride(0),
dw.stride(1),
BLOCK_D,
)
db = None
if bias is not None and ctx.needs_input_grad[2]:
db = torch.empty(d, dtype=x.dtype, device=x.device)
grid_db = (triton.cdiv(d, BLOCK_D),)
causal_conv1d_bwd_db_kernel[grid_db](
d_pre,
db,
b,
l,
d,
d_pre.stride(0),
d_pre.stride(1),
d_pre.stride(2),
db.stride(0),
BLOCK_D,
)
orig_len = ctx.orig_len
dx = dx[..., -orig_len:] if orig_len is not None else dx
return dx, dw, db, None, None
def causal_conv1d_fn(
x: torch.Tensor,
weight: torch.Tensor,
bias: Optional[torch.Tensor] = None,
seq_idx: Optional[torch.Tensor] = None,
initial_states: Optional[torch.Tensor] = None,
return_final_states: bool = False,
final_states_out: Optional[torch.Tensor] = None,
activation: Optional[str] = "silu",
):
"""Causal 1D depthwise convolution (prefill path) -- Triton forward.
ABI-compatible with Dao-AILab/causal-conv1d ``causal_conv1d_fn``.
"""
if activation not in ("silu", "swish", "relu", "identity", None):
raise NotImplementedError(f"activation {activation!r} not supported by causal_conv1d")
x = x.contiguous()
b, d, l = x.shape
width = weight.shape[-1]
weight = weight.contiguous()
bias_t = bias.contiguous() if bias is not None else torch.zeros(d, dtype=x.dtype, device=x.device)
x_in = x
orig_len = l
if initial_states is not None:
x_in = torch.cat([initial_states.to(x.dtype), x], dim=-1).contiguous()
orig_len = l
y = _CausalConv1dFunction.apply(x_in, weight, bias_t, activation, orig_len)
if return_final_states:
final_states = x[..., -width + 1 :].contiguous()
if final_states_out is not None:
final_states_out.copy_(final_states)
return y, final_states_out
return y, final_states
return y
def causal_conv1d_update(
x: torch.Tensor,
conv_state: torch.Tensor,
weight: torch.Tensor,
bias: torch.Tensor,
activation: Optional[str] = "silu",
cache_seqlens: Optional[torch.Tensor] = None,
conv_state_indices: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""Causal 1D depthwise convolution for a single decode step -- Triton.
ABI-compatible with Dao-AILab/causal-conv1d ``causal_conv1d_update``.
Shapes:
x: (batch, dim) single token
conv_state: (batch, dim, width - 1) cached input window
weight: (dim, width)
bias: (dim,)
returns: (batch, dim)
Side-effect: ``conv_state`` is mutated in place (shift left, append x).
"""
x = x.contiguous()
b, d = x.shape
width = weight.shape[-1]
weight = weight.contiguous()
has_bias = bias is not None
bias_t = bias.contiguous() if has_bias else torch.zeros(d, dtype=x.dtype, device=x.device)
apply_silu, apply_relu = _resolve_activation(activation)
out = torch.empty(b, d, dtype=x.dtype, device=x.device)
grid = (b, triton.cdiv(d, BLOCK_D))
causal_conv1d_update_kernel[grid](
conv_state,
x,
weight,
bias_t,
out,
b,
d,
width,
conv_state.stride(0),
conv_state.stride(1),
conv_state.stride(2),
x.stride(0),
x.stride(1),
weight.stride(0),
weight.stride(1),
bias_t.stride(0) if has_bias else 0,
out.stride(0),
out.stride(1),
has_bias,
apply_silu,
apply_relu,
BLOCK_D,
)
return out