///|
/// Materializable additive attention-mask data.
///
/// The shape is `[1, tokens, tokens]` so it broadcasts across attention heads.
/// Allowed score positions contain zero and masked positions contain the
/// caller-provided negative value.
pub struct AttentionMask {
shape_ : @shape.Shape
values_ : Array[Float]
}
///|
pub fn AttentionMask::shape(self : AttentionMask) -> @shape.Shape {
self.shape_
}
///|
pub fn AttentionMask::values(self : AttentionMask) -> Array[Float] {
self.values_.copy()
}
///|
pub fn AttentionMask::additive(
shape : @shape.Shape,
values : Array[Float],
) -> AttentionMask raise @tensor.TensorError {
if shape.rank() != 3 && shape.rank() != 4 {
raise @tensor.TensorError::new(
"additive attention mask must have rank 3 or 4",
)
}
if shape.dimension(shape.rank() - 2) != shape.dimension(shape.rank() - 1) {
raise @tensor.TensorError::new(
"additive self-attention mask query and key dimensions must match",
)
}
if values.length() != shape.element_count() {
raise @tensor.TensorError::new(
"attention mask data length \{values.length()} does not match shape \{shape.to_string()}",
)
}
{ shape_: shape, values_: values.copy() }
}
///|
fn make_attention_mask(
valid_tokens : Array[Bool]?,
tokens : Int,
causal : Bool,
masked_value : Float,
) -> AttentionMask raise @tensor.TensorError {
if tokens <= 0 {
raise @tensor.TensorError::new(
"attention mask must contain at least one token",
)
}
if !(masked_value < 0.0) {
raise @tensor.TensorError::new("attention masked value must be negative")
}
match valid_tokens {
Some(valid) =>
if valid.length() != tokens {
raise @tensor.TensorError::new(
"attention padding mask length must match token count",
)
}
None => ()
}
let values = Array::make(tokens * tokens, Float::from_int(0))
for query in 0.. !valid[key]
None => false
}
if padding_masked || (causal && key > query) {
values[query * tokens + key] = masked_value
}
}
}
let shape = @shape.Shape::new([1, tokens, tokens]) catch {
error => raise @tensor.TensorError::new(error.to_string())
}
{ shape_: shape, values_: values }
}
///|
pub fn AttentionMask::causal(
tokens : Int,
masked_value : Float,
) -> AttentionMask raise @tensor.TensorError {
make_attention_mask(None, tokens, true, masked_value)
}
///|
pub fn AttentionMask::padding(
valid_tokens : Array[Bool],
masked_value : Float,
) -> AttentionMask raise @tensor.TensorError {
make_attention_mask(
Some(valid_tokens),
valid_tokens.length(),
false,
masked_value,
)
}
///|
pub fn AttentionMask::causal_padding(
valid_tokens : Array[Bool],
masked_value : Float,
) -> AttentionMask raise @tensor.TensorError {
make_attention_mask(
Some(valid_tokens),
valid_tokens.length(),
true,
masked_value,
)
}
///|
fn make_batched_attention_mask(
valid_batches : Array[Array[Bool]],
causal : Bool,
masked_value : Float,
) -> AttentionMask raise @tensor.TensorError {
if valid_batches.length() == 0 {
raise @tensor.TensorError::new(
"batched attention mask must contain at least one batch",
)
}
if !(masked_value < 0.0) {
raise @tensor.TensorError::new("attention masked value must be negative")
}
let tokens = valid_batches[0].length()
if tokens == 0 {
raise @tensor.TensorError::new(
"attention mask must contain at least one token",
)
}
for batch, valid_tokens in valid_batches {
if valid_tokens.length() != tokens {
raise @tensor.TensorError::new(
"attention batch \{batch} token count \{valid_tokens.length()} must match \{tokens}",
)
}
}
let values = Array::make(
valid_batches.length() * tokens * tokens,
Float::from_int(0),
)
for batch, valid_tokens in valid_batches {
let batch_offset = batch * tokens * tokens
for query in 0.. query) {
values[batch_offset + query * tokens + key] = masked_value
}
}
}
}
let shape = @shape.Shape::new([valid_batches.length(), 1, tokens, tokens]) catch {
error => raise @tensor.TensorError::new(error.to_string())
}
{ shape_: shape, values_: values }
}
///|
pub fn AttentionMask::batched_padding(
valid_batches : Array[Array[Bool]],
masked_value : Float,
) -> AttentionMask raise @tensor.TensorError {
make_batched_attention_mask(valid_batches, false, masked_value)
}
///|
pub fn AttentionMask::batched_causal_padding(
valid_batches : Array[Array[Bool]],
masked_value : Float,
) -> AttentionMask raise @tensor.TensorError {
make_batched_attention_mask(valid_batches, true, masked_value)
}
///|
/// Fixed-shape self-attention shared by CPU and WebNN backends.
///
/// Input and output use `[tokens, width]` or `[batch, tokens, width]`.
/// Internally, projections are reshaped by head; batched matmul computes every
/// batch and head together. This layer returns only the projected attention;
/// residual addition belongs to a Transformer block.
pub(all) struct SelfAttention[T] {
query : Linear[T]
key : Linear[T]
value : Linear[T]
output : Linear[T]
score_scale : T
additive_mask : T?
heads : Int
}
///|
fn[T : @tensor.TensorOps] require_attention_shape(
name : String,
tensor : T,
expected : @shape.Shape,
) -> Unit raise @tensor.TensorError {
if !tensor.shape().same_as(expected) {
raise @tensor.TensorError::new(
"attention \{name} shape \{tensor.shape().to_string()} must match \{expected.to_string()}",
)
}
}
///|
pub fn[T : @tensor.TensorOps] SelfAttention::forward(
self : SelfAttention[T],
input : T,
) -> T raise {
let input_shape = input.shape()
let rank = input_shape.rank()
if rank != 2 && rank != 3 {
raise @tensor.TensorError::new(
"self-attention input must have shape [tokens, width] or [batch, tokens, width]",
)
}
if self.heads <= 0 {
raise @tensor.TensorError::new("self-attention heads must be positive")
}
let tokens = input_shape.dimension(rank - 2)
let width = input_shape.dimension(rank - 1)
if width % self.heads != 0 {
raise @tensor.TensorError::new(
"self-attention width \{width} must be divisible by \{self.heads} heads",
)
}
let query = self.query.forward(input)
let key = self.key.forward(input)
let value = self.value.forward(input)
require_attention_shape("query projection", query, input_shape)
require_attention_shape("key projection", key, input_shape)
require_attention_shape("value projection", value, input_shape)
let head_size = width / self.heads
let split_shape = if rank == 2 {
@shape.Shape::new([tokens, self.heads, head_size])
} else {
@shape.Shape::new([input_shape.dimension(0), tokens, self.heads, head_size])
}
let query_permutation = if rank == 2 { [1, 0, 2] } else { [0, 2, 1, 3] }
let key_permutation = if rank == 2 { [1, 2, 0] } else { [0, 2, 3, 1] }
let query_heads = query.reshape(split_shape).transpose(query_permutation)
let key_heads = key.reshape(split_shape).transpose(key_permutation)
let value_heads = value.reshape(split_shape).transpose(query_permutation)
let score_shape = if rank == 2 {
@shape.Shape::new([self.heads, tokens, tokens])
} else {
@shape.Shape::new([input_shape.dimension(0), self.heads, tokens, tokens])
}
let scaled_scores = query_heads.matmul(key_heads).mul(self.score_scale)
require_attention_shape("scaled score", scaled_scores, score_shape)
let masked_scores = match self.additive_mask {
Some(mask) => {
let result = scaled_scores.add(mask)
require_attention_shape("masked score", result, score_shape)
result
}
None => scaled_scores
}
let attention = masked_scores.softmax(rank)
let context = attention
.matmul(value_heads)
.transpose(query_permutation)
.reshape(input_shape)
let projected = self.output.forward(context)
require_attention_shape("output projection", projected, input_shape)
projected
}