// dropout.mbt — Inverted dropout regularisation for transformer training.
//
// Inverted dropout (PyTorch / TensorFlow convention):
//
//   Training:  mask[i] ~ Bernoulli(1 - p)
//              out[i]  = (input[i] * mask[i]) / (1 - p)
//   Inference: out[i]  = input[i]                       (identity)
//
// Backward:  d_input[i] = d_output[i] * mask[i] / (1 - p)
// (same scale, so expected gradient magnitude matches forward
// magnitude).
//
// The dropout mask is stored in the cache so that backward uses the
// SAME mask as forward (otherwise the gradient wouldn't correspond
// to the actual sparse activations that were computed).

///|
/// Dropout parameter container.
pub(all) struct Dropout {
  p : Float      // dropout probability (0 = off, 1 = drop all)
  mut training : Bool
}

///|
/// Construct a Dropout with dropout probability `p`. Defaults to 0.5.
pub fn Dropout::new(p? : Float = 0.5F) -> Dropout {
  if p < 0.0F || p > 1.0F {
    abort("Dropout::new: p \{p} out of [0, 1]")
  }
  { p, training: true }
}

///|
/// Forward pass. Returns `(out, mask)` where `mask` is the per-element
/// Bernoulli mask (length n). Use `mask` in the corresponding
/// backward pass to recover the exact same scaling.
pub fn dropout_forward(
  input : Array[Float],
  d : Dropout,
  rng : Xoshiro,
) -> (Array[Float], Array[Float]) {
  let n = input.length()
  let out : Array[Float] = Array::make(n, 0.0F)
  let mask : Array[Float] = Array::make(n, 1.0F)
  if !d.training || d.p == 0.0F {
    // Inference mode / no dropout: identity, mask = all ones.
    for i in 0.. Array[Float] {
  let n = d_output.length()
  let d_input : Array[Float] = Array::make(n, 0.0F)
  if !d.training || d.p == 0.0F {
    // Inference / off: identity.
    for i in 0..