ObsidianSmall-Base / runtime /litgpt /multiscreen_triton.py
Metris's picture
Upload 78 files
236083b verified
Raw
History Blame Contribute Delete
24.7 kB
"""Fused Triton kernels for Multiscreen screening.
Neither forward nor backward materializes the B x H x T x T relevance tensor.
The backward recomputes relevance in tiles, like FlashAttention, and accumulates
the Q/K/V and two scalar screening-parameter gradients directly.
"""
from __future__ import annotations
import os
import torch
import triton
import triton.language as tl
_BLOCK_Q = int(os.environ.get("MULTISCREEN_BLOCK_Q", "32"))
_BLOCK_K = int(os.environ.get("MULTISCREEN_BLOCK_K", "32"))
_NUM_WARPS = int(os.environ.get("MULTISCREEN_NUM_WARPS", "4"))
_NUM_STAGES = int(os.environ.get("MULTISCREEN_NUM_STAGES", "2"))
_GRAD_RELEVANCE_BLOCK_Q = int(
os.environ.get("MULTISCREEN_GRAD_RELEVANCE_BLOCK_Q", "32")
)
@triton.jit
def _screening_window_tables(
window,
softmask_table,
window_grad_table,
sequence_length: tl.constexpr,
block: tl.constexpr,
):
head = tl.program_id(0)
distance = tl.arange(0, block)
head_window = tl.load(window + head).to(tl.float32)
infinite = head_window < 0.0
safe_window = tl.where(infinite, 1.0, head_window)
phase = 3.141592653589793 * distance / safe_window
valid = distance < sequence_length
softmask = tl.where(infinite, 1.0, 0.5 * (tl.cos(phase) + 1.0))
window_grad = tl.where(
infinite,
0.0,
0.5
* tl.sin(phase)
* 3.141592653589793
* distance
/ (safe_window * safe_window),
)
tl.store(softmask_table + head * sequence_length + distance, softmask, mask=valid)
tl.store(
window_grad_table + head * sequence_length + distance,
window_grad,
mask=valid,
)
@triton.jit
def _screening_aggregate_fwd(
query,
key,
value,
acceptance_width,
window,
output,
sequence_length: tl.constexpr,
num_heads: tl.constexpr,
key_dim: tl.constexpr,
value_dim: tl.constexpr,
is_bf16: tl.constexpr,
block_q: tl.constexpr,
block_k: tl.constexpr,
):
query_block = tl.program_id(0)
head = tl.program_id(1)
batch = tl.program_id(2)
query_offsets = query_block * block_q + tl.arange(0, block_q)
key_features = tl.arange(0, key_dim)
value_features = tl.arange(0, value_dim)
base = (batch * num_heads + head) * sequence_length
query_tile = tl.load(
query + (base + query_offsets[:, None]) * key_dim + key_features[None, :],
mask=query_offsets[:, None] < sequence_length,
other=0.0,
)
acceptance = tl.load(acceptance_width + head).to(tl.float32)
head_window = tl.load(window + head).to(tl.float32)
accumulator = tl.zeros((block_q, value_dim), dtype=tl.float32)
infinite_window = head_window < 0.0
safe_window = tl.where(infinite_window, 1.0, head_window)
integer_window = tl.where(
infinite_window, sequence_length, tl.ceil(head_window).to(tl.int32)
)
# Do not cap this by sequence_length: early query blocks begin at a
# negative key offset, and those masked iterations are still needed to
# reach keys 0..query_end.
window_span = integer_window + block_q
first_key = query_block * block_q - integer_window
for key_step in tl.range(0, window_span, block_k):
key_offsets = first_key + key_step + tl.arange(0, block_k)
key_tile = tl.load(
key + (base + key_offsets[:, None]) * key_dim + key_features[None, :],
mask=(key_offsets[:, None] >= 0) & (key_offsets[:, None] < sequence_length),
other=0.0,
)
if is_bf16:
similarity = tl.dot(query_tile, tl.trans(key_tile))
# PyTorch autocast exposes the BF16-rounded similarity to the
# nonlinear screening function. Match that boundary exactly.
similarity = similarity.to(tl.bfloat16).to(tl.float32)
else:
similarity = tl.dot(query_tile, tl.trans(key_tile), input_precision="ieee")
distance = key_offsets[None, :] - query_offsets[:, None]
valid = (
(query_offsets[:, None] < sequence_length)
& (key_offsets[None, :] < sequence_length)
& (distance <= 0)
& (infinite_window | (distance > -safe_window))
)
trimmed = tl.maximum(1.0 - (1.0 - similarity) / acceptance, 0.0)
softmask = tl.where(
infinite_window,
1.0,
0.5 * (tl.cos(3.141592653589793 * distance / safe_window) + 1.0),
)
relevance = tl.where(valid, trimmed * trimmed * softmask, 0.0)
value_tile = tl.load(
value + (base + key_offsets[:, None]) * value_dim + value_features[None, :],
mask=(key_offsets[:, None] >= 0) & (key_offsets[:, None] < sequence_length),
other=0.0,
)
if is_bf16:
accumulator += tl.dot(relevance.to(tl.bfloat16), value_tile)
else:
accumulator += tl.dot(relevance, value_tile, input_precision="ieee")
tl.store(
output + (base + query_offsets[:, None]) * value_dim + value_features[None, :],
accumulator,
mask=query_offsets[:, None] < sequence_length,
)
@triton.jit
def _screening_grad_relevance(
value,
grad_output,
window,
grad_relevance_workspace,
sequence_length: tl.constexpr,
num_heads: tl.constexpr,
value_dim: tl.constexpr,
is_bf16: tl.constexpr,
block_q: tl.constexpr,
block_k: tl.constexpr,
):
"""Compute dL/d(relevance) once for reuse by both backward ownership passes."""
query_block = tl.program_id(0)
head = tl.program_id(1)
batch = tl.program_id(2)
query_offsets = query_block * block_q + tl.arange(0, block_q)
value_features = tl.arange(0, value_dim)
base = (batch * num_heads + head) * sequence_length
query_mask = query_offsets < sequence_length
grad_output_tile = tl.load(
grad_output
+ (base + query_offsets[:, None]) * value_dim
+ value_features[None, :],
mask=query_mask[:, None],
other=0.0,
)
head_window = tl.load(window + head).to(tl.float32)
integer_window = tl.ceil(head_window).to(tl.int32)
window_span = integer_window + block_q
first_key = query_block * block_q - integer_window
for key_step in tl.range(0, window_span, block_k):
key_offsets = first_key + key_step + tl.arange(0, block_k)
key_mask = (key_offsets >= 0) & (key_offsets < sequence_length)
value_tile = tl.load(
value
+ (base + key_offsets[:, None]) * value_dim
+ value_features[None, :],
mask=key_mask[:, None],
other=0.0,
)
if is_bf16:
grad_relevance = tl.dot(grad_output_tile, tl.trans(value_tile))
grad_relevance = grad_relevance.to(tl.bfloat16)
else:
grad_relevance = tl.dot(
grad_output_tile, tl.trans(value_tile), input_precision="ieee"
)
distance = key_offsets[None, :] - query_offsets[:, None]
valid = (
query_mask[:, None]
& key_mask[None, :]
& (distance <= 0)
& (distance > -head_window)
)
triangular_size = sequence_length * (sequence_length + 1) // 2
owner = batch * num_heads + head
workspace_offsets = (
owner * triangular_size
+ query_offsets[:, None] * (query_offsets[:, None] + 1) // 2
+ key_offsets[None, :]
)
tl.store(
grad_relevance_workspace + workspace_offsets,
grad_relevance,
mask=valid,
)
@triton.jit
def _screening_aggregate_bwd(
query,
key,
value,
grad_output,
acceptance_width,
window,
grad_query,
grad_acceptance,
grad_window,
sequence_length: tl.constexpr,
num_heads: tl.constexpr,
key_dim: tl.constexpr,
value_dim: tl.constexpr,
is_bf16: tl.constexpr,
block_q: tl.constexpr,
block_k: tl.constexpr,
):
query_block = tl.program_id(0)
head = tl.program_id(1)
batch = tl.program_id(2)
query_offsets = query_block * block_q + tl.arange(0, block_q)
key_features = tl.arange(0, key_dim)
value_features = tl.arange(0, value_dim)
base = (batch * num_heads + head) * sequence_length
query_mask = query_offsets < sequence_length
query_tile = tl.load(
query + (base + query_offsets[:, None]) * key_dim + key_features[None, :],
mask=query_mask[:, None],
other=0.0,
)
grad_output_tile = tl.load(
grad_output
+ (base + query_offsets[:, None]) * value_dim
+ value_features[None, :],
mask=query_mask[:, None],
other=0.0,
)
acceptance = tl.load(acceptance_width + head).to(tl.float32)
head_window = tl.load(window + head).to(tl.float32)
grad_query_acc = tl.zeros((block_q, key_dim), dtype=tl.float32)
grad_acceptance_acc = tl.zeros((), dtype=tl.float32)
grad_window_acc = tl.zeros((), dtype=tl.float32)
integer_window = tl.ceil(head_window).to(tl.int32)
window_span = integer_window + block_q
first_key = query_block * block_q - integer_window
for key_step in tl.range(0, window_span, block_k):
key_offsets = first_key + key_step + tl.arange(0, block_k)
key_mask = (key_offsets >= 0) & (key_offsets < sequence_length)
key_tile = tl.load(
key + (base + key_offsets[:, None]) * key_dim + key_features[None, :],
mask=key_mask[:, None],
other=0.0,
)
value_tile = tl.load(
value
+ (base + key_offsets[:, None]) * value_dim
+ value_features[None, :],
mask=key_mask[:, None],
other=0.0,
)
if is_bf16:
similarity = tl.dot(query_tile, tl.trans(key_tile))
similarity = similarity.to(tl.bfloat16).to(tl.float32)
else:
similarity = tl.dot(query_tile, tl.trans(key_tile), input_precision="ieee")
distance = key_offsets[None, :] - query_offsets[:, None]
valid = (
query_mask[:, None]
& key_mask[None, :]
& (distance <= 0)
& (distance > -head_window)
)
if is_bf16:
grad_relevance = tl.dot(grad_output_tile, tl.trans(value_tile))
grad_relevance = grad_relevance.to(tl.bfloat16).to(tl.float32)
else:
grad_relevance = tl.dot(
grad_output_tile, tl.trans(value_tile), input_precision="ieee"
)
raw_trim = 1.0 - (1.0 - similarity) / acceptance
trim = tl.maximum(raw_trim, 0.0)
phase = 3.141592653589793 * distance / head_window
softmask = 0.5 * (tl.cos(phase) + 1.0)
relevance = tl.where(valid, trim * trim * softmask, 0.0)
active = valid & (raw_trim > 0.0)
grad_similarity = tl.where(
active,
grad_relevance * 2.0 * trim * softmask / acceptance,
0.0,
)
if is_bf16:
grad_query_acc += tl.dot(grad_similarity.to(tl.bfloat16), key_tile)
else:
grad_query_acc += tl.dot(
grad_similarity, key_tile, input_precision="ieee"
)
grad_acceptance_acc += tl.sum(
tl.where(
active,
grad_relevance
* 2.0
* trim
* softmask
* (1.0 - similarity)
/ (acceptance * acceptance),
0.0,
)
)
grad_softmask_window = (
0.5
* tl.sin(phase)
* 3.141592653589793
* distance
/ (head_window * head_window)
)
grad_window_acc += tl.sum(
tl.where(valid, grad_relevance * trim * trim * grad_softmask_window, 0.0)
)
tl.store(
grad_query
+ (base + query_offsets[:, None]) * key_dim
+ key_features[None, :],
grad_query_acc,
mask=query_mask[:, None],
)
tl.atomic_add(grad_acceptance + head, grad_acceptance_acc)
tl.atomic_add(grad_window + head, grad_window_acc)
@triton.jit
def _screening_aggregate_bwd_kv(
query,
key,
value,
grad_output,
acceptance_width,
window,
grad_key,
grad_value,
sequence_length: tl.constexpr,
num_heads: tl.constexpr,
key_dim: tl.constexpr,
value_dim: tl.constexpr,
is_bf16: tl.constexpr,
block_q: tl.constexpr,
block_k: tl.constexpr,
):
"""Key-major dK/dV pass with exclusive output ownership and no atomics."""
key_block = tl.program_id(0)
head = tl.program_id(1)
batch = tl.program_id(2)
key_offsets = key_block * block_k + tl.arange(0, block_k)
key_features = tl.arange(0, key_dim)
value_features = tl.arange(0, value_dim)
base = (batch * num_heads + head) * sequence_length
key_mask = key_offsets < sequence_length
key_tile = tl.load(
key + (base + key_offsets[:, None]) * key_dim + key_features[None, :],
mask=key_mask[:, None],
other=0.0,
)
value_tile = tl.load(
value + (base + key_offsets[:, None]) * value_dim + value_features[None, :],
mask=key_mask[:, None],
other=0.0,
)
acceptance = tl.load(acceptance_width + head).to(tl.float32)
head_window = tl.load(window + head).to(tl.float32)
grad_key_acc = tl.zeros((block_k, key_dim), dtype=tl.float32)
grad_value_acc = tl.zeros((block_k, value_dim), dtype=tl.float32)
# A key contributes only to causal queries from its own position through
# key + window - 1. The extra block_k covers all keys in this tile.
query_span = tl.ceil(head_window).to(tl.int32) + block_k
first_query = key_block * block_k
for query_step in tl.range(0, query_span, block_q):
query_offsets = first_query + query_step + tl.arange(0, block_q)
query_mask = query_offsets < sequence_length
query_tile = tl.load(
query + (base + query_offsets[:, None]) * key_dim + key_features[None, :],
mask=query_mask[:, None],
other=0.0,
)
grad_output_tile = tl.load(
grad_output
+ (base + query_offsets[:, None]) * value_dim
+ value_features[None, :],
mask=query_mask[:, None],
other=0.0,
)
if is_bf16:
similarity = tl.dot(query_tile, tl.trans(key_tile))
similarity = similarity.to(tl.bfloat16).to(tl.float32)
else:
similarity = tl.dot(
query_tile, tl.trans(key_tile), input_precision="ieee"
)
distance = key_offsets[None, :] - query_offsets[:, None]
valid = (
query_mask[:, None]
& key_mask[None, :]
& (distance <= 0)
& (distance > -head_window)
)
if is_bf16:
grad_relevance = tl.dot(grad_output_tile, tl.trans(value_tile))
grad_relevance = grad_relevance.to(tl.bfloat16).to(tl.float32)
else:
grad_relevance = tl.dot(
grad_output_tile, tl.trans(value_tile), input_precision="ieee"
)
raw_trim = 1.0 - (1.0 - similarity) / acceptance
trim = tl.maximum(raw_trim, 0.0)
phase = 3.141592653589793 * distance / head_window
softmask = 0.5 * (tl.cos(phase) + 1.0)
relevance = tl.where(valid, trim * trim * softmask, 0.0)
grad_similarity = tl.where(
valid & (raw_trim > 0.0),
grad_relevance * 2.0 * trim * softmask / acceptance,
0.0,
)
if is_bf16:
grad_key_acc += tl.dot(
tl.trans(grad_similarity).to(tl.bfloat16), query_tile
)
grad_value_acc += tl.dot(
tl.trans(relevance).to(tl.bfloat16), grad_output_tile
)
else:
grad_key_acc += tl.dot(
tl.trans(grad_similarity), query_tile, input_precision="ieee"
)
grad_value_acc += tl.dot(
tl.trans(relevance), grad_output_tile, input_precision="ieee"
)
tl.store(
grad_key + (base + key_offsets[:, None]) * key_dim + key_features[None, :],
grad_key_acc,
mask=key_mask[:, None],
)
tl.store(
grad_value
+ (base + key_offsets[:, None]) * value_dim
+ value_features[None, :],
grad_value_acc,
mask=key_mask[:, None],
)
@triton.jit
def _screening_aggregate_bwd_v(
query,
key,
grad_output,
acceptance_width,
window,
grad_value,
sequence_length: tl.constexpr,
num_heads: tl.constexpr,
key_dim: tl.constexpr,
value_dim: tl.constexpr,
is_bf16: tl.constexpr,
block_q: tl.constexpr,
block_k: tl.constexpr,
block_v: tl.constexpr,
):
key_blocks = tl.cdiv(sequence_length, block_k)
value_blocks = tl.cdiv(value_dim, block_v)
combined_block = tl.program_id(0)
key_block = combined_block // value_blocks
value_block = combined_block - key_block * value_blocks
head = tl.program_id(1)
batch = tl.program_id(2)
key_offsets = key_block * block_k + tl.arange(0, block_k)
value_features = value_block * block_v + tl.arange(0, block_v)
key_features = tl.arange(0, key_dim)
base = (batch * num_heads + head) * sequence_length
key_mask = key_offsets < sequence_length
value_mask = value_features < value_dim
key_tile = tl.load(
key + (base + key_offsets[:, None]) * key_dim + key_features[None, :],
mask=key_mask[:, None],
other=0.0,
)
acceptance = tl.load(acceptance_width + head).to(tl.float32)
head_window = tl.load(window + head).to(tl.float32)
grad_value_acc = tl.zeros((block_k, block_v), dtype=tl.float32)
query_span = tl.ceil(head_window).to(tl.int32) + block_k
first_query = key_block * block_k
for query_step in tl.range(0, query_span, block_q):
query_offsets = first_query + query_step + tl.arange(0, block_q)
query_mask = query_offsets < sequence_length
query_tile = tl.load(
query + (base + query_offsets[:, None]) * key_dim + key_features[None, :],
mask=query_mask[:, None],
other=0.0,
)
grad_output_tile = tl.load(
grad_output
+ (base + query_offsets[:, None]) * value_dim
+ value_features[None, :],
mask=query_mask[:, None] & value_mask[None, :],
other=0.0,
)
if is_bf16:
similarity = tl.dot(query_tile, tl.trans(key_tile))
similarity = similarity.to(tl.bfloat16).to(tl.float32)
else:
similarity = tl.dot(
query_tile, tl.trans(key_tile), input_precision="ieee"
)
distance = key_offsets[None, :] - query_offsets[:, None]
valid = (
query_mask[:, None]
& key_mask[None, :]
& (distance <= 0)
& (distance > -head_window)
)
raw_trim = 1.0 - (1.0 - similarity) / acceptance
trim = tl.maximum(raw_trim, 0.0)
phase = 3.141592653589793 * distance / head_window
softmask = 0.5 * (tl.cos(phase) + 1.0)
relevance = tl.where(valid, trim * trim * softmask, 0.0)
if is_bf16:
grad_value_acc += tl.dot(
tl.trans(relevance).to(tl.bfloat16), grad_output_tile
)
else:
grad_value_acc += tl.dot(
tl.trans(relevance), grad_output_tile, input_precision="ieee"
)
tl.store(
grad_value
+ (base + key_offsets[:, None]) * value_dim
+ value_features[None, :],
grad_value_acc,
mask=key_mask[:, None] & value_mask[None, :],
)
def _validate_inputs(
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
) -> tuple[int, int, int, int, int]:
if not (
query.is_cuda
and query.dtype in {torch.bfloat16, torch.float32}
and key.dtype == query.dtype
and value.dtype == query.dtype
and query.is_contiguous()
and key.is_contiguous()
and value.is_contiguous()
):
raise ValueError("The Multiscreen Triton kernel requires contiguous CUDA BF16/FP32 Q, K, and V tensors.")
batch, heads, sequence_length, key_dim = query.shape
value_dim = value.shape[-1]
if key_dim != 16 or value_dim not in {32, 64, 128}:
raise ValueError(f"Unsupported Multiscreen Triton shape: key_dim={key_dim}, value_dim={value_dim}.")
return batch, heads, sequence_length, key_dim, value_dim
def _screening_forward(
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
acceptance_width: torch.Tensor,
window: torch.Tensor,
) -> torch.Tensor:
batch, heads, sequence_length, key_dim, value_dim = _validate_inputs(query, key, value)
output = torch.empty_like(value)
block_q = _BLOCK_Q
block_k = _BLOCK_K
grid = (triton.cdiv(sequence_length, block_q), heads, batch)
_screening_aggregate_fwd[grid](
query,
key,
value,
acceptance_width,
window,
output,
sequence_length=sequence_length,
num_heads=heads,
key_dim=key_dim,
value_dim=value_dim,
is_bf16=query.dtype == torch.bfloat16,
block_q=block_q,
block_k=block_k,
num_warps=_NUM_WARPS,
num_stages=_NUM_STAGES,
)
return output
class _ScreeningAggregate(torch.autograd.Function):
@staticmethod
def forward(ctx, query, key, value, acceptance_width, window):
output = _screening_forward(query, key, value, acceptance_width, window)
ctx.save_for_backward(query, key, value, acceptance_width, window)
return output
@staticmethod
def backward(ctx, grad_output):
query, key, value, acceptance_width, window = ctx.saved_tensors
batch, heads, sequence_length, key_dim, value_dim = _validate_inputs(
query, key, value
)
grad_output = grad_output.contiguous()
grad_query = torch.empty_like(query)
# Accumulate in the model dtype. Ada and newer GPUs implement BF16
# atomics, and keeping these two large buffers in BF16 is what makes
# batch-32 practical on an 8 GB card.
grad_key = torch.zeros_like(key)
grad_value = torch.zeros_like(value)
grad_acceptance = torch.zeros_like(acceptance_width, dtype=torch.float32)
grad_window = torch.zeros_like(window, dtype=torch.float32)
block_q = _BLOCK_Q
block_k = _BLOCK_K
grid = (triton.cdiv(sequence_length, block_q), heads, batch)
_screening_aggregate_bwd[grid](
query,
key,
value,
grad_output,
acceptance_width,
window,
grad_query,
grad_acceptance,
grad_window,
sequence_length=sequence_length,
num_heads=heads,
key_dim=key_dim,
value_dim=value_dim,
is_bf16=query.dtype == torch.bfloat16,
block_q=block_q,
block_k=block_k,
num_warps=_NUM_WARPS,
num_stages=_NUM_STAGES,
)
kv_grid = (triton.cdiv(sequence_length, block_k), heads, batch)
_screening_aggregate_bwd_kv[kv_grid](
query,
key,
value,
grad_output,
acceptance_width,
window,
grad_key,
grad_value,
sequence_length=sequence_length,
num_heads=heads,
key_dim=key_dim,
value_dim=value_dim,
is_bf16=query.dtype == torch.bfloat16,
block_q=block_q,
block_k=block_k,
num_warps=_NUM_WARPS,
num_stages=_NUM_STAGES,
)
return (
grad_query,
grad_key,
grad_value,
grad_acceptance.to(acceptance_width.dtype),
grad_window.to(window.dtype),
)
def screening_aggregate_triton(
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
acceptance_width: torch.Tensor,
window: torch.Tensor,
) -> torch.Tensor:
"""Fused paper-equivalent screening with a memory-linear custom backward."""
if torch.is_grad_enabled() and any(
tensor.requires_grad
for tensor in (query, key, value, acceptance_width, window)
):
return _ScreeningAggregate.apply(
query, key, value, acceptance_width, window
)
return _screening_forward(query, key, value, acceptance_width, window)