feat: add state_layout.cuh — CUDA header with layout constants and assemble_state()

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-04-20 14:41:16 +02:00
parent c60cd98a35
commit 62650aa53b

View File

@@ -0,0 +1,71 @@
// state_layout.cuh — single source of truth for GPU state vector layout.
// Included by experience_kernels.cu and any kernel that assembles states.
// Must match crates/ml-core/src/state_layout.rs exactly.
#pragma once
// ── Dimensions ──
#define SL_STATE_DIM 96
#define SL_MARKET_DIM 42
#define SL_OFI_DIM 20
#define SL_MTF_DIM 16
#define SL_PORTFOLIO_BASE_DIM 8
#define SL_PORTFOLIO_PLAN_DIM 6
#define SL_PADDING_DIM 4
// ── Offsets ──
#define SL_MARKET_START 0
#define SL_OFI_START (SL_MARKET_START + SL_MARKET_DIM)
#define SL_MTF_START (SL_OFI_START + SL_OFI_DIM)
#define SL_PORTFOLIO_START (SL_MTF_START + SL_MTF_DIM)
#define SL_PLAN_ISV_START (SL_PORTFOLIO_START + SL_PORTFOLIO_BASE_DIM)
#define SL_PADDING_START (SL_PLAN_ISV_START + SL_PORTFOLIO_PLAN_DIM)
// ── cuBLAS alignment ──
#define SL_STATE_DIM_PADDED 128
// ── Compile-time checks ──
_Static_assert(SL_PADDING_START + SL_PADDING_DIM == SL_STATE_DIM,
"State layout dimensions must sum to SL_STATE_DIM");
_Static_assert(SL_STATE_DIM % 8 == 0,
"SL_STATE_DIM must be 8-aligned for tensor core cuBLAS");
// ── Shared state assembly function ──
// Called by both experience_state_gather (training) and backtest_state_gather (validation).
// Each caller prepares its own data and calls this to produce identical layouts.
//
// All pointers are per-thread (one state row). Caller is responsible for
// computing the correct source pointers for this thread's episode/window.
__device__ __forceinline__ void assemble_state(
float* __restrict__ out,
const float* __restrict__ market,
const float* __restrict__ ofi,
const float mtf[SL_MTF_DIM],
const float portfolio[SL_PORTFOLIO_BASE_DIM],
const float plan_isv[SL_PORTFOLIO_PLAN_DIM]
) {
// Market features [0..42)
for (int k = 0; k < SL_MARKET_DIM; k++)
out[SL_MARKET_START + k] = market[k];
// OFI [42..62)
for (int k = 0; k < SL_OFI_DIM; k++)
out[SL_OFI_START + k] = ofi[k];
// MTF [62..78)
for (int k = 0; k < SL_MTF_DIM; k++)
out[SL_MTF_START + k] = mtf[k];
// Portfolio base [78..86)
for (int k = 0; k < SL_PORTFOLIO_BASE_DIM; k++)
out[SL_PORTFOLIO_START + k] = portfolio[k];
// Plan/ISV [86..92)
for (int k = 0; k < SL_PORTFOLIO_PLAN_DIM; k++)
out[SL_PLAN_ISV_START + k] = plan_isv[k];
// Padding [92..96) — zero
for (int k = 0; k < SL_PADDING_DIM; k++)
out[SL_PADDING_START + k] = 0.0f;
}