// lstm_qnetwork_continuous.mbt — Recurrent continuous-action Q-network
// for partially-observable environments (v0.62.0).
//
// Architecture: MLP w1 (concat(state, action) -> hidden) -> ReLU ->
// LSTM cell (over time, tracking both hidden state h_t and cell state
// c_t) -> MLP w2 (hidden -> scalar Q).
//
//   [state_t ; action_t]    ∈ R^{state_dim + action_dim}
//   x_proj_t = ReLU(w1 · [state_t ; action_t] + b1)         ∈ R^{hidden}
//   (h_t, c_t) = LSTM_cell(x_proj_t, h_{t-1}, c_{t-1})
//   q_t = w2 · h_t + b2                                     ∈ R
//
// This is the LSTM counterpart of `GRUQNetworkContinuous` (v0.59.0) for
// the DDPG_LSTM agent. The LSTM hidden state mixes past state-action
// pairs so the critic can break the Markov assumption.

///|
/// Recurrent continuous-action Q-network. MLP hidden dim equals LSTM
/// hidden dim so the ReLU-projected [state; action] and the LSTM cell
/// input/output match. `mlp_b2` is a scalar (the critic outputs a
/// single Q-value per timestep).
pub struct LSTMQNetworkContinuous {
  state_dim : Int
  action_dim : Int
  hidden : Int
  mlp_w1 : Array[Array[Float]]
  mlp_b1 : Array[Float]
  lstm : LstmCellParam
  mlp_w2 : Array[Array[Float]]
  mut mlp_b2 : Float
}

///|
/// Build a fresh LSTMQNetworkContinuous. MLP weights use Xavier-normal
/// init scaled by sqrtf(2 / fan_in) (He-style for ReLU). LSTM cell
/// uses its own init. Zero biases.
pub fn LSTMQNetworkContinuous::new(
  state_dim : Int,
  action_dim : Int,
  hidden : Int,
  seed : UInt64,
) -> LSTMQNetworkContinuous {
  let in_dim = state_dim + action_dim
  let rng1 = Xoshiro::from_state(seed, seed + 1UL, seed + 2UL, seed + 3UL)
  let std1 = sqrtf(2.0F / Float::from_int(in_dim))
  let mlp_w1 = xavier_normal(hidden, in_dim, std1, rng1)
  let mlp_b1 : Array[Float] = Array::make(hidden, 0.0F)
  let lstm = LstmCellParam::new(hidden, hidden, seed + 4UL)
  let rng2 = Xoshiro::from_state(seed + 5UL, seed + 6UL, seed + 7UL, seed + 8UL)
  let std2 = sqrtf(2.0F / Float::from_int(hidden))
  let mlp_w2 = xavier_normal(1, hidden, std2, rng2)
  { state_dim, action_dim, hidden, mlp_w1, mlp_b1, lstm, mlp_w2, mlp_b2: 0.0F }
}

///|
/// Single-step forward. Returns the scalar Q-value for the given
/// (state, action) pair mixed with the previous hidden + cell state,
/// plus the new hidden + cell state for the next step.
pub fn lstm_qnetwork_continuous_step(
  qnet : LSTMQNetworkContinuous,
  state : Array[Float],
  action : Array[Float],
  hidden : Array[Float],
  cell : Array[Float],
) -> (Float, Array[Float], Array[Float]) {
  // sa = concat(state, action)
  let sa = vec_concat(state, action)
  // x_proj = ReLU(w1 · sa + b1)
  let x_proj_pre = matvec(qnet.mlp_w1, qnet.mlp_b1, sa)
  let x_proj = relu_forward(x_proj_pre)
  // (hidden_next, cell_next) = LSTM_cell(x_proj, hidden, cell)
  let (hidden_next, cache) = lstm_cell_forward(x_proj, hidden, cell, qnet.lstm)
  let cell_next = cache.c_t
  // q = w2 · hidden_next + b2
  let q_pre = matvec(qnet.mlp_w2, [qnet.mlp_b2], hidden_next)
  (q_pre[0], hidden_next, cell_next)
}

///|
/// Sequence forward. `state_seq` is flat row-major `[seq_len × state_dim]`,
/// `action_seq` is flat row-major `[seq_len × action_dim]`. Returns
/// `(q_seq, final_hidden, final_cell)` where `q_seq` is flat `[seq_len]`.
pub fn lstm_qnetwork_continuous_seq_forward(
  qnet : LSTMQNetworkContinuous,
  state_seq : Array[Float],
  action_seq : Array[Float],
  seq_len : Int,
  hidden_init : Array[Float],
  cell_init : Array[Float],
) -> (Array[Float], Array[Float], Array[Float]) {
  let q_seq : Array[Float] = Array::make(seq_len, 0.0F)
  let mut hidden = hidden_init
  let mut cell = cell_init
  for t in 0..