diff --git a/crates/ml/src/cuda_pipeline/batched_forward.rs b/crates/ml/src/cuda_pipeline/batched_forward.rs index d83d1b28b..3be500d57 100644 --- a/crates/ml/src/cuda_pipeline/batched_forward.rs +++ b/crates/ml/src/cuda_pipeline/batched_forward.rs @@ -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, + next_states_f32: &CudaSlice, // [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::()) 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) -> 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) -> 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) -> 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) -> 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 { + &self.v_logits_bf16 + } + + /// Reference to online branch logits BF16 buffer. + pub fn b_logits_bf16_buf(&self) -> &CudaSlice { + &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