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