fix(sp4): Task A7 fix-up #2 — Curiosity sub-buffer support, all 8 groups wired

Curiosity stores params/grads/Adam state as 4 non-contiguous sub-buffers
(w1, b1, w2, b2 — each its own CudaSlice<f32>). The single-pointer kernel
signature couldn't describe them. The first fix-up wired 4 of 5 aux
groups; Curiosity (group 7) was deferred.

This commit extends the param_group_oracle kernel to accept a sub-buffer
table (mapped-pinned u64 ptr-arrays + i32 counts), with `n_sub=1` for
groups 0-6 (existing behavior) and `n_sub=4` for Curiosity. Pass A/B/C
(p99 histograms) iterate sub-buffers via a new
`sp4_histogram_p99_multi<BLOCK_SIZE>` template that mirrors the original
three-pass structure but loops sub-buffers within the max-reduce + binning
passes; Pass 3 (cumulative-from-top) divides by `total_count`. Pass D
(4-way reduce) iterates sub-buffers in the accumulator loop; Pass E
(L1 trunk only) reads `grads_ptrs[0]` (group 0 has n_sub=1 by construction).

`Sp4ParamGroupBufs` redefined from a flat quartet to `Vec<Sp4SubBuffer>`.
Convenience constructors `Sp4ParamGroupBufs::single(...)` and
`Sp4ParamGroupBufs::empty()` keep call sites ergonomic; `total_count()`
for kernel arg. Switched Copy → Clone since the descriptor now owns a Vec.
`SP4AuxBuffers` extended with `curiosity` field (5-tuple from 4-tuple).

Added 16 public accessors to GpuCuriosityTrainer (grad + Adam state) and
8 to CuriosityWeightSet (weights + lengths) for the 4 sub-buffers ×
4 signal types.

`oracle_subbuf_table_buf: MappedU64Buffer` (4×4 = 16 u64s) and
`oracle_subbuf_counts_buf: MappedI32Buffer` (4 i32s) allocated at
construction. Launcher overwrites entries [0..n_sub) per group launch
via volatile writes; zeros the unused tail (defence-in-depth so a kernel
bug reading past `n_sub` lands on count=0 no-op). Inter-launch
`stream.synchronize()` added so the next iteration's host writes don't
race with the in-flight kernel's coalesced loads from the persistent
table. Cold-path producer; per-launch sync cost is negligible vs the
kernel work.

`build_sp4_aux_buffers` signature changed: takes
`Option<&CuriosityWeightSet>` and `Option<&GpuCuriosityTrainer>` since
Curiosity state lives outside FusedTrainingCtx (owned by
GpuExperienceCollector). Layer B's training-loop caller threads them
in from `collector.curiosity_weight_set()` + the collector's
`curiosity_trainer` field; both must be Some together (caller
responsibility — they live on the same collector). Passing None for
either yields an empty curiosity descriptor and the launcher silently
skips group 7.

`param_group_buffers` return type changed from
`Option<(u64, u64, u64, u64, usize, i32, i32)>` to
`Option<(Sp4ParamGroupBufs, i32, i32)>`. All groups now return Some(...)
(Curiosity included); None reserved for forward-compat.

GPU test extended: group 7 exercises 4 sub-buffers of distinct shapes
[1024, 32, 1024, 32] (w1/b1/w2/b2-like sizes scaled to keep test runtime
small while still exercising the multi-sub-buffer iteration), each with
its own seed offset so distinct sub-buffers have distinct distributions
— catches buffer-mixup bugs in the kernel's sub-buffer iteration.
Reference computation builds union vectors and computes p99/WD_RATE
over the union. Test launcher refactored to `&[TestSubBuffer]` slice
matching the production kernel's table-packing layout.

All 8 SP4 param-groups now produce real outputs in Layer A. The
launcher's `count == 0` short-circuit retained for the optional aux-
trainer fallback (init failures on gpu_iqn / gpu_attention / curiosity).

`MappedU64Buffer` gained a manual Debug impl (warn missing_debug_impls)
for parity with MappedU32Buffer.

cargo check -p ml --lib --tests clean. Workspace clean. Layer B can now
safely consume all 8 ISV[WEIGHT_BOUND/ADAM_M_BOUND/ADAM_V_BOUND/WD_RATE[group]]
slots.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-05-01 00:39:32 +02:00
parent 64298b34c0
commit 8f93c11525
9 changed files with 730 additions and 256 deletions

View File

@@ -934,4 +934,45 @@ impl GpuCuriosityTrainer {
pub fn step_count(&self) -> i32 {
self.step
}
// ── SP4 Task A7 fix-up #2: per-sub-buffer device-pointer + length ─────
// accessors for the param-group oracle. Curiosity stores its 4 grad +
// Adam state slices (matching `[w1, b1, w2, b2]`) as separate
// `CudaSlice<f32>`s; the oracle kernel iterates them via the multi-
// sub-buffer table.
/// Raw device pointer to `grad_w1 [CUR_W1_LEN]`.
pub fn grad_w1_ptr(&self) -> u64 { self.grad_w1.raw_ptr() }
/// Element count of `grad_w1`.
pub fn grad_w1_len(&self) -> usize { self.grad_w1.len() }
/// Raw device pointer to `grad_b1 [CUR_B1_LEN]`.
pub fn grad_b1_ptr(&self) -> u64 { self.grad_b1.raw_ptr() }
/// Element count of `grad_b1`.
pub fn grad_b1_len(&self) -> usize { self.grad_b1.len() }
/// Raw device pointer to `grad_w2 [CUR_W2_LEN]`.
pub fn grad_w2_ptr(&self) -> u64 { self.grad_w2.raw_ptr() }
/// Element count of `grad_w2`.
pub fn grad_w2_len(&self) -> usize { self.grad_w2.len() }
/// Raw device pointer to `grad_b2 [CUR_B2_LEN]`.
pub fn grad_b2_ptr(&self) -> u64 { self.grad_b2.raw_ptr() }
/// Element count of `grad_b2`.
pub fn grad_b2_len(&self) -> usize { self.grad_b2.len() }
/// Raw device pointer to `adam_m_w1 [CUR_W1_LEN]` (first moment).
pub fn adam_m_w1_ptr(&self) -> u64 { self.adam_m_w1.raw_ptr() }
/// Raw device pointer to `adam_m_b1 [CUR_B1_LEN]`.
pub fn adam_m_b1_ptr(&self) -> u64 { self.adam_m_b1.raw_ptr() }
/// Raw device pointer to `adam_m_w2 [CUR_W2_LEN]`.
pub fn adam_m_w2_ptr(&self) -> u64 { self.adam_m_w2.raw_ptr() }
/// Raw device pointer to `adam_m_b2 [CUR_B2_LEN]`.
pub fn adam_m_b2_ptr(&self) -> u64 { self.adam_m_b2.raw_ptr() }
/// Raw device pointer to `adam_v_w1 [CUR_W1_LEN]` (second moment).
pub fn adam_v_w1_ptr(&self) -> u64 { self.adam_v_w1.raw_ptr() }
/// Raw device pointer to `adam_v_b1 [CUR_B1_LEN]`.
pub fn adam_v_b1_ptr(&self) -> u64 { self.adam_v_b1.raw_ptr() }
/// Raw device pointer to `adam_v_w2 [CUR_W2_LEN]`.
pub fn adam_v_w2_ptr(&self) -> u64 { self.adam_v_w2.raw_ptr() }
/// Raw device pointer to `adam_v_b2 [CUR_B2_LEN]`.
pub fn adam_v_b2_ptr(&self) -> u64 { self.adam_v_b2.raw_ptr() }
}

View File

