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