causal-conv1d / tests /test_causal_conv1d_pytorch.py
Ashiedu's picture
Initial kernel bundle upload: causal-conv1d
8996239 verified
Raw
History Blame Contribute Delete
9.47 kB
import torch
import causal_conv1d
DEVICE = "cuda:0"
def _ref_causal_conv1d(x, weight, bias=None, activation="silu"):
"""Manual reference: causal depthwise conv1d with left pad width-1."""
b, d, seq_len = x.shape
width = weight.shape[-1]
w = weight.unsqueeze(1).contiguous()
xp = torch.nn.functional.pad(x, (width - 1, 0))
out = torch.nn.functional.conv1d(xp, w, bias, groups=d)
if activation in ("silu", "swish"):
out = torch.nn.functional.silu(out)
elif activation == "relu":
out = torch.nn.functional.relu(out)
return out
def test_causal_conv1d_fn_matches_reference_fp32():
"""Triton forward matches F.conv1d reference for all activations (FP32)."""
torch.manual_seed(0)
b, d, seq_len, width = 2, 8, 16, 4
x = torch.randn(b, d, seq_len, dtype=torch.float32, device=DEVICE)
weight = torch.randn(d, width, dtype=torch.float32, device=DEVICE)
bias = torch.randn(d, dtype=torch.float32, device=DEVICE)
for act in ("silu", "relu", "identity"):
out = causal_conv1d.causal_conv1d_fn(x, weight, bias, activation=act)
ref = _ref_causal_conv1d(x, weight, bias, activation=act)
assert out.shape == (b, d, seq_len)
assert out.is_cuda, f"output must be on GPU for activation={act}"
torch.testing.assert_close(out, ref, atol=1e-4, rtol=1e-4)
def test_causal_conv1d_fn_matches_reference_fp16():
"""Triton forward matches F.conv1d reference in FP16 (production dtype)."""
torch.manual_seed(0)
b, d, seq_len, width = 2, 16, 32, 4
x = torch.randn(b, d, seq_len, dtype=torch.float16, device=DEVICE)
weight = torch.randn(d, width, dtype=torch.float16, device=DEVICE)
bias = torch.randn(d, dtype=torch.float16, device=DEVICE)
out = causal_conv1d.causal_conv1d_fn(x, weight, bias, activation="silu")
ref = _ref_causal_conv1d(x, weight, bias, activation="silu")
assert out.shape == (b, d, seq_len)
assert out.dtype == torch.float16
# FP16 tolerance: wider than FP32 due to packed-math accumulation.
torch.testing.assert_close(out, ref, atol=1e-2, rtol=1e-2)
def test_causal_conv1d_fn_initial_states():
"""Triton forward handles prepended initial_states correctly."""
torch.manual_seed(1)
b, d, seq_len, width = 2, 8, 16, 4
x = torch.randn(b, d, seq_len, device=DEVICE)
weight = torch.randn(d, width, device=DEVICE)
bias = torch.randn(d, device=DEVICE)
init = torch.randn(b, d, width - 1, device=DEVICE)
out, final = causal_conv1d.causal_conv1d_fn(
x, weight, bias, initial_states=init, return_final_states=True
)
ref = _ref_causal_conv1d(torch.cat([init, x], dim=-1), weight, bias)
torch.testing.assert_close(out, ref, atol=1e-4, rtol=1e-4)
torch.testing.assert_close(final, x[..., -width + 1 :], atol=1e-5, rtol=0)
def test_causal_conv1d_update_matches_prefill():
"""Decode-step update must equal the prefill output for the last token."""
torch.manual_seed(2)
b, d, seq_len, width = 1, 8, 6, 4
x = torch.randn(b, d, seq_len, device=DEVICE)
weight = torch.randn(d, width, device=DEVICE)
bias = torch.randn(d, device=DEVICE)
# Prefill output via fn over the whole sequence.
prefill = causal_conv1d.causal_conv1d_fn(x, weight, bias, activation="silu")
last_token = prefill[:, :, -1] # (b, d)
# Decode: feed tokens one at a time, carrying the conv state.
# The kernel mutates conv_state in place (shift left, append x).
conv_state = x[:, :, : width - 1].contiguous().clone() # (b, d, width-1)
for t in range(width - 1, seq_len):
token = x[:, :, t] # (b, d)
out = causal_conv1d.causal_conv1d_update(
token, conv_state, weight, bias, activation="silu"
)
torch.testing.assert_close(out, prefill[:, :, t], atol=1e-4, rtol=1e-4)
# Final decode-step output equals the prefill last token.
torch.testing.assert_close(out, last_token, atol=1e-4, rtol=1e-4)
def test_causal_conv1d_backward_fp32():
"""Autograd backward produces finite gradients for all inputs."""
torch.manual_seed(3)
b, d, seq_len, width = 1, 8, 16, 4
x = torch.randn(b, d, seq_len, dtype=torch.float32, device=DEVICE, requires_grad=True)
weight = torch.randn(d, width, dtype=torch.float32, device=DEVICE, requires_grad=True)
bias = torch.randn(d, dtype=torch.float32, device=DEVICE, requires_grad=True)
out = causal_conv1d.causal_conv1d_fn(x, weight, bias, activation="silu")
loss = out.float().square().mean()
loss.backward()
torch.cuda.synchronize()
assert torch.isfinite(x.grad).all(), "dx contains non-finite values"
assert torch.isfinite(weight.grad).all(), "dweight contains non-finite values"
assert torch.isfinite(bias.grad).all(), "dbias contains non-finite values"
def test_causal_conv1d_backward_fp16():
"""Autograd backward produces finite gradients in FP16."""
torch.manual_seed(4)
b, d, seq_len, width = 1, 16, 32, 4
x = torch.randn(b, d, seq_len, dtype=torch.float16, device=DEVICE, requires_grad=True)
weight = torch.randn(d, width, dtype=torch.float16, device=DEVICE, requires_grad=True)
bias = torch.randn(d, dtype=torch.float16, device=DEVICE, requires_grad=True)
out = causal_conv1d.causal_conv1d_fn(x, weight, bias, activation="silu")
loss = out.float().square().mean()
loss.backward()
torch.cuda.synchronize()
assert torch.isfinite(x.grad).all(), "dx contains non-finite values (fp16)"
assert torch.isfinite(weight.grad).all(), "dweight contains non-finite values (fp16)"
assert torch.isfinite(bias.grad).all(), "dbias contains non-finite values (fp16)"
def test_causal_conv1d_update_fp16():
"""Decode update works in FP16."""
torch.manual_seed(5)
b, d, width = 1, 16, 4
x_tok = torch.randn(b, d, dtype=torch.float16, device=DEVICE)
conv_state = torch.randn(b, d, width - 1, dtype=torch.float16, device=DEVICE)
weight = torch.randn(d, width, dtype=torch.float16, device=DEVICE)
bias = torch.randn(d, dtype=torch.float16, device=DEVICE)
out = causal_conv1d.causal_conv1d_update(
x_tok, conv_state, weight, bias, activation="silu"
)
assert out.shape == (b, d)
assert out.dtype == torch.float16
assert torch.isfinite(out).all()
def test_causal_conv1d_update_mutates_state_inplace():
"""conv_state is mutated in-place: shift-left + append x.
Transformers' fast-path decode passes conv_state to causal_conv1d_update
and expects the operation to update that state in-place. If the kernel
does not mutate, the second decoded token will see stale state.
Expected mutation:
state[:, :, 0:W-2] = old_state[:, :, 1:W-1] (shift left)
state[:, :, W-1] = x[:] (append new token)
"""
torch.manual_seed(6)
b, d, width = 2, 8, 4
x_tok = torch.randn(b, d, dtype=torch.float32, device=DEVICE)
conv_state = torch.randn(b, d, width - 1, dtype=torch.float32, device=DEVICE)
weight = torch.randn(d, width, dtype=torch.float32, device=DEVICE)
bias = torch.randn(d, dtype=torch.float32, device=DEVICE)
# Snapshot state before the call
state_before = conv_state.clone()
# Expected state after mutation:
# state[:, :, 0] = old_state[:, :, 1]
# state[:, :, 1] = old_state[:, :, 2] (= x_tok, since W-1=3 positions)
# Wait — state shape is (B, D, W-1) = (2, 8, 3). The kernel shifts
# state[:, :, 0:W-2] = state[:, :, 1:W-1] then sets state[:, :, W-2] = x.
# So: state[:, :, 0] = old[:, :, 1], state[:, :, 1] = old[:, :, 2],
# state[:, :, 2] = x_tok (W-2 = 1, so only index 1 shifts, index 2 gets x)
expected_state = state_before.clone()
expected_state[:, :, : width - 2] = state_before[:, :, 1 :]
expected_state[:, :, width - 2] = x_tok
# Capture data_ptr before call to verify in-place mutation
ptr_before = conv_state.data_ptr()
# Run the update
out = causal_conv1d.causal_conv1d_update(
x_tok, conv_state, weight, bias, activation="identity"
)
torch.cuda.synchronize()
# Verify the tensor is the SAME allocation (not replaced)
assert conv_state.data_ptr() == ptr_before, (
"conv_state was replaced with a new tensor, not mutated in-place"
)
# Verify in-place mutation: conv_state should now equal expected_state
torch.testing.assert_close(
conv_state,
expected_state,
atol=0,
rtol=0,
msg="conv_state was NOT mutated in-place by causal_conv1d_update",
)
# Verify output is finite
assert torch.isfinite(out).all(), "output contains non-finite values"
# Verify multi-step decode: run 3 more tokens, checking state mutation each time
for step in range(3):
x_tok = torch.randn(b, d, dtype=torch.float32, device=DEVICE)
state_before = conv_state.clone()
expected_state = state_before.clone()
expected_state[:, :, : width - 2] = state_before[:, :, 1 :]
expected_state[:, :, width - 2] = x_tok
out = causal_conv1d.causal_conv1d_update(
x_tok, conv_state, weight, bias, activation="identity"
)
torch.cuda.synchronize()
torch.testing.assert_close(
conv_state,
expected_state,
atol=0,
rtol=0,
msg=f"conv_state mutation failed at step {step + 1}",
)
assert torch.isfinite(out).all()