// reinforce.mbt — REINFORCE policy gradient on a small GridWorld (v0.33.0).
//
// Vanilla policy gradient (Monte Carlo):
//   ∇_θ J(θ) = E_τ [ Σ_t ∇_θ log π_θ(a_t | s_t) · G_t ]
// where G_t = Σ_{k >= t} γ^(k-t) · r_k is the discounted return.
//
// Components:
//   - `GridWorld`        — small deterministic MDP (n×n, 4 actions)
//   - `LinearSoftmaxPolicy` — linear logits W · x, softmax, no hidden layer
//   - `Episode`          — list of (state, action, reward) tuples
//   - `policy_rollout`   — sample one episode under the current policy
//   - `compute_returns`  — Monte-Carlo discounted returns
//   - `policy_gradient_update` — REINFORCE gradient step
//   - `train_reinforce`  — K-episode training loop

///|
/// Action encoding: 0=up, 1=down, 2=left, 3=right.
pub let rl_action_up : Int = 0

///|
pub let rl_action_down : Int = 1

///|
pub let rl_action_left : Int = 2

///|
pub let rl_action_right : Int = 3

///|
/// Deterministic 2D GridWorld. `n_states = n_rows * n_cols`.
/// State is `(row, col)` flattened to `row * n_cols + col`.
pub struct GridWorld {
  n_rows : Int
  n_cols : Int
  start : Int
  goal : Int
  step_penalty : Float
  goal_reward : Float
  max_steps : Int
}

///|
/// Build a simple GridWorld with start at (0, 0), goal at
/// (n_rows-1, n_cols-1), step penalty -0.1, goal reward +1.0.
pub fn GridWorld::new(n_rows : Int, n_cols : Int, max_steps : Int) -> GridWorld {
  {
    n_rows,
    n_cols,
    start: 0,
    goal: (n_rows - 1) * n_cols + (n_cols - 1),
    step_penalty: -0.1F,
    goal_reward: 1.0F,
    max_steps,
  }
}

///|
pub fn GridWorld::n_states(self : GridWorld) -> Int {
  self.n_rows * self.n_cols
}

///|
/// One-hot encode a state index into a length `n` vector.
pub fn rl_one_hot(state : Int, n : Int) -> Array[Float] {
  let x : Array[Float] = Array::make(n, 0.0F)
  if state >= 0 && state < n {
    x[state] = 1.0F
  }
  x
}

///|
/// Apply action. Returns (next_state, reward, done).
/// Going off the grid keeps the agent in place and adds step_penalty.
pub fn GridWorld::step(self : GridWorld, state : Int, action : Int) -> (Int, Float, Bool) {
  let row = state / self.n_cols
  let col = state % self.n_cols
  let mut new_row = row
  let mut new_col = col
  if action == rl_action_up {
    if row > 0 {
      new_row = row - 1
    }
  } else if action == rl_action_down {
    if row < self.n_rows - 1 {
      new_row = row + 1
    }
  } else if action == rl_action_left {
    if col > 0 {
      new_col = col - 1
    }
  } else if action == rl_action_right {
    if col < self.n_cols - 1 {
      new_col = col + 1
    }
  }
  let next_state = new_row * self.n_cols + new_col
  if next_state == self.goal {
    (next_state, self.goal_reward, true)
  } else {
    (next_state, self.step_penalty, false)
  }
}

///|
/// Linear softmax policy: logits = W · x, π = softmax(logits).
pub struct LinearSoftmaxPolicy {
  n_states : Int
  n_actions : Int
  w : Array[Array[Float]]
}

///|
pub fn LinearSoftmaxPolicy::new(
  n_states : Int,
  n_actions : Int,
  seed : UInt64,
) -> LinearSoftmaxPolicy {
  let std = sqrtf(1.0F / Float::from_int(n_states))
  let rng = Xoshiro::from_state(seed, seed + 1UL, seed + 2UL, seed + 3UL)
  let w = xavier_normal(n_actions, n_states, std, rng)
  { n_states, n_actions, w }
}

///|
/// Compute logits and numerically stable softmax probabilities.
pub fn policy_forward(
  policy : LinearSoftmaxPolicy,
  x : Array[Float],
) -> (Array[Float], Array[Float]) {
  let n_actions = policy.n_actions
  let logits : Array[Float] = Array::make(n_actions, 0.0F)
  for i in 0.. max_l {
      max_l = logits[i]
    }
  }
  let probs : Array[Float] = Array::make(n_actions, 0.0F)
  let mut sum_e = 0.0F
  for i in 0.. (Int, Float) {
  let n = probs.length()
  let (u_raw, _) = box_muller(rng)
  let u = Float::from_double(u_raw).abs()
  let mut thr = u
  if thr > 1.0F {
    thr = 1.0F - 1.0e-7F
  } else if thr < 0.0F {
    thr = 0.0F
  }
  let mut cum = 0.0F
  let mut a = n - 1
  for i in 0..= thr {
      a = i
      break
    }
  }
  let lp = if probs[a] > 0.0F {
    logf(probs[a])
  } else {
    -20.0F
  }
  (a, lp)
}

///|
/// Episode record: states, actions, rewards.
pub(all) struct Episode {
  states : Array[Int]
  actions : Array[Int]
  rewards : Array[Float]
}

///|
/// Roll out one episode under the current policy.
pub fn policy_rollout(
  env : GridWorld,
  policy : LinearSoftmaxPolicy,
  max_steps : Int,
  rng : Xoshiro,
) -> Episode {
  let states : Array[Int] = []
  let actions : Array[Int] = []
  let rewards : Array[Float] = []
  let mut state = env.start
  let mut done = false
  let mut t = 0
  while !done && t < max_steps {
    let x = rl_one_hot(state, env.n_states())
    let (_logits, probs) = policy_forward(policy, x)
    let (a, _lp) = sample_categorical(probs, rng)
    let (next_state, r, d) = env.step(state, a)
    states.push(state)
    actions.push(a)
    rewards.push(r)
    state = next_state
    done = d
    t = t + 1
  }
  { states, actions, rewards }
}

///|
/// Monte-Carlo discounted returns G_t = Σ_{k >= t} γ^(k-t) · r_k.
pub fn compute_returns(rewards : Array[Float], gamma : Float) -> Array[Float] {
  let t = rewards.length()
  let g : Array[Float] = Array::make(t, 0.0F)
  let mut running = 0.0F
  let mut i = t - 1
  while i >= 0 {
    running = rewards[i] + gamma * running
    g[i] = running
    i = i - 1
  }
  g
}

///|
/// Total undiscounted return of an episode.
pub fn episode_return(episode : Episode) -> Float {
  let mut s = 0.0F
  for r in episode.rewards {
    s = s + r
  }
  s
}

///|
/// REINFORCE policy gradient step. Modifies `policy.w` in place.
/// d_logits[i] = (1{i == a_t} - π(a|s_t)) · G_t
/// d_W[i, j] += d_logits[i] · x_t[j]
pub fn policy_gradient_update(
  policy : LinearSoftmaxPolicy,
  episode : Episode,
  returns : Array[Float],
  lr : Float,
) -> Unit {
  let t = episode.states.length()
  let n_a = policy.n_actions
  let n_s = policy.n_states
  for step in 0.. Float {
  let rng = Xoshiro::from_state(seed, seed + 7UL, seed + 13UL, seed + 17UL)
  let mut total_return = 0.0F
  for _ep in 0..