// t5_relative_position.mbt — T5-style relative position bias
// (Raffel 2020 §3.2.3) for self-attention.
//
// Unlike absolute position embeddings (added to the input), T5 uses
// learned *relative* position biases that are added to the attention
// scores:
//
//   scores += bias[h, i, j]
//
// where the bias is indexed by the *relative* offset (j - i) clamped
// to `[-max_distance, max_distance]`. There is one bias table per
// attention head, shape `(num_heads, 2*max_distance+1)`.
//
// For a query at position `i` attending to a key at position `j`:
//
//   offset = j - i
//   bucket = clamp(offset, -max_distance, max_distance) + max_distance
//   bias[h, i, j] = weights[h, bucket]
//
// The bias table is added to the existing scores *before* softmax.
// This implementation uses a flat row-major layout with one shared
// bucket range per head (no log-bucket scheme for simplicity).

///|
/// T5-style relative position bias parameter container.
pub struct T5RelativePosition {
  num_heads : Int
  max_distance : Int
  // bias table: shape (num_heads, 2 * max_distance + 1).
  weight : Array[Float]
}

///|
/// Construct a T5 relative position bias. `max_distance` controls
/// how far apart positions can be before the bias saturates (T5
/// paper uses 128; we use a smaller default for demos).
pub fn T5RelativePosition::new(
  num_heads : Int,
  max_distance : Int,
  seed : UInt64,
) -> T5RelativePosition {
  let n_buckets = 2 * max_distance + 1
  let rng = Xoshiro::from_state(seed, seed + 1UL, seed + 2UL, seed + 3UL)
  let std = sqrtf(1.0F / Float::from_int(n_buckets))
  let weight : Array[Float] = Array::make(num_heads * n_buckets, 0.0F)
  let mut i = 0
  while i < num_heads * n_buckets {
    let (z, _) = box_muller(rng)
    weight[i] = Float::from_double(z) * std
    i = i + 1
  }
  { num_heads, max_distance, weight }
}

///|
/// Compute the per-head bias matrix of shape
/// (num_heads × seq_len × seq_len). `bias[h, i, j]` is indexed by
/// the clamped relative offset (j - i).
pub fn t5_compute_bias(
  rp : T5RelativePosition,
  seq_len : Int,
) -> Array[Float] {
  let out : Array[Float] = Array::make(
    rp.num_heads * seq_len * seq_len, 0.0F,
  )
  let n_buckets = 2 * rp.max_distance + 1
  for h in 0.. rp.max_distance {
          rp.max_distance
        } else {
          offset
        }
        let bucket = clamped + rp.max_distance
        out[h * seq_len * seq_len + i * seq_len + j] =
          rp.weight[h * n_buckets + bucket]
      }
    }
  }
  out
}

///|
/// Backward: returns `d_weight` of shape (num_heads, 2*max_distance+1).
///
/// `d_bias[h, i, j]` is the upstream gradient on the bias entry at
/// (h, i, j). Sum over (i, j) bucketed by `bucket(j - i)`:
///
///   d_weight[h, b] = sum_{i, j: bucket(j-i)==b} d_bias[h, i, j]
pub fn t5_backward(
  rp : T5RelativePosition,
  d_bias : Array[Float],
  seq_len : Int,
) -> Array[Float] {
  let n_buckets = 2 * rp.max_distance + 1
  let d_weight : Array[Float] = Array::make(
    rp.num_heads * n_buckets, 0.0F,
  )
  for h in 0.. rp.max_distance {
          rp.max_distance
        } else {
          offset
        }
        let bucket = clamped + rp.max_distance
        d_weight[h * n_buckets + bucket] = d_weight[h * n_buckets + bucket] +
          d_bias[h * seq_len * seq_len + i * seq_len + j]
      }
    }
  }
  d_weight
}