feat(bf16): target forward BF16 + accessor methods
- forward_target_bf16() for target network inference path - BF16 activation pointer accessors for wiring into training path - States F32→BF16 conversion at system boundary (experience collector output) Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -606,6 +606,132 @@ impl CublasForward {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// BF16 target network forward: next_states(F32) → BF16 GemmEx → target logits(BF16).
|
||||
///
|
||||
/// Uses target-specific internal BF16 scratch buffers (`tgt_h_s1_bf16`, etc.)
|
||||
/// and writes logits to `tgt_v_logits_bf16` / `tgt_b_logits_bf16`.
|
||||
/// Does NOT save activations for backward (inference only).
|
||||
#[allow(dead_code, clippy::too_many_arguments)]
|
||||
pub fn forward_target_bf16(
|
||||
&self,
|
||||
stream: &Arc<CudaStream>,
|
||||
next_states_f32: &CudaSlice<f32>, // [B, SD] F32 from experience collector
|
||||
bf16_w_ptrs: &[u64; 20], // BF16 target weight pointers
|
||||
) -> Result<(), MLError> {
|
||||
let b = self.batch_size;
|
||||
let n_states = b * self.state_dim;
|
||||
|
||||
// ── Step 1: Convert states F32 → BF16 (reuse online states_bf16 as scratch) ──
|
||||
{
|
||||
let src_ptr = raw_f32_ptr(next_states_f32, stream);
|
||||
let dst_ptr = raw_u16_ptr(&self.states_bf16, stream);
|
||||
let n_i32 = n_states as i32;
|
||||
let blocks = ((n_states + 255) / 256) as u32;
|
||||
unsafe {
|
||||
stream
|
||||
.launch_builder(&self.f32_to_bf16_kernel)
|
||||
.arg(&src_ptr)
|
||||
.arg(&dst_ptr)
|
||||
.arg(&n_i32)
|
||||
.launch(LaunchConfig {
|
||||
grid_dim: (blocks, 1, 1),
|
||||
block_dim: (256, 1, 1),
|
||||
shared_mem_bytes: 0,
|
||||
})
|
||||
.map_err(|e| MLError::ModelError(format!("f32_to_bf16 tgt_states: {e}")))?;
|
||||
}
|
||||
}
|
||||
|
||||
let states_bf16_ptr = raw_u16_ptr(&self.states_bf16, stream);
|
||||
let tgt_h_s1_ptr = raw_u16_ptr(&self.tgt_h_s1_bf16, stream);
|
||||
let tgt_h_s2_ptr = raw_u16_ptr(&self.tgt_h_s2_bf16, stream);
|
||||
let tgt_h_v_ptr = raw_u16_ptr(&self.tgt_h_v_bf16, stream);
|
||||
|
||||
// ── Step 2: Shared trunk layer 1 ────
|
||||
self.gemmex_bf16(bf16_w_ptrs[0], states_bf16_ptr, tgt_h_s1_ptr,
|
||||
self.shared_h1, b, self.state_dim, "bf16_tgt_h_s1")?;
|
||||
self.launch_add_bias_relu_bf16_raw(stream, tgt_h_s1_ptr, bf16_w_ptrs[1], self.shared_h1, b)?;
|
||||
|
||||
// ── Step 3: Shared trunk layer 2 ────
|
||||
self.gemmex_bf16(bf16_w_ptrs[2], tgt_h_s1_ptr, tgt_h_s2_ptr,
|
||||
self.shared_h2, b, self.shared_h1, "bf16_tgt_h_s2")?;
|
||||
self.launch_add_bias_relu_bf16_raw(stream, tgt_h_s2_ptr, bf16_w_ptrs[3], self.shared_h2, b)?;
|
||||
|
||||
// ── Step 4: Value head layer 1 ────
|
||||
self.gemmex_bf16(bf16_w_ptrs[4], tgt_h_s2_ptr, tgt_h_v_ptr,
|
||||
self.value_h, b, self.shared_h2, "bf16_tgt_h_v")?;
|
||||
self.launch_add_bias_relu_bf16_raw(stream, tgt_h_v_ptr, bf16_w_ptrs[5], self.value_h, b)?;
|
||||
|
||||
// ── Step 5: Value head layer 2 (logits, no ReLU) ────
|
||||
let tgt_v_logits_ptr = raw_u16_ptr(&self.tgt_v_logits_bf16, stream);
|
||||
self.gemmex_bf16(bf16_w_ptrs[6], tgt_h_v_ptr, tgt_v_logits_ptr,
|
||||
self.num_atoms, b, self.value_h, "bf16_tgt_v_logits")?;
|
||||
self.launch_add_bias_bf16_raw(stream, tgt_v_logits_ptr, bf16_w_ptrs[7], self.num_atoms, b)?;
|
||||
|
||||
// ── Step 6: Branch heads ────
|
||||
let branch_sizes = [self.branch_0_size, self.branch_1_size, self.branch_2_size];
|
||||
let branch_w_base = [8_usize, 12, 16];
|
||||
let tgt_h_b_ptr = raw_u16_ptr(&self.tgt_h_b_bf16, stream);
|
||||
let tgt_b_logits_ptr = raw_u16_ptr(&self.tgt_b_logits_bf16, stream);
|
||||
let na = self.num_atoms;
|
||||
|
||||
let mut logit_byte_offset: u64 = 0;
|
||||
for d in 0..3 {
|
||||
let n_d = branch_sizes[d];
|
||||
let w_fc_idx = branch_w_base[d];
|
||||
let b_fc_idx = w_fc_idx + 1;
|
||||
let w_out_idx = w_fc_idx + 2;
|
||||
let b_out_idx = w_fc_idx + 3;
|
||||
|
||||
self.gemmex_bf16(bf16_w_ptrs[w_fc_idx], tgt_h_s2_ptr, tgt_h_b_ptr,
|
||||
self.adv_h, b, self.shared_h2, "bf16_tgt_h_bd")?;
|
||||
self.launch_add_bias_relu_bf16_raw(stream, tgt_h_b_ptr, bf16_w_ptrs[b_fc_idx], self.adv_h, b)?;
|
||||
|
||||
let adv_out_ptr = tgt_b_logits_ptr + logit_byte_offset;
|
||||
self.gemmex_bf16(bf16_w_ptrs[w_out_idx], tgt_h_b_ptr, adv_out_ptr,
|
||||
n_d * na, b, self.adv_h, "bf16_tgt_adv_logits")?;
|
||||
self.launch_add_bias_bf16_raw(stream, adv_out_ptr, bf16_w_ptrs[b_out_idx], n_d * na, b)?;
|
||||
|
||||
logit_byte_offset += (b * n_d * na * std::mem::size_of::<u16>()) as u64;
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// ══════════════════════════════════════════════════════════════════════════
|
||||
// BF16 logit buffer accessors (for loss kernel wiring)
|
||||
// ══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
/// Raw pointer to online value logits BF16 buffer [B, NA].
|
||||
pub fn v_logits_bf16_ptr(&self, stream: &Arc<CudaStream>) -> u64 {
|
||||
raw_u16_ptr(&self.v_logits_bf16, stream)
|
||||
}
|
||||
|
||||
/// Raw pointer to online branch logits BF16 buffer [B, (B0+B1+B2)*NA].
|
||||
pub fn b_logits_bf16_ptr(&self, stream: &Arc<CudaStream>) -> u64 {
|
||||
raw_u16_ptr(&self.b_logits_bf16, stream)
|
||||
}
|
||||
|
||||
/// Raw pointer to target value logits BF16 buffer [B, NA].
|
||||
pub fn tgt_v_logits_bf16_ptr(&self, stream: &Arc<CudaStream>) -> u64 {
|
||||
raw_u16_ptr(&self.tgt_v_logits_bf16, stream)
|
||||
}
|
||||
|
||||
/// Raw pointer to target branch logits BF16 buffer [B, (B0+B1+B2)*NA].
|
||||
pub fn tgt_b_logits_bf16_ptr(&self, stream: &Arc<CudaStream>) -> u64 {
|
||||
raw_u16_ptr(&self.tgt_b_logits_bf16, stream)
|
||||
}
|
||||
|
||||
/// Reference to online value logits BF16 buffer.
|
||||
pub fn v_logits_bf16_buf(&self) -> &CudaSlice<u16> {
|
||||
&self.v_logits_bf16
|
||||
}
|
||||
|
||||
/// Reference to online branch logits BF16 buffer.
|
||||
pub fn b_logits_bf16_buf(&self) -> &CudaSlice<u16> {
|
||||
&self.b_logits_bf16
|
||||
}
|
||||
|
||||
/// Run value head forward only: h_s2 → W_v1 → ReLU → W_v2 → v_logits.
|
||||
///
|
||||
/// Used by ensemble heads (1..K-1) to compute per-head value logits
|
||||
|
||||
Reference in New Issue
Block a user