// noisy_linear.mbt — NoisyLinear layer (v0.43.0).
//
// Fortunato et al. 2018 "Noisy Networks for Exploration" —
// factorised-Gaussian parametric noise added to a Linear layer's
// weights and biases. Used to replace epsilon-greedy exploration
// in DQN (see v0.43.2 NoisyDQN).
//
// Factorised Gaussian noise (per input/output dimension):
//   ε_i ~ N(0, 1)              for i in [0, in_features)
//   ε_j ~ N(0, 1)              for j in [0, out_features)
//   f(x) = sign(x) · sqrt(|x|)  (paper's reparameterisation; CPU-friendly)
//   ε_w[o, i] = f(ε_i)[i] · f(ε_j)[o]    (rank-1 outer product)
//   ε_b[o]   = f(ε_j)[o]
//
// Effective weight:
//   w[o, i] = weight_mu[o, i] + weight_sigma[o, i] · ε_w[o, i]
//   b[o]   = bias_mu[o]       + bias_sigma[o]     · ε_b[o]
//
// Forward: y[n, o] = sum_i w[o, i] · x[n, i] + b[o]
//
// Reuses `next_f32` (Box-Muller via Xoshiro RNG) for ε_i and ε_j.

///|
/// NoisyLinear parameter container. Stores the mean and std of
/// weights/biases; noise is sampled at forward time.
pub struct NoisyLinearParam {
  // μ_weight : [out_features, in_features]
  weight_mu : Array[Float]
  // σ_weight : [out_features, in_features]  (non-negative)
  weight_sigma : Array[Float]
  // μ_bias : [out_features]
  bias_mu : Array[Float]
  // σ_bias : [out_features]  (non-negative)
  bias_sigma : Array[Float]
  in_features : Int
  out_features : Int
}

///|
/// Build a NoisyLinearParam. `mu_*` / `sigma_*` arrays are not copied.
/// All `sigma_*` entries should be ≥ 0 (typical init: 0.017).
pub fn NoisyLinearParam::new(
  weight_mu : Array[Float],
  weight_sigma : Array[Float],
  bias_mu : Array[Float],
  bias_sigma : Array[Float],
  in_features : Int,
  out_features : Int,
) -> NoisyLinearParam {
  {
    weight_mu,
    weight_sigma,
    bias_mu,
    bias_sigma,
    in_features,
    out_features,
  }
}

///|
/// Init μ / σ weights with uniform in [-1/sqrt(in), +1/sqrt(in)]
/// (PyTorch's NoisyLinear default bound). σ weights init to 0.017.
/// bias μ init to 0, σ init to 0.017.
pub fn NoisyLinearParam::init(
  in_features : Int,
  out_features : Int,
  rng : Xoshiro,
) -> NoisyLinearParam {
  let bound : Float = 1.0F /
    Float::from_int(in_features).to_double().sqrt().to_float()
  let w_mu : Array[Float] = Array::make(out_features * in_features, 0.0F)
  let w_sig : Array[Float] = Array::make(out_features * in_features, 0.017F)
  let b_mu : Array[Float] = Array::make(out_features, 0.0F)
  let b_sig : Array[Float] = Array::make(out_features, 0.017F)
  for k in 0.. Float {
  if x >= 0.0F {
    x.to_double().sqrt().to_float()
  } else {
    -(-x).to_double().sqrt().to_float()
  }
}

///|
/// Sample N standard-normals via Box-Muller (consumes 2·N draws).
pub fn sample_gaussians(n : Int, rng : Xoshiro) -> Array[Float] {
  let out : Array[Float] = Array::make(n, 0.0F)
  let mut k = 0
  while k < n {
    let u1 = next_f32(rng)
    let u2 = next_f32(rng)
    // Box-Muller: z0 = sqrt(-2 ln u1) cos(2π u2)
    let safe_u1 = if u1 < 1.0e-7F { 1.0e-7F } else { u1 }
    let r : Float = (-2.0F *
      Float::from_double(@math.ln(safe_u1.to_double()))).to_double().sqrt().to_float()
    let theta = 2.0F * 3.14159265F * u2
    let z0 : Float = r * Float::from_double(@math.cos(theta.to_double()))
    out[k] = z0
    if k + 1 < n {
      // z1 = sqrt(-2 ln u1) sin(2π u2)
      let z1 : Float = r * Float::from_double(@math.sin(theta.to_double()))
      out[k + 1] = z1
    }
    k = k + 2
  }
  out
}

///|
/// Forward pass with sampled noise. `rng` produces two N(0,1) sets
/// (ε_i of size in_features, ε_j of size out_features) per call;
/// the rest is deterministic.
///
/// `input` : length = n * in_features
/// returns: length = n * out_features
pub fn noisy_linear_forward(
  input : Array[Float],
  n : Int,
  param : NoisyLinearParam,
  rng : Xoshiro,
) -> Array[Float] {
  let in_f = param.in_features
  let out_f = param.out_features
  // Sample ε_i (in_f) and ε_j (out_f).
  let eps_i_raw = sample_gaussians(in_f, rng)
  let eps_j_raw = sample_gaussians(out_f, rng)
  // Apply f(x) = sign(x)·sqrt(|x|).
  let eps_i : Array[Float] = Array::make(in_f, 0.0F)
  let eps_j : Array[Float] = Array::make(out_f, 0.0F)
  for k in 0.. Array[Float] {
  let in_f = param.in_features
  let out_f = param.out_features
  let out : Array[Float] = Array::make(n * out_f, 0.0F)
  for batch in 0.. Int {
  2 * (param.weight_mu.length() + param.bias_mu.length())
}