@@ -1330,30 +1330,81 @@ impl Drop for EventTrackingGuard<'_> {
// represent contiguously (Curiosity is the architectural hold-out — see
// `param_group_buffers` doc-comment).
/// SP4 Task A7 fix-up: device-pointer + length quartet for one aux
/// trainer's `(params, grads, adam_m, adam_v)` buffers. All four buffers
/// must be the same length (Adam state mirrors the params shape;
/// `count` records that single shared element count). Built host-side
/// from the trainer's `params_ptr()` / `grads_ptr()` / `adam_m_ptr()` /
/// `adam_v_ptr()` accessors and consumed by `param_group_buffers`.
/// SP4 Task A7 fix-up #2: one sub-buffer's `(params, grads, adam_m, adam_v)`
/// device pointers + element count. All four buffers must be the same length
/// (Adam state mirrors the params shape; `count` records that single shared
/// element count). For groups with a single contiguous params buffer (groups
/// 0-6) a `Sp4ParamGroupBufs` holds exactly one of these. For Curiosity
/// (group 7), four — one per `[w1, b1, w2, b2]` sub-tensor.
#[derive(Copy, Clone, Debug)]
pub struct Sp4ParamGroupBufs {
/// Raw device pointer to the contiguous flat-f32 params buffer.
pub struct Sp4SubBuffer {
/// Raw device pointer to this sub-buffer's flat-f32 params slice.
pub params_ptr: u64,
/// Raw device pointer to the reduced-gradient buffer (same shape).
/// Raw device pointer to the matched-shape reduced-gradient slice.
pub grads_ptr: u64,
/// Raw device pointer to the Adam first-moment buffer.
/// Raw device pointer to the Adam first-moment slice.
pub adam_m_ptr: u64,
/// Raw device pointer to the Adam second-moment buffer.
/// Raw device pointer to the Adam second-moment slice.
pub adam_v_ptr: u64,
/// Element count (f32 entries) shared by all four buffers.
/// Element count (f32 entries) shared by all four slices.
pub count: usize,
}
/// SP4 Task A7 fix-up: aux-trainer buffer pointers for the param-group
/// SP4 Task A7 fix-up: per-param-group buffer descriptor for the Pearl B
/// oracle. Most groups (0-6) have a single contiguous flat-f32 params buffer
/// per signal type — `sub_buffers.len() == 1`. Curiosity (group 7) splits
/// its weights into 4 sub-buffers (`w1`, `b1`, `w2`, `b2`) each in its own
/// `CudaSlice<f32>`; `sub_buffers.len() == 4`. The oracle kernel iterates
/// sub-buffers within each pass, treating their union as a single logical
/// distribution for p99/WD_RATE/L1 computation.
///
/// Built host-side from each trainer's `params_ptr()` / `grads_ptr()` /
/// `adam_m_ptr()` / `adam_v_ptr()` accessors (or per-sub-buffer accessors
/// for Curiosity) and consumed by `param_group_buffers` /
/// `launch_sp4_param_group_oracles_all_groups`.
#[derive(Clone, Debug)]
pub struct Sp4ParamGroupBufs {
/// Sub-buffers in matched-stride order: index `i` corresponds to the
/// same logical sub-tensor across (params, grads, adam_m, adam_v).
/// All four pointer arrays MUST be the same length (= `sub_buffers.len()`)
/// and have matching `count` per index.
pub sub_buffers: Vec<Sp4SubBuffer>,
}
impl Sp4ParamGroupBufs {
/// Convenience constructor for groups with a single contiguous buffer
/// per signal type (groups 0-6).
pub fn single(
params_ptr: u64,
grads_ptr: u64,
adam_m_ptr: u64,
adam_v_ptr: u64,
count: usize,
) -> Self {
Sp4ParamGroupBufs {
sub_buffers: vec![Sp4SubBuffer { params_ptr, grads_ptr, adam_m_ptr, adam_v_ptr, count }],
}
}
/// Empty descriptor (no sub-buffers). Trips the launcher's `total_count==0`
/// short-circuit so the kernel launch is silently skipped — used for
/// optional aux trainers that aren't present in the current configuration.
pub fn empty() -> Self {
Sp4ParamGroupBufs { sub_buffers: Vec::new() }
}
/// Total element count across all sub-buffers (for kernel `total_count`
/// arg + p99 normalisation).
pub fn total_count(&self) -> usize {
self.sub_buffers.iter().map(|s| s.count).sum()
}
}
/// SP4 Task A7 fix-up #2: aux-trainer buffer pointers for the param-group
/// statistics oracle. `FusedTrainingCtx` builds this from its child
/// trainers (`gpu_iqn`, `gpu_iql`, `gpu_iql_low`, `gpu_attention`) and
/// passes it to `launch_sp4_param_group_oracles_all_groups`.
/// trainers (`gpu_iqn`, `gpu_iql`, `gpu_iql_low`, `gpu_attention`,
/// `gpu_curiosity`) and passes it to
/// `launch_sp4_param_group_oracles_all_groups`.
///
/// The mapping `aux_buffers.{field} → ParamGroup::{variant}` is fixed by
/// `param_group_buffers`:
@@ -1361,9 +1412,9 @@ pub struct Sp4ParamGroupBufs {
/// - `iql_high` → group 4 (`ParamGroup::IqlHigh`)
/// - `iql_low` → group 5 (`ParamGroup::IqlLow`)
/// - `attn` → group 6 (`ParamGroup::Attn`)
/// - (group 7 `ParamGroup::Curiosity` is omitted — non-contiguous
/// `[w1, b1, w2, b2]` layout cannot be described by a single
/// `Sp4ParamGroupBufs`; see `param_group_buffers` doc-comment.)
/// - `curiosity` → group 7 (`ParamGroup::Curiosity`) — multi-sub-buffer
/// descriptor (`Sp4ParamGroupBufs::sub_buffers.len() == 4` for
/// `[w1, b1, w2, b2]`); the kernel iterates them as one logical group.
///
/// All five aux trainers are *unconditional* in `FusedTrainingCtx` —
/// IQN/Attn used to be `Option<_>` for graceful-degrade reasons that
@@ -1371,7 +1422,7 @@ pub struct Sp4ParamGroupBufs {
/// optional aux trainers, the launcher would need a per-group "present"
/// flag; for now `aux_buffers.iqn`, `aux_buffers.iql_high`, etc. are
/// always live device pointers.
#[derive(Copy, Clone, Debug)]
#[derive(Clone, Debug)]
pub struct SP4AuxBuffers {
/// IQN online-network params + grads + Adam state.
pub iqn: Sp4ParamGroupBufs,
@@ -1381,6 +1432,10 @@ pub struct SP4AuxBuffers {
pub iql_low: Sp4ParamGroupBufs,
/// Multi-head attention params + grads + Adam state.
pub attn: Sp4ParamGroupBufs,
/// Curiosity forward-model params + grads + Adam state. Multi-sub-buffer
/// descriptor: 4 entries for `[w1, b1, w2, b2]`. The oracle kernel
/// treats the union as one logical group.
pub curiosity: Sp4ParamGroupBufs,
}
@@ -2860,6 +2915,22 @@ pub struct GpuDqnTrainer {
/// host reads via mapped-pinned device-ptr.
pub(crate) producer_step_scratch_buf: MappedF32Buffer,
/// SP4 Task A7 fix-up #2: oracle sub-buffer device-pointer table.
/// Mapped-pinned u64 buffer holding 4 contiguous K-wide ptr arrays
/// (params, grads, adam_m, adam_v) where `K = SP4_ORACLE_TABLE_MAX_SUB`.
/// Total: 4×K u64s. The kernel reads each ptr array starting at its
/// own device-pointer offset (`oracle_subbuf_table_buf.dev_ptr +
/// k_idx * K * 8`). Host populates entries [0..n_sub) once per group
/// launch and the kernel iterates them via `n_sub` arg.
pub(crate) oracle_subbuf_table_buf: MappedU64Buffer,
/// SP4 Task A7 fix-up #2: oracle sub-buffer per-entry counts table.
/// Companion to `oracle_subbuf_table_buf`. Mapped-pinned i32 buffer
/// of `K = SP4_ORACLE_TABLE_MAX_SUB` entries. Host populates [0..n_sub)
/// per group launch; the kernel reads via the device pointer as
/// `const int*`.
pub(crate) oracle_subbuf_counts_buf: MappedI32Buffer,
/// SP2: fused NaN-check buffer pointer table. 12 device pointers (u64) for
/// slots 24-35. Host+device-visible mapped-pinned buffer; populated via
/// host-side write (no HtoD copy) per
@@ -9013,19 +9084,15 @@ impl GpuDqnTrainer {
/// AdamW weight_decay and L1 unchanged. SP4 consumer migration follows
/// once all producers (A5-A9) land.
///
/// **Aux trainers wired via `SP4AuxBuffers`** (Task A7 fix-up) —
/// `param_group_buffers` consults the supplied `aux_buffers` for groups
/// 3-6 (IQN, IQL-hi, IQL-lo, Attn) and the existing main-DQN param
/// slicing for groups 0-2. Group 7 (Curiosity) is the lone hold-out:
/// `GpuCuriosityTrainer` stores its params/grads/Adam state as four
/// separate `[w1, b1, w2, b2]` sub-buffers (non-contiguous), so a
/// single `(params_ptr, count)` tuple cannot describe the slice the
/// kernel reads. Reconciling that requires either a per-layer launch
/// loop (4× the kernel cost) or relaying out the trainer to a flat
/// params buffer; both are scoped beyond Task A7's fix-up. Until
/// then, group 7 still returns `None` and the launcher silently skips
/// it — Layer B must therefore guard against ISV[170] being the
/// natural-zero floor for Curiosity-related clamps.
/// **Aux trainers wired via `SP4AuxBuffers`** (Task A7 fix-up + fix-up #2)
/// `param_group_buffers` consults the supplied `aux_buffers` for groups
/// 3-7 (IQN, IQL-hi, IQL-lo, Attn, Curiosity) and the existing main-DQN
/// param slicing for groups 0-2. Curiosity is described as a multi-sub-
/// buffer descriptor (`Sp4ParamGroupBufs::sub_buffers.len() == 4` for
/// `[w1, b1, w2, b2]`); the kernel iterates all sub-buffers per pass
/// and treats the union as a single logical group for p99/WD_RATE.
/// Pass E (L1 trunk lambda) is dispatched only for group 0 which is
/// always single-sub-buffer.
pub fn launch_sp4_param_group_oracles_all_groups(
&self,
aux_buffers: &SP4AuxBuffers,
@@ -9049,6 +9116,10 @@ impl GpuDqnTrainer {
const SCRATCH_BASE_WD: usize = 29;
const SCRATCH_L1_TRUNK: usize = 37;
// Sub-buffer table layout: 4 ptr-arrays of length K each.
// Must match `SP4_ORACLE_TABLE_MAX_SUB` in the constructor.
const K_MAX: usize = 4;
// Same shared-memory budget as the rest of the SP4 producer family
// (sp4_histogram_p99 contract): 8 warps × 256 bins × sizeof(int).
// Pass D's 4 sequential block-reduces reuse this same dynamic shmem
@@ -9064,31 +9135,76 @@ impl GpuDqnTrainer {
let scratch_dev = self.producer_step_scratch_buf.dev_ptr;
// Track which groups got a real launch — Phase 2 only post-processes
// those (skipping aux groups whose buffers aren't yet wired).
// those (skipping any group whose buffers aren't present).
let mut launched: [bool; SP4_PARAM_GROUP_COUNT] = [false; SP4_PARAM_GROUP_COUNT];
// Device-pointer offsets into the persistent oracle sub-buffer table.
// `oracle_subbuf_table_buf` is `4 × K_MAX` u64s (= 32 u64s = 128 B).
// Each ptr-array slice starts `K_MAX × 8` bytes apart.
let table_dev = self.oracle_subbuf_table_buf.dev_ptr;
let stride_bytes = (K_MAX * std::mem::size_of::<u64>()) as u64;
let params_ptrs_dev = table_dev;
let grads_ptrs_dev = table_dev + stride_bytes;
let adam_m_ptrs_dev = table_dev + 2 * stride_bytes;
let adam_v_ptrs_dev = table_dev + 3 * stride_bytes;
let counts_dev = self.oracle_subbuf_counts_buf.dev_ptr;
// ── Phase 1: launch one kernel per group (when buffers available) ──
for g_idx in 0..SP4_PARAM_GROUP_COUNT {
let g = ParamGroup::ALL[g_idx];
let bufs = match self.param_group_buffers(g, aux_buffers) {
let (bufs, k_in, h_dim) = match self.param_group_buffers(g, aux_buffers) {
Some(b) => b,
None => {
// Group 7 (Curiosity) only — its non-contiguous
// `[w1, b1, w2, b2]` sub-buffer layout can't be
// described by a single `(params_ptr, count)` tuple.
// See `param_group_buffers` for the architectural
// hold-out details.
continue;
}
None => continue,
};
let (params_ptr, grads_ptr, m_ptr, v_ptr, count, k_in, h_dim) = bufs;
// Sanity: a zero-element slice would yield a degenerate p99 = 0
let n_sub = bufs.sub_buffers.len();
if n_sub == 0 { continue; }
if n_sub > K_MAX {
return Err(MLError::ModelError(format!(
"param_group_oracle[group={g_idx}]: n_sub={n_sub} exceeds \
SP4_ORACLE_TABLE_MAX_SUB={K_MAX} extend the table buffers \
or reduce sub-buffer count",
)));
}
let total_count = bufs.total_count();
// Sanity: a zero-element union would yield degenerate p99 = 0
// inside the kernel and short-circuit Pearls A+D below; cheaper
// to skip the launch entirely.
if count == 0 { continue; }
let count_i32 = count as i32;
let k_in_i32 = k_in;
let h_dim_i32 = h_dim;
if total_count == 0 { continue; }
// Populate the host-visible side of the mapped-pinned tables for
// this group's sub-buffers. Layout matches the device-pointer
// arithmetic above (4 ptr-arrays of length K_MAX, packed).
//
// Safety: `oracle_subbuf_table_buf.host_ptr` is valid for
// `4*K_MAX` u64s; `oracle_subbuf_counts_buf.host_ptr` is valid
// for K_MAX i32s. Volatile writes ensure the kernel-visible
// mapping observes the freshest values once we hit the launch.
unsafe {
let table_host = self.oracle_subbuf_table_buf.host_ptr;
let counts_host = self.oracle_subbuf_counts_buf.host_ptr;
for (s, sb) in bufs.sub_buffers.iter().enumerate() {
std::ptr::write_volatile(table_host.add(0 * K_MAX + s), sb.params_ptr);
std::ptr::write_volatile(table_host.add(1 * K_MAX + s), sb.grads_ptr);
std::ptr::write_volatile(table_host.add(2 * K_MAX + s), sb.adam_m_ptr);
std::ptr::write_volatile(table_host.add(3 * K_MAX + s), sb.adam_v_ptr);
std::ptr::write_volatile(counts_host.add(s), sb.count as i32);
}
// Zero the unused tail entries so a bug that reads past
// `n_sub` lands in known-safe territory (count=0 inner loop
// is a no-op).
for s in n_sub..K_MAX {
std::ptr::write_volatile(table_host.add(0 * K_MAX + s), 0u64);
std::ptr::write_volatile(table_host.add(1 * K_MAX + s), 0u64);
std::ptr::write_volatile(table_host.add(2 * K_MAX + s), 0u64);
std::ptr::write_volatile(table_host.add(3 * K_MAX + s), 0u64);
std::ptr::write_volatile(counts_host.add(s), 0i32);
}
}
let n_sub_i32 = n_sub as i32;
let total_count_i32 = total_count as i32;
let k_in_i32 = k_in;
let h_dim_i32 = h_dim;
let weight_idx_i32 = (SCRATCH_BASE_W + g_idx) as i32;
let adam_m_idx_i32 = (SCRATCH_BASE_M + g_idx) as i32;
let adam_v_idx_i32 = (SCRATCH_BASE_V + g_idx) as i32;
@@ -9099,18 +9215,21 @@ impl GpuDqnTrainer {
-1
};
// Safety: kernel signature `(const float*, const float*,
// const float*, const float*, int, int, int, float*, int, int,
// int, int, int)` matches the 13 args below; all device pointers
// come from `param_group_buffers` which validates ownership.
// Safety: kernel signature `(const u64*, const u64*, const u64*,
// const u64*, const int*, int, int, int, int, float*, int, int,
// int, int, int)` matches the 15 args below; all device pointers
// come from `param_group_buffers` (sub-buffers) and the
// construction-time mapped-pinned tables (ptr arrays + counts).
unsafe {
self.stream
.launch_builder(&self.param_group_oracle_update)
.arg(&params_ptr)
.arg(&grads_ptr)
.arg(&m_ptr)
.arg(&v_ptr)
.arg(&count_i32)
.arg(&params_ptrs_dev)
.arg(&grads_ptrs_dev)
.arg(&adam_m_ptrs_dev)
.arg(&adam_v_ptrs_dev)
.arg(&counts_dev)
.arg(&n_sub_i32)
.arg(&total_count_i32)
.arg(&k_in_i32)
.arg(&h_dim_i32)
.arg(&scratch_dev)
@@ -9124,6 +9243,18 @@ impl GpuDqnTrainer {
"param_group_oracle_update[group={g_idx}] launch: {e}"
)))?;
}
// The kernel reads the table BEFORE the next host overwrite
// because each launch is serialised on the same stream and the
// next iteration's table population happens after kernel
// dispatch. To avoid host stomping the table while the prior
// launch is still reading, sync between launches.
//
// Cold-path producer; the per-launch sync cost is negligible vs
// the kernel work. Without it, the next iteration's host writes
// would race with the in-flight kernel's coalesced loads.
self.stream.synchronize().map_err(|e| MLError::ModelError(format!(
"param_group_oracle_update[group={g_idx}] inter-launch sync: {e}",
)))?;
launched[g_idx] = true;
}
// Single end-of-loop sync — all launches are serialised on the same
@@ -9189,16 +9320,14 @@ impl GpuDqnTrainer {
Ok(())
}
/// SP4 Layer A Task A7 helper: resolve the per-param-group buffer tuple
/// for the Pearl B fused statistics oracle.
/// SP4 Layer A Task A7 helper: resolve the per-param-group buffer
/// descriptor for the Pearl B fused statistics oracle.
///
/// Returns `Some((params_ptr, grads_ptr, adam_m_ptr, adam_v_ptr, count,
/// k_in, h_dim))` when the trainer owns all four buffers for the group
/// **and the trainer's params buffer is contiguous**; returns `None`
/// when the underlying trainer's storage layout is not amenable to a
/// single `(params_ptr, count)` slice — currently only
/// `ParamGroup::Curiosity` (its `[w1, b1, w2, b2]` layout is four
/// separate sub-buffers, not a single flat slice).
/// Returns `Some((bufs, k_in, h_dim))` for every group; the descriptor
/// holds 1 sub-buffer entry for groups 0-6 (single contiguous params
/// buffer) and 4 entries for Curiosity (one per `[w1, b1, w2, b2]`
/// sub-tensor). `None` reserved for forward-compat in case a trainer
/// is unavailable in a future configuration.
///
/// `k_in`/`h_dim` are non-zero only for `ParamGroup::DqnTrunk` (group 0)
/// — the trunk Pass-E entropy-deficit computation needs to know the
@@ -9216,14 +9345,14 @@ impl GpuDqnTrainer {
/// - **Branch heads** (tensors `[17..33)`): W/B_b{0..3}fc +
/// W/B_b{0..3}out concatenated. Same range as `branch_adam_m_ptr`.
///
/// For aux trainers (IQN, IQL-hi, IQL-lo, Attn) the caller supplies
/// pointers via `aux_buffers`; `FusedTrainingCtx` builds that struct
/// from the device-pointer accessors on each child trainer.
/// For aux trainers (IQN, IQL-hi, IQL-lo, Attn, Curiosity) the caller
/// supplies the descriptor via `aux_buffers`; `FusedTrainingCtx` builds
/// that struct from each child trainer's device-pointer accessors.
fn param_group_buffers(
&self,
group: crate::cuda_pipeline::sp4_isv_slots::ParamGroup,
aux_buffers: &SP4AuxBuffers,
) -> Option<(u64, u64, u64, u64, usize, i32, i32)> {
) -> Option<(Sp4ParamGroupBufs, i32, i32)> {
use crate::cuda_pipeline::sp4_isv_slots::ParamGroup;
let f32_sz = std::mem::size_of::<f32>() as u64;
@@ -9241,11 +9370,6 @@ impl GpuDqnTrainer {
let grads_ptr = self.ptrs.grad_buf;
let m_ptr = self.m_buf.raw_ptr();
let v_ptr = self.v_buf.raw_ptr();
// `s1_input_dim` is computed inside `compute_param_sizes`
// (state_dim or bottleneck_dim + portfolio_dim depending on
// bottleneck activation). Recompute here from
// tensor [0]'s size: w_a_h_s1 = shared_h1 × s1_input_dim,
// so s1_input_dim = param_sizes[0] / shared_h1.
let h_dim = self.config.shared_h1 as i32;
let s1_input_dim = if self.config.shared_h1 > 0 {
(param_sizes[0] / self.config.shared_h1) as i32
@@ -9253,11 +9377,12 @@ impl GpuDqnTrainer {
0
};
let k_in = s1_input_dim;
Some((params_ptr, grads_ptr, m_ptr, v_ptr, count, k_in, h_dim))
Some((
Sp4ParamGroupBufs::single(params_ptr, grads_ptr, m_ptr, v_ptr, count),
k_in, h_dim,
))
}
ParamGroup::DqnValue => {
// Value-head slice = tensors [13..17). k_in/h_dim = 0
// (Pass E gated to group 0 only).
let start = padded_byte_offset(&param_sizes, 13);
let end = padded_byte_offset(&param_sizes, 17);
let count = ((end - start) / f32_sz) as usize;
@@ -9265,10 +9390,12 @@ impl GpuDqnTrainer {
let grads_ptr = self.ptrs.grad_buf + start;
let m_ptr = self.m_buf.raw_ptr() + start;
let v_ptr = self.v_buf.raw_ptr() + start;
Some((params_ptr, grads_ptr, m_ptr, v_ptr, count, 0, 0))
Some((
Sp4ParamGroupBufs::single(params_ptr, grads_ptr, m_ptr, v_ptr, count),
0, 0,
))
}
ParamGroup::DqnBranches => {
// Branch-heads slice = tensors [17..33). k_in/h_dim = 0.
let start = padded_byte_offset(&param_sizes, 17);
let end = padded_byte_offset(&param_sizes, 33);
let count = ((end - start) / f32_sz) as usize;
@@ -9276,42 +9403,16 @@ impl GpuDqnTrainer {
let grads_ptr = self.ptrs.grad_buf + start;
let m_ptr = self.m_buf.raw_ptr() + start;
let v_ptr = self.v_buf.raw_ptr() + start;
Some((params_ptr, grads_ptr, m_ptr, v_ptr, count, 0, 0))
}
ParamGroup::Iqn => {
let b = &aux_buffers.iqn;
Some((b.params_ptr, b.grads_ptr, b.adam_m_ptr, b.adam_v_ptr,
b.count, 0, 0))
}
ParamGroup::IqlHigh => {
let b = &aux_buffers.iql_high;
Some((b.params_ptr, b.grads_ptr, b.adam_m_ptr, b.adam_v_ptr,
b.count, 0, 0))
}
ParamGroup::IqlLow => {
let b = &aux_buffers.iql_low;
Some((b.params_ptr, b.grads_ptr, b.adam_m_ptr, b.adam_v_ptr,
b.count, 0, 0))
}
ParamGroup::Attn => {
let b = &aux_buffers.attn;
Some((b.params_ptr, b.grads_ptr, b.adam_m_ptr, b.adam_v_ptr,
b.count, 0, 0))
}
ParamGroup::Curiosity => {
// `GpuCuriosityTrainer` stores its params/grads/Adam state as
// four separate `[w1, b1, w2, b2]` sub-buffers
// (non-contiguous). A single `(params_ptr, count)` tuple
// can't describe the slice the kernel reads. Resolving
// this requires either a per-layer launch loop (4× the
// kernel cost) or a contiguous-flat re-layout of the
// trainer; both are deeper architectural changes scoped
// beyond Task A7's fix-up. Layer B must therefore guard
// against ISV[170+...] being the natural-zero floor for
// Curiosity-related clamps.
let _ = aux_buffers; // suppress unused warning when only Curiosity is hit
None
Some((
Sp4ParamGroupBufs::single(params_ptr, grads_ptr, m_ptr, v_ptr, count),
0, 0,
))
}
ParamGroup::Iqn => Some((aux_buffers.iqn.clone(), 0, 0)),
ParamGroup::IqlHigh => Some((aux_buffers.iql_high.clone(), 0, 0)),
ParamGroup::IqlLow => Some((aux_buffers.iql_low.clone(), 0, 0)),
ParamGroup::Attn => Some((aux_buffers.attn.clone(), 0, 0)),
ParamGroup::Curiosity => Some((aux_buffers.curiosity.clone(), 0, 0)),
}
}
@@ -12678,6 +12779,28 @@ impl GpuDqnTrainer {
let producer_step_scratch_buf = unsafe { MappedF32Buffer::new(SP4_PRODUCER_COUNT) }
.map_err(|e| MLError::ModelError(format!("SP4 producer_step_scratch_buf alloc: {e}")))?;
// SP4 Task A7 fix-up #2: oracle sub-buffer descriptor tables.
// `SP4_ORACLE_TABLE_MAX_SUB = 4` covers Curiosity's 4 sub-buffers
// (groups 0-6 use only the first slot). Tables are reused across
// all 8 per-step group launches — host overwrites entries [0..n_sub)
// before each launch; the kernel reads them via the device pointer.
// Allocated once at construction so per-step allocator pressure is
// avoided. Shared across all launches because the launcher serialises
// them on the same stream (the `__threadfence_system()` per pass +
// single end-of-loop sync makes the table layout invisible to a
// concurrent reader).
const SP4_ORACLE_TABLE_MAX_SUB: usize = 4;
let oracle_subbuf_table_buf =
unsafe { MappedU64Buffer::new(4 * SP4_ORACLE_TABLE_MAX_SUB) }
.map_err(|e| MLError::ModelError(format!(
"SP4 oracle_subbuf_table_buf alloc: {e}"
)))?;
let oracle_subbuf_counts_buf =
unsafe { MappedI32Buffer::new(SP4_ORACLE_TABLE_MAX_SUB) }
.map_err(|e| MLError::ModelError(format!(
"SP4 oracle_subbuf_counts_buf alloc: {e}"
)))?;
// v8: PopArt running statistics buffers
let popart_mean = stream.alloc_zeros::<f32>(1)
.map_err(|e| MLError::ModelError(format!("alloc popart_mean: {e}")))?;
@@ -14066,6 +14189,8 @@ impl GpuDqnTrainer {
wiener_state_buf,
clamp_engage_per_block_buf,
producer_step_scratch_buf,
oracle_subbuf_table_buf,
oracle_subbuf_counts_buf,
nan_check_buf_ptrs,
nan_check_buf_lens,
pruning_epoch: prune_ep,

View File

@@ -444,6 +444,28 @@ impl CuriosityWeightSet {
})?,
})
}
// ── SP4 Task A7 fix-up #2: per-sub-buffer device-pointer + length ─────
// accessors for the param-group oracle. Curiosity stores its 4 weight
// tensors (w1, b1, w2, b2) as separate `CudaSlice<f32>`s; the oracle
// kernel iterates them via the multi-sub-buffer table.
/// Raw device pointer to the `w1 [CUR_HIDDEN, CUR_INPUT]` weight slice.
pub fn w1_ptr(&self) -> u64 { self.w1.raw_ptr() }
/// Element count of the `w1` slice.
pub fn w1_len(&self) -> usize { self.w1.len() }
/// Raw device pointer to the `b1 [CUR_HIDDEN]` bias slice.
pub fn b1_ptr(&self) -> u64 { self.b1.raw_ptr() }
/// Element count of the `b1` slice.
pub fn b1_len(&self) -> usize { self.b1.len() }
/// Raw device pointer to the `w2 [CUR_OUTPUT, CUR_HIDDEN]` weight slice.
pub fn w2_ptr(&self) -> u64 { self.w2.raw_ptr() }
/// Element count of the `w2` slice.
pub fn w2_len(&self) -> usize { self.w2.len() }
/// Raw device pointer to the `b2 [CUR_OUTPUT]` bias slice.
pub fn b2_ptr(&self) -> u64 { self.b2.raw_ptr() }
/// Element count of the `b2` slice.
pub fn b2_len(&self) -> usize { self.b2.len() }
}
/// GPU buffers holding all 6 weight tensors of the PPO Actor (PolicyNetwork).

View File

@@ -514,6 +514,15 @@ pub struct MappedU64Buffer {
pub len: usize,
}
impl std::fmt::Debug for MappedU64Buffer {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("MappedU64Buffer")
.field("dev_ptr", &self.dev_ptr)
.field("len", &self.len)
.finish()
}
}
unsafe impl Send for MappedU64Buffer {}
unsafe impl Sync for MappedU64Buffer {}

View File

@@ -63,12 +63,24 @@ __device__ __forceinline__ float block_reduce_sum_256(float val, float* s_buf, i
return s_buf[0];
}
// SP4 Task A7 fix-up #2 (2026-04-30): the kernel now accepts an array of
// sub-buffer pointers + counts so groups with non-contiguous storage
// (Curiosity's `[w1, b1, w2, b2]` layout) can be described as a single
// logical group. For groups 0-6 the launcher passes `n_sub=1` and the
// behaviour is identical to the original single-pointer signature. For
// group 7 (Curiosity) it passes `n_sub=4` and the union of the four
// sub-buffers is treated as one distribution for Pass A/B/C p99 and
// Pass D (`Σ w·g`, `Σ w²`, `Σ|g|`, `Σ|w|`). Pass E (trunk-only) operates
// on the contiguous trunk grad-w slice via the `g_idx == 0` arm and is
// guaranteed to see `n_sub=1` (Curiosity has `l1_lambda_scratch_idx==-1`).
extern "C" __global__ void param_group_oracle_update(
const float* __restrict__ params,
const float* __restrict__ grads,
const float* __restrict__ adam_m,
const float* __restrict__ adam_v,
int count, /* slice length (elements) */
const unsigned long long* __restrict__ params_ptrs, /* [n_sub] device pointers */
const unsigned long long* __restrict__ grads_ptrs, /* [n_sub] device pointers */
const unsigned long long* __restrict__ adam_m_ptrs, /* [n_sub] device pointers */
const unsigned long long* __restrict__ adam_v_ptrs, /* [n_sub] device pointers */
const int* __restrict__ sub_counts, /* [n_sub] per-sub-buffer element counts */
int n_sub, /* number of sub-buffers (1 for groups 0-6, 4 for group 7) */
int total_count, /* sum of sub_counts (used for p99 normalisation + WD inv_count) */
int k_in, /* trunk input dim (Pass E only; 0 for non-trunk) */
int h_dim, /* trunk output dim of layer (Pass E only) */
float* __restrict__ scratch_buf,
@@ -83,7 +95,7 @@ extern "C" __global__ void param_group_oracle_update(
// ── Pass A: WEIGHT_BOUND[group] = p99(|params|) ─────────────────────
{
float p99 = sp4_histogram_p99<256>(params, count);
float p99 = sp4_histogram_p99_multi<256>(params_ptrs, sub_counts, n_sub, total_count);
if (tid == 0) {
scratch_buf[weight_scratch_idx] = p99;
__threadfence_system(); // make host-mapped write visible
@@ -93,7 +105,7 @@ extern "C" __global__ void param_group_oracle_update(
// ── Pass B: ADAM_M_BOUND[group] = p99(|adam_m|) ─────────────────────
{
float p99 = sp4_histogram_p99<256>(adam_m, count);
float p99 = sp4_histogram_p99_multi<256>(adam_m_ptrs, sub_counts, n_sub, total_count);
if (tid == 0) {
scratch_buf[adam_m_scratch_idx] = p99;
__threadfence_system();
@@ -103,7 +115,7 @@ extern "C" __global__ void param_group_oracle_update(
// ── Pass C: ADAM_V_BOUND[group] = p99(|adam_v|) ─────────────────────
{
float p99 = sp4_histogram_p99<256>(adam_v, count);
float p99 = sp4_histogram_p99_multi<256>(adam_v_ptrs, sub_counts, n_sub, total_count);
if (tid == 0) {
scratch_buf[adam_v_scratch_idx] = p99;
__threadfence_system();
@@ -112,11 +124,10 @@ extern "C" __global__ void param_group_oracle_update(
__syncthreads();
// ── Pass D: WD_RATE[group] = |Σ w·g| / max(Σ w², EPS_DIV) ──────────
// Side products: mean|g| = Σ|g|/count, mean|w| = Σ|w|/count for Pass E.
//
// 4 parallel block-wide reductions: each thread maintains 4 register
// accumulators (sum_wg, sum_w2, sum_abs_g, sum_abs_w), then we reduce
// each tree-style through reused shared memory.
// Side products: mean|g| = Σ|g|/total_count, mean|w| = Σ|w|/total_count
// for Pass E. Iterates `n_sub` sub-buffers in the per-thread accumulator
// loop; for `n_sub=1` this is equivalent to the original single-buffer
// implementation.
//
// Reuses the histogram dynamic shmem region (`s_warp_tiles`) as the
// per-thread accumulator buffer; the Pass C `__syncthreads` above
@@ -132,13 +143,18 @@ extern "C" __global__ void param_group_oracle_update(
float t_sum_abs_g = 0.0f;
float t_sum_abs_w = 0.0f;
for (int i = tid; i < count; i += 256) {
float w = params[i];
float g = grads[i];
t_sum_wg += w * g;
t_sum_w2 += w * w;
t_sum_abs_g += fabsf(g);
t_sum_abs_w += fabsf(w);
for (int s = 0; s < n_sub; ++s) {
const float* w_buf = reinterpret_cast<const float*>(params_ptrs[s]);
const float* g_buf = reinterpret_cast<const float*>(grads_ptrs[s]);
const int cnt_s = sub_counts[s];
for (int i = tid; i < cnt_s; i += 256) {
float w = w_buf[i];
float g = g_buf[i];
t_sum_wg += w * g;
t_sum_w2 += w * w;
t_sum_abs_g += fabsf(g);
t_sum_abs_w += fabsf(w);
}
}
// Four sequential block-reduces sharing the same 256-float shmem region.
@@ -155,7 +171,7 @@ extern "C" __global__ void param_group_oracle_update(
__shared__ float s_mean_abs_w;
__shared__ float s_wd_rate;
if (tid == 0) {
const float inv_count = (count > 0) ? (1.0f / (float)count) : 0.0f;
const float inv_count = (total_count > 0) ? (1.0f / (float)total_count) : 0.0f;
s_mean_abs_g = sum_abs_g * inv_count;
s_mean_abs_w = sum_abs_w * inv_count;
const float denom = fmaxf(sum_w2, EPS_DIV);
@@ -173,6 +189,12 @@ extern "C" __global__ void param_group_oracle_update(
// (h_dim-1)*k_in + i]` in row-major. We compute the per-feature L2
// norm by squaring + summing across `h_dim` then sqrt.
//
// Pass E reads the contiguous trunk grad-w sub-buffer at
// `grads_ptrs[0]` directly — group 0 is contiguous so `n_sub == 1` and
// the `[h_dim × k_in]` weight-matrix layout sits in sub_counts[0].
// For multi-sub-buffer groups (Curiosity) `l1_lambda_scratch_idx==-1`
// and Pass E is gated off entirely.
//
// Feature normalization → discrete distribution; Shannon entropy H
// → entropy_deficit = (log K H) / log K ∈ [0, 1] where K = k_in
// (1 = perfectly concentrated on one feature; 0 = uniform). The L1
@@ -192,6 +214,11 @@ extern "C" __global__ void param_group_oracle_update(
}
__syncthreads();
// Read the contiguous trunk grads sub-buffer. Pass E is only
// dispatched for group 0 (DqnTrunk) which has `n_sub==1`, so
// `grads_ptrs[0]` is the canonical trunk grad-w pointer.
const float* grads_e = reinterpret_cast<const float*>(grads_ptrs[0]);
// Each thread accumulates squared gradients for a strided slice
// of the [k_in × h_dim] grad_w matrix (= grads[0..k_in*h_dim)).
// Per-feature accumulation is naturally race-free across threads
@@ -212,7 +239,7 @@ extern "C" __global__ void param_group_oracle_update(
for (int h = tid; h < h_dim; h += 256) {
int idx = h * k_in + i;
if (idx < kh) {
float gv = grads[idx];
float gv = grads_e[idx];
partial += gv * gv;
}
}

View File

@@ -118,3 +118,110 @@ __device__ float sp4_histogram_p99(const float* __restrict__ buf, int count) {
__syncthreads();
return s_step_max;
}
// ── Multi-sub-buffer variant (SP4 Task A7 fix-up #2) ─────────────────────
//
// `sp4_histogram_p99_multi` treats the union of `n_sub` sub-buffers as one
// logical distribution. Sub-buffer pointers + counts are read from device
// arrays (typically mapped-pinned tables populated host-side) so the
// caller can dispatch with `n_sub=1` (the contiguous-buffer common case;
// SP4 groups 0-6) or `n_sub=4` (Curiosity's `[w1, b1, w2, b2]` layout).
//
// Pass structure mirrors the single-buffer template:
// - Pass 1: max-reduce across all `n_sub` sub-buffers.
// - Pass 2: linear-bin all `n_sub` sub-buffers into the same per-warp
// tile space; binned counts reflect the union distribution.
// - Pass 3: cumulative-from-top sees `total_count` (sum of sub_counts);
// yields the same p99 semantics as the single-buffer template
// when n_sub == 1.
//
// `sub_buf_ptrs`/`sub_counts` must be readable from the device (mapped-
// pinned device-pointer is the canonical case). `n_sub` is a small
// integer (≤ 8 in practice) so the per-pass `for (s)` loops add no
// material overhead vs the inner per-element loops.
template <int BLOCK_SIZE>
__device__ float sp4_histogram_p99_multi(
const unsigned long long* __restrict__ sub_buf_ptrs,
const int* __restrict__ sub_counts,
int n_sub,
int total_count
) {
static_assert(BLOCK_SIZE >= 32 && (BLOCK_SIZE % 32) == 0,
"BLOCK_SIZE must be a positive multiple of 32 (warp size)");
__shared__ int s_bins[SP4_HIST_BINS];
__shared__ float s_step_max;
const int tid = threadIdx.x;
// ── Pass 1: block-wide max-reduce of |buf| across all sub-buffers ──
float local_max = 0.0f;
for (int s = 0; s < n_sub; ++s) {
const float* buf_s = reinterpret_cast<const float*>(sub_buf_ptrs[s]);
const int cnt_s = sub_counts[s];
for (int i = tid; i < cnt_s; i += BLOCK_SIZE) {
local_max = fmaxf(local_max, fabsf(buf_s[i]));
}
}
s_bins[tid] = __float_as_int(local_max);
__syncthreads();
for (int s = BLOCK_SIZE / 2; s > 0; s >>= 1) {
if (tid < s) {
float a = __int_as_float(s_bins[tid]);
float b = __int_as_float(s_bins[tid + s]);
s_bins[tid] = __float_as_int(fmaxf(a, b));
}
__syncthreads();
}
if (tid == 0) s_step_max = __int_as_float(s_bins[0]);
__syncthreads();
if (s_step_max == 0.0f) return 0.0f;
const float step_max = s_step_max;
const float bin_width = step_max / (float)SP4_HIST_BINS;
// ── Pass 2: linear-bin all sub-buffers into per-warp tiles ─────────
extern __shared__ int s_warp_tiles[];
const int warp_id = tid >> 5;
const int lane = tid & 31;
const int warps = BLOCK_SIZE >> 5;
for (int b = lane; b < SP4_HIST_BINS; b += 32) {
s_warp_tiles[warp_id * SP4_HIST_BINS + b] = 0;
}
__syncwarp();
for (int s = 0; s < n_sub; ++s) {
const float* buf_s = reinterpret_cast<const float*>(sub_buf_ptrs[s]);
const int cnt_s = sub_counts[s];
for (int i = tid; i < cnt_s; i += BLOCK_SIZE) {
int bin_idx = (int)floorf(fabsf(buf_s[i]) / bin_width);
if (bin_idx >= SP4_HIST_BINS) bin_idx = SP4_HIST_BINS - 1;
s_warp_tiles[warp_id * SP4_HIST_BINS + bin_idx] += 1;
}
}
__syncthreads();
// Tree-reduce warp tiles into s_bins[].
for (int b = tid; b < SP4_HIST_BINS; b += BLOCK_SIZE) {
int sum = 0;
for (int w = 0; w < warps; ++w) sum += s_warp_tiles[w * SP4_HIST_BINS + b];
s_bins[b] = sum;
}
__syncthreads();
// ── Pass 3: cumulative-from-top → p99 over the union distribution ──
if (tid == 0) {
const int target = (total_count + 99) / 100; // ceil(total_count / 100)
int cumul = 0;
int p99_bin = SP4_HIST_BINS - 1;
for (int b = SP4_HIST_BINS - 1; b >= 0; --b) {
cumul += s_bins[b];
if (cumul >= target) { p99_bin = b; break; }
}
s_step_max = (float)(p99_bin + 1) * bin_width;
}
__syncthreads();
return s_step_max;
}

View File

@@ -2307,75 +2307,118 @@ impl FusedTrainingCtx {
.map_err(|e| anyhow::anyhow!("stochastic depth mask: {e}"))
}
/// SP4 Layer A Task A7 fix-up: build the aux-trainer buffer descriptor
/// `SP4AuxBuffers` that `GpuDqnTrainer::launch_sp4_param_group_oracles_all_groups`
/// consumes for groups 3-6 (IQN, IQL hi/lo, Attn).
/// SP4 Layer A Task A7 fix-up + fix-up #2: build the aux-trainer buffer
/// descriptor `SP4AuxBuffers` that
/// `GpuDqnTrainer::launch_sp4_param_group_oracles_all_groups` consumes
/// for groups 3-7 (IQN, IQL hi/lo, Attn, Curiosity).
///
/// `gpu_iqn` and `gpu_attention` are `Option<_>` purely as init-failure
/// fallback (graceful degrade) and are otherwise always present in
/// production; if either is `None`, this helper emits a zero-count
/// descriptor for that group so the launcher's `count == 0` short-circuit
/// silently skips the kernel launch (matches the existing skip-then-no-op
/// path for non-launched groups). Group 7 (Curiosity) is omitted from
/// `SP4AuxBuffers` per the architectural hold-out documented on
/// `param_group_buffers` — its `[w1, b1, w2, b2]` layout cannot be
/// described by a single contiguous slice.
/// production; if either is `None`, this helper emits an empty
/// descriptor for that group so the launcher's `total_count == 0`
/// short-circuit silently skips the kernel launch.
///
/// **Curiosity (group 7)** lives outside `FusedTrainingCtx` —
/// `CuriosityWeightSet` is owned by the experience collector and the
/// matching `GpuCuriosityTrainer` (grad + Adam state) is also there.
/// Layer B's training-loop caller threads them in via the
/// `curiosity_weights` and `curiosity_trainer` parameters; both must
/// be `Some` together (caller responsibility — they live on the same
/// collector) and the descriptor describes the union of all four
/// `[w1, b1, w2, b2]` sub-buffers as one logical group. Passing
/// `None` for either yields an empty curiosity descriptor and the
/// launcher silently skips group 7.
#[allow(dead_code)] // Layer B will consume this; lint suppressed until then.
pub(crate) fn build_sp4_aux_buffers(&self)
-> crate::cuda_pipeline::gpu_dqn_trainer::SP4AuxBuffers
pub(crate) fn build_sp4_aux_buffers(
&self,
curiosity_weights: Option<&crate::cuda_pipeline::gpu_weights::CuriosityWeightSet>,
curiosity_trainer: Option<&crate::cuda_pipeline::gpu_curiosity_trainer::GpuCuriosityTrainer>,
) -> crate::cuda_pipeline::gpu_dqn_trainer::SP4AuxBuffers
{
use crate::cuda_pipeline::gpu_dqn_trainer::{SP4AuxBuffers, Sp4ParamGroupBufs};
// Zero-count placeholder (count=0 trips the launcher's skip branch).
let placeholder = Sp4ParamGroupBufs {
params_ptr: 0,
grads_ptr: 0,
adam_m_ptr: 0,
adam_v_ptr: 0,
count: 0,
};
use crate::cuda_pipeline::gpu_dqn_trainer::{SP4AuxBuffers, Sp4ParamGroupBufs, Sp4SubBuffer};
let iqn = if let Some(iqn) = self.gpu_iqn.as_ref() {
Sp4ParamGroupBufs {
params_ptr: iqn.online_params_ptr(),
grads_ptr: iqn.online_grad_ptr(),
adam_m_ptr: iqn.adam_m_ptr(),
adam_v_ptr: iqn.adam_v_ptr(),
count: iqn.online_params_len(),
}
Sp4ParamGroupBufs::single(
iqn.online_params_ptr(),
iqn.online_grad_ptr(),
iqn.adam_m_ptr(),
iqn.adam_v_ptr(),
iqn.online_params_len(),
)
} else {
placeholder
Sp4ParamGroupBufs::empty()
};
let iql_high = Sp4ParamGroupBufs {
params_ptr: self.gpu_iql.params_ptr(),
grads_ptr: self.gpu_iql.grads_ptr(),
adam_m_ptr: self.gpu_iql.adam_m_ptr(),
adam_v_ptr: self.gpu_iql.adam_v_ptr(),
count: self.gpu_iql.params_len(),
};
let iql_high = Sp4ParamGroupBufs::single(
self.gpu_iql.params_ptr(),
self.gpu_iql.grads_ptr(),
self.gpu_iql.adam_m_ptr(),
self.gpu_iql.adam_v_ptr(),
self.gpu_iql.params_len(),
);
let iql_low = Sp4ParamGroupBufs {
params_ptr: self.gpu_iql_low.params_ptr(),
grads_ptr: self.gpu_iql_low.grads_ptr(),
adam_m_ptr: self.gpu_iql_low.adam_m_ptr(),
adam_v_ptr: self.gpu_iql_low.adam_v_ptr(),
count: self.gpu_iql_low.params_len(),
};
let iql_low = Sp4ParamGroupBufs::single(
self.gpu_iql_low.params_ptr(),
self.gpu_iql_low.grads_ptr(),
self.gpu_iql_low.adam_m_ptr(),
self.gpu_iql_low.adam_v_ptr(),
self.gpu_iql_low.params_len(),
);
let attn = if let Some(attn) = self.gpu_attention.as_ref() {
Sp4ParamGroupBufs {
params_ptr: attn.params_ptr(),
grads_ptr: attn.grads_ptr(),
adam_m_ptr: attn.adam_m_ptr(),
adam_v_ptr: attn.adam_v_ptr(),
count: attn.params_len(),
}
Sp4ParamGroupBufs::single(
attn.params_ptr(),
attn.grads_ptr(),
attn.adam_m_ptr(),
attn.adam_v_ptr(),
attn.params_len(),
)
} else {
placeholder
Sp4ParamGroupBufs::empty()
};
SP4AuxBuffers { iqn, iql_high, iql_low, attn }
// Curiosity multi-sub-buffer descriptor. Both halves of state
// (params from CuriosityWeightSet; grad/Adam from GpuCuriosityTrainer)
// must be supplied — Layer B's caller fetches both from the same
// GpuExperienceCollector. If either is missing (curiosity disabled
// at collector init), emit an empty descriptor.
let curiosity = match (curiosity_weights, curiosity_trainer) {
(Some(w), Some(t)) => Sp4ParamGroupBufs {
sub_buffers: vec![
Sp4SubBuffer {
params_ptr: w.w1_ptr(),
grads_ptr: t.grad_w1_ptr(),
adam_m_ptr: t.adam_m_w1_ptr(),
adam_v_ptr: t.adam_v_w1_ptr(),
count: w.w1_len(),
},
Sp4SubBuffer {
params_ptr: w.b1_ptr(),
grads_ptr: t.grad_b1_ptr(),
adam_m_ptr: t.adam_m_b1_ptr(),
adam_v_ptr: t.adam_v_b1_ptr(),
count: w.b1_len(),
},
Sp4SubBuffer {
params_ptr: w.w2_ptr(),
grads_ptr: t.grad_w2_ptr(),
adam_m_ptr: t.adam_m_w2_ptr(),
adam_v_ptr: t.adam_v_w2_ptr(),
count: w.w2_len(),
},
Sp4SubBuffer {
params_ptr: w.b2_ptr(),
grads_ptr: t.grad_b2_ptr(),
adam_m_ptr: t.adam_m_b2_ptr(),
adam_v_ptr: t.adam_v_b2_ptr(),
count: w.b2_len(),
},
],
},
_ => Sp4ParamGroupBufs::empty(),
};
SP4AuxBuffers { iqn, iql_high, iql_low, attn, curiosity }
}
/// Submit IQL advantage-weighted TD error modulation (captured as iql_modulate_child).

View File

@@ -22,7 +22,7 @@
use std::sync::Arc;
use cudarc::driver::{CudaContext, CudaFunction, CudaStream, LaunchConfig, PushKernelArg};
use ml::cuda_pipeline::mapped_pinned::{MappedF32Buffer, MappedI32Buffer};
use ml::cuda_pipeline::mapped_pinned::{MappedF32Buffer, MappedI32Buffer, MappedU64Buffer};
/// Test-only cubin for the SP4 histogram-p99 wrapper kernel. Built by
/// `crates/ml/build.rs`; the cubin path is the standard `OUT_DIR` slot
@@ -618,21 +618,37 @@ fn sample_p99(xs: &[f32]) -> f32 {
sorted[(sorted.len() * 99) / 100]
}
/// One sub-buffer's host-side data for a kernel-direct test launch.
/// Mirrors the production `Sp4SubBuffer` Rust type but holds host
/// vectors instead of device pointers. The test launcher allocates
/// matching mapped-pinned device-visible buffers and populates them.
struct TestSubBuffer {
params_host: Vec<f32>,
grads_host: Vec<f32>,
adam_m_host: Vec<f32>,
adam_v_host: Vec<f32>,
}
impl TestSubBuffer {
fn count(&self) -> usize {
let n = self.params_host.len();
assert_eq!(self.grads_host.len(), n);
assert_eq!(self.adam_m_host.len(), n);
assert_eq!(self.adam_v_host.len(), n);
n
}
}
/// Launch the production `param_group_oracle_update` kernel kernel-direct
/// with controlled `(params, grads, adam_m, adam_v)` slices and write into a
/// 47-slot scratch buffer at the production-layout slot offsets for group
/// `g_idx`. Returns the host-readable scratch buffer for assertion.
/// across one or more sub-buffers (groups 0-6: 1 sub-buffer; group 7: 4
/// sub-buffers). Returns the host-readable scratch buffer for assertion.
///
/// `k_in`/`h_dim` non-zero only for `g_idx == 0` — Pass E (L1 lambda)
/// gating on `l1_lambda_scratch_idx >= 0`. For other groups, pass 0/0/-1.
#[allow(clippy::too_many_arguments)]
fn launch_sp4_param_group_oracle_for_group(
stream: &Arc<CudaStream>,
g_idx: usize,
params_host: &[f32],
grads_host: &[f32],
adam_m_host: &[f32],
adam_v_host: &[f32],
sub_buffers: &[TestSubBuffer],
k_in: i32,
h_dim: i32,
) -> Vec<f32> {
@@ -643,48 +659,84 @@ fn launch_sp4_param_group_oracle_for_group(
const SCRATCH_BASE_WD: usize = 29;
const SCRATCH_L1_TRUNK: usize = 37;
const SHARED_BYTES: u32 = (256 / 32) * 256 * 4;
const K_MAX: usize = 4; // matches production `SP4_ORACLE_TABLE_MAX_SUB`
let n = params_host.len();
assert_eq!(grads_host.len(), n);
assert_eq!(adam_m_host.len(), n);
assert_eq!(adam_v_host.len(), n);
let n_sub = sub_buffers.len();
assert!(n_sub >= 1 && n_sub <= K_MAX,
"n_sub={n_sub} must be in [1, {K_MAX}]");
let total_count: usize = sub_buffers.iter().map(|sb| sb.count()).sum();
assert!(total_count > 0, "total_count must be positive");
let kernel = load_sp4_param_group_oracle_kernel(stream);
// Safety: CUDA context active on this thread (resolved via stream).
let params_buf = unsafe { MappedF32Buffer::new(n) }.expect("alloc params");
let grads_buf = unsafe { MappedF32Buffer::new(n) }.expect("alloc grads");
let adam_m_buf = unsafe { MappedF32Buffer::new(n) }.expect("alloc adam_m");
let adam_v_buf = unsafe { MappedF32Buffer::new(n) }.expect("alloc adam_v");
params_buf.write_from_slice(params_host);
grads_buf.write_from_slice(grads_host);
adam_m_buf.write_from_slice(adam_m_host);
adam_v_buf.write_from_slice(adam_v_host);
// Allocate per-sub-buffer mapped-pinned device-visible storage and
// write the host data through. We hold the buffers alive in a Vec
// for the duration of the kernel launch.
let mut buf_owners: Vec<(MappedF32Buffer, MappedF32Buffer, MappedF32Buffer, MappedF32Buffer)>
= Vec::with_capacity(n_sub);
for sb in sub_buffers {
let n = sb.count();
let p_buf = unsafe { MappedF32Buffer::new(n) }.expect("alloc params");
let g_buf = unsafe { MappedF32Buffer::new(n) }.expect("alloc grads");
let m_buf = unsafe { MappedF32Buffer::new(n) }.expect("alloc adam_m");
let v_buf = unsafe { MappedF32Buffer::new(n) }.expect("alloc adam_v");
p_buf.write_from_slice(&sb.params_host);
g_buf.write_from_slice(&sb.grads_host);
m_buf.write_from_slice(&sb.adam_m_host);
v_buf.write_from_slice(&sb.adam_v_host);
buf_owners.push((p_buf, g_buf, m_buf, v_buf));
}
// Build the device-pointer table: 4 ptr-arrays of length K_MAX,
// packed contiguously in a single mapped-pinned u64 buffer. Counts
// sit in a parallel mapped-pinned i32 buffer of length K_MAX.
let table_buf = unsafe { MappedU64Buffer::new(4 * K_MAX) }.expect("alloc table");
let counts_buf = unsafe { MappedI32Buffer::new(K_MAX) }.expect("alloc counts");
{
// Build host-side images then write them through.
let mut table_img: Vec<u64> = vec![0; 4 * K_MAX];
let mut counts_img: Vec<i32> = vec![0; K_MAX];
for (s, owners) in buf_owners.iter().enumerate() {
table_img[0 * K_MAX + s] = owners.0.dev_ptr;
table_img[1 * K_MAX + s] = owners.1.dev_ptr;
table_img[2 * K_MAX + s] = owners.2.dev_ptr;
table_img[3 * K_MAX + s] = owners.3.dev_ptr;
counts_img[s] = sub_buffers[s].count() as i32;
}
table_buf.write_from_slice(&table_img);
counts_buf.write_from_slice(&counts_img);
}
let scratch_buf =
unsafe { MappedF32Buffer::new(SP4_PRODUCER_COUNT) }.expect("alloc scratch");
let count_i32 = i32::try_from(n).expect("count fits in i32");
let n_sub_i32 = n_sub as i32;
let total_count_i32 = total_count as i32;
let weight_idx = (SCRATCH_BASE_W + g_idx) as i32;
let adam_m_idx = (SCRATCH_BASE_M + g_idx) as i32;
let adam_v_idx = (SCRATCH_BASE_V + g_idx) as i32;
let wd_rate_idx = (SCRATCH_BASE_WD + g_idx) as i32;
let l1_idx: i32 = if g_idx == 0 { SCRATCH_L1_TRUNK as i32 } else { -1 };
let params_dev = params_buf.dev_ptr;
let grads_dev = grads_buf.dev_ptr;
let adam_m_dev = adam_m_buf.dev_ptr;
let adam_v_dev = adam_v_buf.dev_ptr;
let table_dev = table_buf.dev_ptr;
let stride_bytes = (K_MAX * std::mem::size_of::<u64>()) as u64;
let params_ptrs_dev = table_dev;
let grads_ptrs_dev = table_dev + stride_bytes;
let adam_m_ptrs_dev = table_dev + 2 * stride_bytes;
let adam_v_ptrs_dev = table_dev + 3 * stride_bytes;
let counts_dev = counts_buf.dev_ptr;
let scratch_dev = scratch_buf.dev_ptr;
unsafe {
stream
.launch_builder(&kernel)
.arg(&params_dev)
.arg(&grads_dev)
.arg(&adam_m_dev)
.arg(&adam_v_dev)
.arg(&count_i32)
.arg(&params_ptrs_dev)
.arg(&grads_ptrs_dev)
.arg(&adam_m_ptrs_dev)
.arg(&adam_v_ptrs_dev)
.arg(&counts_dev)
.arg(&n_sub_i32)
.arg(&total_count_i32)
.arg(&k_in)
.arg(&h_dim)
.arg(&scratch_dev)
@@ -794,35 +846,60 @@ fn sp4_param_group_oracle_per_group_writes_distinct_isv_slots() {
let sigma_m = 0.05 * scale;
let sigma_v = 0.001 * scale;
let seed_w = 0xA7_5A_F1_57_u64 ^ ((g_idx as u64) << 32) ^ 0x01;
let seed_g = 0xA7_5A_F1_57_u64 ^ ((g_idx as u64) << 32) ^ 0x02;
let seed_m = 0xA7_5A_F1_57_u64 ^ ((g_idx as u64) << 32) ^ 0x03;
let seed_v = 0xA7_5A_F1_57_u64 ^ ((g_idx as u64) << 32) ^ 0x04;
// Group 7 (Curiosity) exercises the multi-sub-buffer path with 4
// sub-buffers of distinct shapes (w1/b1/w2/b2-like sizes). All
// other groups use a single contiguous N-element sub-buffer.
let sub_shapes: Vec<usize> = if g_idx == 7 {
// Mirror Curiosity production layout: w1=[128×45]=5760, b1=128,
// w2=[42×128]=5376, b2=42 (total 11306) — but scale to fit the
// test budget. Use 1024/32/1024/32 to total 2112 elements (kept
// small for fast tests; still exercises the multi-sub-buffer
// iteration in Pass A/B/C/D).
vec![1024, 32, 1024, 32]
} else {
vec![N]
};
let params_host = boxmuller_abs_normal(seed_w, N, sigma_w);
// grads: signed (mix of positive/negative) so Σ w·g doesn't trivially
// collapse to Σ w·|g|. Reuse boxmuller for magnitudes, alternate
// signs by index parity → roughly zero-mean but non-trivial Σ w·g.
let grads_mag = boxmuller_abs_normal(seed_g, N, sigma_g);
let grads_host: Vec<f32> = grads_mag
.iter()
.enumerate()
.map(|(i, &g)| if i % 2 == 0 { g } else { -g })
.collect();
let adam_m_host = boxmuller_abs_normal(seed_m, N, sigma_m);
let adam_v_host = boxmuller_abs_normal(seed_v, N, sigma_v);
let mut sub_buffers: Vec<TestSubBuffer> = Vec::with_capacity(sub_shapes.len());
for (s, &shape_n) in sub_shapes.iter().enumerate() {
// Per-sub-buffer seed offsets so distinct sub-buffers have
// distinct distributions (catches buffer-mixup bugs in the
// kernel's sub-buffer iteration).
let seed_w = 0xA7_5A_F1_57_u64 ^ ((g_idx as u64) << 32) ^ ((s as u64) << 16) ^ 0x01;
let seed_g = 0xA7_5A_F1_57_u64 ^ ((g_idx as u64) << 32) ^ ((s as u64) << 16) ^ 0x02;
let seed_m = 0xA7_5A_F1_57_u64 ^ ((g_idx as u64) << 32) ^ ((s as u64) << 16) ^ 0x03;
let seed_v = 0xA7_5A_F1_57_u64 ^ ((g_idx as u64) << 32) ^ ((s as u64) << 16) ^ 0x04;
// For group 0 (DqnTrunk) the kernel uses k_in × h_dim = N total to
// drive Pass E. We pick k_in = 64, h_dim = 64 so k_in × h_dim = N.
let params_host = boxmuller_abs_normal(seed_w, shape_n, sigma_w);
// grads: signed (mix of positive/negative) so Σ w·g doesn't
// trivially collapse to Σ w·|g|. Alternate signs by index parity.
let grads_mag = boxmuller_abs_normal(seed_g, shape_n, sigma_g);
let grads_host: Vec<f32> = grads_mag
.iter()
.enumerate()
.map(|(i, &g)| if i % 2 == 0 { g } else { -g })
.collect();
let adam_m_host = boxmuller_abs_normal(seed_m, shape_n, sigma_m);
let adam_v_host = boxmuller_abs_normal(seed_v, shape_n, sigma_v);
sub_buffers.push(TestSubBuffer {
params_host,
grads_host,
adam_m_host,
adam_v_host,
});
}
// For group 0 (DqnTrunk) the kernel uses k_in × h_dim total
// elements (= sub_buffers[0].count() since group 0 is single-sub-
// buffer) to drive Pass E. We pick k_in = 64, h_dim = 64 so the
// matrix has N=4096 elements.
let (k_in, h_dim) = if g_idx == 0 { (64_i32, 64_i32) } else { (0_i32, 0_i32) };
let scratch = launch_sp4_param_group_oracle_for_group(
&stream,
g_idx,
&params_host,
&grads_host,
&adam_m_host,
&adam_v_host,
&sub_buffers,
k_in,
h_dim,
);
@@ -847,12 +924,27 @@ fn sp4_param_group_oracle_per_group_writes_distinct_isv_slots() {
}
}
// Build the union distributions for reference computation.
// p99 is computed over the union of all sub-buffer entries — this
// matches the kernel's `sp4_histogram_p99_multi` semantics.
let union_params: Vec<f32> = sub_buffers.iter()
.flat_map(|sb| sb.params_host.iter().copied()).collect();
let union_grads: Vec<f32> = sub_buffers.iter()
.flat_map(|sb| sb.grads_host.iter().copied()).collect();
let union_adam_m: Vec<f32> = sub_buffers.iter()
.flat_map(|sb| sb.adam_m_host.iter().copied()).collect();
let union_adam_v: Vec<f32> = sub_buffers.iter()
.flat_map(|sb| sb.adam_v_host.iter().copied()).collect();
let union_total_count: usize = union_params.len();
// ── Pass A: WEIGHT_BOUND[g] p99 ──
let true_p99_w = sample_p99(&params_host);
let true_p99_w = sample_p99(&union_params);
let kernel_p99_w = scratch[SCRATCH_BASE_W + g_idx];
let rel_w = ((kernel_p99_w - true_p99_w) / true_p99_w).abs();
println!(
"SP4 param_group[{g_idx}] WEIGHT — true_p99={true_p99_w:.5}, kernel_p99={kernel_p99_w:.5}, rel_err={rel_w:.5}",
"SP4 param_group[{g_idx}] WEIGHT — true_p99={true_p99_w:.5}, kernel_p99={kernel_p99_w:.5}, rel_err={rel_w:.5} (n_sub={} of len {})",
sub_buffers.len(),
sub_buffers.iter().map(|sb| sb.count().to_string()).collect::<Vec<_>>().join("+"),
);
assert!(kernel_p99_w > 0.0, "WEIGHT p99 must be positive");
assert!(
@@ -861,7 +953,7 @@ fn sp4_param_group_oracle_per_group_writes_distinct_isv_slots() {
);
// ── Pass B: ADAM_M_BOUND[g] p99 ──
let true_p99_m = sample_p99(&adam_m_host);
let true_p99_m = sample_p99(&union_adam_m);
let kernel_p99_m = scratch[SCRATCH_BASE_M + g_idx];
let rel_m = ((kernel_p99_m - true_p99_m) / true_p99_m).abs();
println!(
@@ -874,7 +966,7 @@ fn sp4_param_group_oracle_per_group_writes_distinct_isv_slots() {
);
// ── Pass C: ADAM_V_BOUND[g] p99 ──
let true_p99_v = sample_p99(&adam_v_host);
let true_p99_v = sample_p99(&union_adam_v);
let kernel_p99_v = scratch[SCRATCH_BASE_V + g_idx];
let rel_v = ((kernel_p99_v - true_p99_v) / true_p99_v).abs();
println!(
@@ -889,13 +981,12 @@ fn sp4_param_group_oracle_per_group_writes_distinct_isv_slots() {
// ── Pass D: WD_RATE[g] = |Σ w·g| / max(Σ w², EPS_DIV) ──
// Analytically computed in f64 host-side then cast to f32 to bound
// accumulation order error. The kernel does the same reduction in
// f32 only, so a 2% tolerance comfortably covers summation drift
// across 4096 elements.
// f32 only over the union of sub-buffers.
let mut sum_wg: f64 = 0.0;
let mut sum_w2: f64 = 0.0;
for i in 0..N {
let w = params_host[i] as f64;
let g = grads_host[i] as f64;
for i in 0..union_total_count {
let w = union_params[i] as f64;
let g = union_grads[i] as f64;
sum_wg += w * g;
sum_w2 += w * w;
}
@@ -923,7 +1014,13 @@ fn sp4_param_group_oracle_per_group_writes_distinct_isv_slots() {
if g_idx == 0 {
let k_in_us = k_in as usize;
let h_dim_us = h_dim as usize;
assert_eq!(k_in_us * h_dim_us, N);
// Group 0 (DqnTrunk) is single-sub-buffer by construction; the
// trunk grad-w matrix sits in `sub_buffers[0]` and `union_*`
// are identical.
assert_eq!(sub_buffers.len(), 1);
assert_eq!(k_in_us * h_dim_us, union_total_count);
let grads_host = &sub_buffers[0].grads_host;
let params_host = &sub_buffers[0].params_host;
// Per-feature L2 norms (column-wise across the [h_dim × k_in]
// row-major grad matrix → feature `i` indexed at
@@ -951,10 +1048,11 @@ fn sp4_param_group_oracle_per_group_writes_distinct_isv_slots() {
}
deficit = ((log_k - entropy) / log_k).clamp(0.0, 1.0);
}
let total_count_f64 = union_total_count as f64;
let mean_abs_g: f64 =
grads_host.iter().map(|&g| (g as f64).abs()).sum::<f64>() / (N as f64);
grads_host.iter().map(|&g| (g as f64).abs()).sum::<f64>() / total_count_f64;
let mean_abs_w: f64 =
params_host.iter().map(|&w| (w as f64).abs()).sum::<f64>() / (N as f64);
params_host.iter().map(|&w| (w as f64).abs()).sum::<f64>() / total_count_f64;
let scale = mean_abs_g / mean_abs_w.max(EPS_DIV);
let true_l1 = (scale * deficit) as f32;

File diff suppressed because one or more lines are too long