From 62650aa53b9f695db08bac634b44197f3fa2ff38 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Mon, 20 Apr 2026 14:41:16 +0200 Subject: [PATCH] =?UTF-8?q?feat:=20add=20state=5Flayout.cuh=20=E2=80=94=20?= =?UTF-8?q?CUDA=20header=20with=20layout=20constants=20and=20assemble=5Fst?= =?UTF-8?q?ate()?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-Authored-By: Claude Opus 4.6 (1M context) --- crates/ml/src/cuda_pipeline/state_layout.cuh | 71 ++++++++++++++++++++ 1 file changed, 71 insertions(+) create mode 100644 crates/ml/src/cuda_pipeline/state_layout.cuh diff --git a/crates/ml/src/cuda_pipeline/state_layout.cuh b/crates/ml/src/cuda_pipeline/state_layout.cuh new file mode 100644 index 000000000..2d98868bb --- /dev/null +++ b/crates/ml/src/cuda_pipeline/state_layout.cuh @@ -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; +}