// dt_embeddings.mbt — Decision Transformer primitive: tokenization of
// (rtg, state, action, timestep) trajectories into d_model-dimensional
// tokens (v0.76.0).
//
// Reference: Chen et al. 2021 "Decision Transformer: Reinforcement
// Learning through Sequence Modeling". The original DT uses GPT-style
// causal attention; here we substitute a GTrXL block (Parisotto et al.
// 2020) as the per-token recurrent memory primitive. This keeps DT in
// the same architecture family as the v0.72–v0.75 GTrXL primitives.
//
// Per-timestep token layout:
// token_t = Linear_rtg(R_t) + Linear_state(s_t) + Linear_action(a_t) + Embed_timestep(t)
//
// All three modality projections (rtg_w, state_w, action_w) are linear
// without bias. The timestep embedding is a small lookup table
// (max_timestep slots) of size d_model each. The sum of the four terms
// becomes the input to the GTrXL block.
//
// Scope of v0.76.0:
// - DTEmbeddings struct + constructor
// - dt_embed_trajectory: stitch (rtg_seq, state_seq, action_seq,
// timesteps) into a flat [seq_len × d_model] token sequence
// - dt_embed_single_token: helper for inference (one token at a time,
// no batch dim)
// - Returns-to-go are NOT computed here — that's TrajectoryBuffer's
// job (v0.78.0).
///|
/// Decision Transformer token embeddings. Three modality projections
/// (rtg, state, action) + a timestep embedding lookup. d_model is
/// shared across all four (each adds into the same d_model vec).
pub struct DTEmbeddings {
state_dim : Int
action_dim : Int
d_model : Int
max_timestep : Int
rtg_w : Array[Array[Float]]
state_w : Array[Array[Float]]
action_w : Array[Array[Float]]
timestep_emb : Array[Array[Float]]
}
///|
/// Build fresh DTEmbeddings.
/// - rtg_w: (d_model × 1), each row scaled by sqrtf(2 / 1) = sqrt(2)
/// - state_w: (d_model × state_dim), Xavier-normal scaled by sqrt(2 / state_dim)
/// - action_w: (d_model × action_dim), Xavier-normal scaled by sqrt(2 / action_dim)
/// - timestep_emb: (max_timestep × d_model), Xavier-normal scaled by 0.02 (small,
/// like original Transformer init for positional embeddings)
pub fn DTEmbeddings::new(
state_dim : Int,
action_dim : Int,
d_model : Int,
max_timestep : Int,
seed : UInt64,
) -> DTEmbeddings {
let rng1 = Xoshiro::from_state(seed, seed + 1UL, seed + 2UL, seed + 3UL)
let std_rtg = sqrtf(2.0F / 1.0F)
let rtg_w = xavier_normal(d_model, 1, std_rtg, rng1)
let rng2 = Xoshiro::from_state(seed + 4UL, seed + 5UL, seed + 6UL, seed + 7UL)
let std_state = sqrtf(2.0F / Float::from_int(state_dim))
let state_w = xavier_normal(d_model, state_dim, std_state, rng2)
let rng3 = Xoshiro::from_state(seed + 8UL, seed + 9UL, seed + 10UL, seed + 11UL)
let std_action = sqrtf(2.0F / Float::from_int(action_dim))
let action_w = xavier_normal(d_model, action_dim, std_action, rng3)
let rng4 = Xoshiro::from_state(seed + 12UL, seed + 13UL, seed + 14UL, seed + 15UL)
let std_ts = 0.02F
let timestep_emb = xavier_normal(max_timestep, d_model, std_ts, rng4)
// xavier_normal returns (rows, cols). We want (max_timestep × d_model).
// If std_ts * sqrt(2 / fan_in) where fan_in = d_model is too small,
// fall back to a hand-rolled init that produces max_timestep × d_model.
// xavier_normal returns (rows × cols); we passed (max_timestep, d_model, std_ts, rng4),
// so the result is already max_timestep rows × d_model cols. Good.
{
state_dim,
action_dim,
d_model,
max_timestep,
rtg_w,
state_w,
action_w,
timestep_emb,
}
}
///|
/// Helper: project a single scalar `x` (treated as 1-vec) through `w`
/// (shape d_model × 1) with optional bias (length 1). Returns
/// Array[Float] of length d_model.
fn dt_project_scalar(
w : Array[Array[Float]],
x : Float,
) -> Array[Float] {
let out : Array[Float] = Array::make(w.length(), 0.0F)
for i in 0.. Array[Float] {
let out : Array[Float] = Array::make(w.length(), 0.0F)
for i in 0.. Array[Float] {
let rtg_proj = dt_project_scalar(emb.rtg_w, rtg_t)
let state_proj = dt_project_vector(emb.state_w, state_t)
let action_proj = dt_project_vector(emb.action_w, action_t)
let token : Array[Float] = Array::make(emb.d_model, 0.0F)
for k in 0.. Array[Float] {
let tokens : Array[Float] = Array::make(seq_len * emb.d_model, 0.0F)
for t in 0..