Skip to main content

entrenar/train/transformer_trainer/
cuda_trainer.rs

1//! GPU-resident transformer trainer (ALB-040)
2//!
3//! Wires the existing `CudaTransformerBlock` forward/backward/optimizer_step
4//! into the pretraining path. Follows the proven `classify_pipeline.rs` pattern.
5//!
6//! # Architecture
7//!
8//! ```text
9//! CudaTransformerTrainer
10//! ├── model: Transformer                 (CPU — embed + save)
11//! ├── cuda_trainer: CudaTrainer          (GPU device context)
12//! ├── cuda_blocks: Vec<CudaBlock>            (fp32 or NF4)
13//! ├── cuda_grad_workspace: CudaGradWorkspace
14//! ├── gpu_training: GpuPretrainState     (layer_inputs, grad bufs, opt states)
15//! ├── lm_head_weight_gpu: GpuBuffer      (V × H on GPU)
16//! ├── lm_head_grad_gpu: GpuBuffer        (V × H gradient scratch)
17//! ├── lm_head_m/v: GpuBuffer             (AdamW moment states)
18//! └── config: TransformerTrainConfig
19//! ```
20//!
21//! # Transfer budget (C-GPUTRAIN-002, updated KAIZEN-050/052)
22//!
23//! 1 PCIe transfer per training step (+ tiny control transfers):
24//! 1. H2D: hidden states after embedding (seq×H×4 bytes)
25//! 2. H2D: target_ids for fused cross-entropy (seq×4 bytes — ~512B)
26//! 3. D2H: loss_partials from fused cross-entropy (seq×4 bytes — ~512B)
27//!
28//! Eliminated by KAIZEN-050:
29//! - D2H logits (was seq×V×4 = 77.8MB for Qwen3-4B)
30//! - H2D grad_logits (was seq×V×4 = 77.8MB)
31//!
32//! Eliminated by KAIZEN-052:
33//! - grad_gpu buffer allocation (was seq×V×4 = 77.8MB per step)
34
35#[cfg(feature = "cuda")]
36use trueno_gpu::driver::{CudaStream, GpuBuffer};
37
38#[cfg(feature = "cuda")]
39use crate::autograd::cuda_backward::{gemm_backward_a, gemm_backward_b, rms_norm_backward};
40#[cfg(feature = "cuda")]
41use crate::autograd::cuda_forward::{
42    gemm_forward, pre_warm_forward_kernels, rms_norm_forward, rms_norm_forward_with_eps,
43};
44#[cfg(feature = "cuda")]
45use crate::autograd::cuda_optim::{
46    adamw_step_cuda, clip_scale_reduce_cuda, fused_cross_entropy_cuda, gradient_clip_cuda,
47    gradient_clip_gpu_scale_cuda, squared_sum_collect, squared_sum_cuda, squared_sum_launch_cuda,
48    squared_sum_launch_into, FusedClipState,
49};
50#[cfg(feature = "cuda")]
51use crate::autograd::cuda_training::{cuda_training_available, CudaTrainer};
52#[cfg(feature = "cuda")]
53use crate::autograd::precision::GradScaler;
54#[cfg(feature = "cuda")]
55use crate::autograd::Tensor;
56#[cfg(feature = "cuda")]
57use crate::io::{save_model, Model, ModelFormat, ModelMetadata, SaveConfig};
58#[cfg(feature = "cuda")]
59use crate::optim::{AdamW, Optimizer};
60#[cfg(feature = "cuda")]
61use crate::train::MetricsTracker;
62#[cfg(feature = "cuda")]
63use crate::transformer::{
64    CudaBlock, CudaBlockScratch, CudaGradWorkspace, CudaLoraGradWorkspace, CudaTransformerBlock,
65    GpuBlockOptimizerState, GpuLoraOptimizerState, Transformer,
66};
67
68#[cfg(feature = "cuda")]
69use super::batch::LMBatch;
70#[cfg(feature = "cuda")]
71use super::config::TransformerTrainConfig;
72#[cfg(feature = "cuda")]
73use super::step_profiler::StepProfiler;
74
75/// Compute gradient L2 norm of the shared workspace via GPU reduction (KAIZEN-054).
76///
77/// Uses `squared_sum_cuda` per buffer (~1KB D2H each) instead of downloading entire
78/// gradient buffers to CPU (was 58 MB+ per block, disabled in ALB-067).
79///
80/// Free function to avoid borrow conflicts with `&mut self`.
81#[cfg(feature = "cuda")]
82fn compute_workspace_clip_scale_gpu(
83    ws: &CudaGradWorkspace,
84    max_norm: f32,
85    stream: &CudaStream,
86) -> (f32, f32) {
87    use crate::autograd::cuda_optim::PendingSquaredSum;
88
89    let all_bufs: [&GpuBuffer<f32>; 9] = [
90        &ws.grad_w_q,
91        &ws.grad_w_k,
92        &ws.grad_w_v,
93        &ws.grad_w_o,
94        &ws.grad_gate,
95        &ws.grad_up,
96        &ws.grad_down,
97        &ws.grad_input_norm,
98        &ws.grad_post_attn_norm,
99    ];
100
101    // KAIZEN-055: Launch all 9 squared_sum kernels back-to-back without syncing.
102    // Single sync after all launches — reduces 9 pipeline flushes to 1 per block.
103    let mut pending: Vec<PendingSquaredSum> = Vec::with_capacity(9);
104    for buf in &all_bufs {
105        let n = buf.len() as u32;
106        if n == 0 {
107            continue;
108        }
109        if let Ok(p) = squared_sum_launch_cuda(buf, n, stream) {
110            pending.push(p);
111        }
112    }
113
114    // Single sync point for all 9 kernel launches.
115    if stream.synchronize().is_err() {
116        return (1.0, 0.0);
117    }
118
119    // Collect results: download partial sums (~1KB each) and combine.
120    // C-CLIP-001: squared_sum_collect returns sum(x²) = ||g||².
121    // Accumulate directly — do NOT re-square (entrenar#311 fix).
122    let mut total_sq = 0.0f64;
123    for p in &pending {
124        if let Ok(sq_norm) = squared_sum_collect(p) {
125            total_sq += f64::from(sq_norm); // sq_norm is already ||g||²
126        }
127    }
128
129    let grad_norm = total_sq.sqrt() as f32; // L2 norm = sqrt(sum of squared norms)
130    let scale = if grad_norm > max_norm { max_norm / grad_norm } else { 1.0 };
131    (scale, grad_norm)
132}
133
134/// Clip all gradient buffers in the shared workspace using GPU-computed L2 norm (KAIZEN-054).
135///
136/// R-004: Returns pre-clip gradient L2 norm for observability logging.
137#[cfg(feature = "cuda")]
138fn clip_workspace_gradients(ws: &mut CudaGradWorkspace, max_norm: f32, stream: &CudaStream) -> f32 {
139    let (scale, grad_norm) = compute_workspace_clip_scale_gpu(ws, max_norm, stream);
140    if (scale - 1.0).abs() < 1e-7 {
141        return grad_norm;
142    }
143
144    let n_wq = ws.grad_w_q.len() as u32;
145    let n_wk = ws.grad_w_k.len() as u32;
146    let n_wv = ws.grad_w_v.len() as u32;
147    let n_wo = ws.grad_w_o.len() as u32;
148    let n_gate = ws.grad_gate.len() as u32;
149    let n_up = ws.grad_up.len() as u32;
150    let n_down = ws.grad_down.len() as u32;
151    let n_inorm = ws.grad_input_norm.len() as u32;
152    let n_panorm = ws.grad_post_attn_norm.len() as u32;
153
154    let _ = gradient_clip_cuda(&mut ws.grad_w_q, scale, n_wq, stream);
155    let _ = gradient_clip_cuda(&mut ws.grad_w_k, scale, n_wk, stream);
156    let _ = gradient_clip_cuda(&mut ws.grad_w_v, scale, n_wv, stream);
157    let _ = gradient_clip_cuda(&mut ws.grad_w_o, scale, n_wo, stream);
158    let _ = gradient_clip_cuda(&mut ws.grad_gate, scale, n_gate, stream);
159    let _ = gradient_clip_cuda(&mut ws.grad_up, scale, n_up, stream);
160    let _ = gradient_clip_cuda(&mut ws.grad_down, scale, n_down, stream);
161    let _ = gradient_clip_cuda(&mut ws.grad_input_norm, scale, n_inorm, stream);
162    let _ = gradient_clip_cuda(&mut ws.grad_post_attn_norm, scale, n_panorm, stream);
163    grad_norm
164}
165
166/// ALB-078: Fused gradient clipping — entire pipeline stays on GPU.
167///
168/// Replaces `clip_workspace_gradients` by eliminating the stream.synchronize()
169/// and D2H partial-sum download. All computation happens on GPU:
170///
171/// 1. 9× SquaredSumKernel → write partials to pre-allocated contiguous buffer
172/// 2. 1× ClipScaleReduceKernel → reduce partials, compute scale on GPU
173/// 3. 9× GradientClipGpuScaleKernel → read scale from GPU, apply to gradients
174///
175/// Zero sync points, zero D2H transfers per block.
176#[cfg(feature = "cuda")]
177fn fused_clip_workspace_gradients(
178    ws: &mut CudaGradWorkspace,
179    max_norm: f32,
180    state: &FusedClipState,
181    stream: &CudaStream,
182) {
183    let all_bufs: [&GpuBuffer<f32>; 9] = [
184        &ws.grad_w_q,
185        &ws.grad_w_k,
186        &ws.grad_w_v,
187        &ws.grad_w_o,
188        &ws.grad_gate,
189        &ws.grad_up,
190        &ws.grad_down,
191        &ws.grad_input_norm,
192        &ws.grad_post_attn_norm,
193    ];
194
195    // Phase 1: Launch 9 squared_sum kernels into contiguous partials buffer.
196    // Each writes to state.partials_buf at its pre-computed offset.
197    for (i, buf) in all_bufs.iter().enumerate() {
198        let n = buf.len() as u32;
199        if n == 0 {
200            continue;
201        }
202        let output_ptr = state.partials_buf.as_ptr() + u64::from(state.offsets[i]) * 4;
203        let _ = squared_sum_launch_into(buf, n, output_ptr, stream);
204    }
205
206    // Phase 2: Reduce all partials and compute clip_scale on GPU.
207    // Stream ordering guarantees all squared_sum kernels complete before this runs.
208    let _ = clip_scale_reduce_cuda(
209        &state.partials_buf,
210        state.total_partials,
211        max_norm,
212        &state.scale_buf,
213        stream,
214    );
215
216    // Phase 3: Apply clip scale to all 9 gradient buffers.
217    // Scale is read from GPU memory — no D2H needed.
218    let scale_ptr = state.scale_buf.as_ptr(); // output[0] = clip_scale
219    let mut all_bufs_mut: [&mut GpuBuffer<f32>; 9] = [
220        &mut ws.grad_w_q,
221        &mut ws.grad_w_k,
222        &mut ws.grad_w_v,
223        &mut ws.grad_w_o,
224        &mut ws.grad_gate,
225        &mut ws.grad_up,
226        &mut ws.grad_down,
227        &mut ws.grad_input_norm,
228        &mut ws.grad_post_attn_norm,
229    ];
230    for buf in &mut all_bufs_mut {
231        let n = buf.len() as u32;
232        if n == 0 {
233            continue;
234        }
235        let _ = gradient_clip_gpu_scale_cuda(buf, scale_ptr, n, stream);
236    }
237}
238
239/// R-004: Compute gradient L2 norm without clipping (for observability only).
240///
241/// Uses GPU reduction (KAIZEN-054). Only ~9KB D2H per call.
242#[cfg(feature = "cuda")]
243#[allow(dead_code)]
244fn compute_workspace_grad_norm(ws: &CudaGradWorkspace, stream: &CudaStream) -> f32 {
245    let (_, norm) = compute_workspace_clip_scale_gpu(ws, f32::MAX, stream);
246    norm
247}
248
249/// ALB-072: Unscale all gradient buffers in the shared workspace by `inv_scale`.
250///
251/// In fp16 AMP, the fused cross-entropy kernel multiplies loss_scale into the
252/// gradient output. All subsequent backward gradients carry this scaling. The
253/// GPU block optimizer (AdamW) must receive unscaled gradients — otherwise the
254/// second moment `v` overflows f32, producing NaN in early layers.
255///
256/// This is the GPU-side equivalent of `GradScaler::unscale_and_check()` used
257/// for CPU embedding gradients.
258#[cfg(feature = "cuda")]
259#[allow(dead_code)]
260fn unscale_workspace_gradients(ws: &mut CudaGradWorkspace, inv_scale: f32, stream: &CudaStream) {
261    if (inv_scale - 1.0).abs() < 1e-7 {
262        return;
263    }
264
265    let n_wq = ws.grad_w_q.len() as u32;
266    let n_wk = ws.grad_w_k.len() as u32;
267    let n_wv = ws.grad_w_v.len() as u32;
268    let n_wo = ws.grad_w_o.len() as u32;
269    let n_gate = ws.grad_gate.len() as u32;
270    let n_up = ws.grad_up.len() as u32;
271    let n_down = ws.grad_down.len() as u32;
272    let n_inorm = ws.grad_input_norm.len() as u32;
273    let n_panorm = ws.grad_post_attn_norm.len() as u32;
274
275    let _ = gradient_clip_cuda(&mut ws.grad_w_q, inv_scale, n_wq, stream);
276    let _ = gradient_clip_cuda(&mut ws.grad_w_k, inv_scale, n_wk, stream);
277    let _ = gradient_clip_cuda(&mut ws.grad_w_v, inv_scale, n_wv, stream);
278    let _ = gradient_clip_cuda(&mut ws.grad_w_o, inv_scale, n_wo, stream);
279    let _ = gradient_clip_cuda(&mut ws.grad_gate, inv_scale, n_gate, stream);
280    let _ = gradient_clip_cuda(&mut ws.grad_up, inv_scale, n_up, stream);
281    let _ = gradient_clip_cuda(&mut ws.grad_down, inv_scale, n_down, stream);
282    let _ = gradient_clip_cuda(&mut ws.grad_input_norm, inv_scale, n_inorm, stream);
283    let _ = gradient_clip_cuda(&mut ws.grad_post_attn_norm, inv_scale, n_panorm, stream);
284}
285
286/// GPU-resident training state for pretraining.
287///
288/// # Contract (C-GPUTRAIN-001)
289///
290/// - `layer_inputs.len() == num_layers`
291/// - All buffers preallocated at init; zero GPU allocations during training
292/// - `step` increments monotonically
293#[cfg(feature = "cuda")]
294struct GpuPretrainState {
295    /// Saved layer inputs for backward [num_layers][seq_len * hidden_size]
296    layer_inputs: Vec<GpuBuffer<f32>>,
297    /// Which layer inputs were saved during forward (activation checkpointing).
298    /// When checkpointing is enabled, only checkpoint boundary layers are saved.
299    /// Non-saved layers are recomputed from the nearest checkpoint before backward.
300    saved_layer_mask: Vec<bool>,
301    /// Temporary buffer for activation recomputation [seq_len * hidden_size].
302    /// Used as the initial input when recomputing from a checkpoint boundary.
303    /// Only allocated when activation checkpointing is enabled.
304    recompute_buf: Option<GpuBuffer<f32>>,
305    /// Final RMSNorm weight on GPU [hidden_size]
306    final_norm_weight: GpuBuffer<f32>,
307    /// Final block output (pre-norm) for RMSNorm backward [seq_len * hidden_size]
308    blocks_output: GpuBuffer<f32>,
309    /// Alternating gradient buffer A [seq_len * hidden_size]
310    grad_buf_a: GpuBuffer<f32>,
311    /// Alternating gradient buffer B [seq_len * hidden_size]
312    grad_buf_b: GpuBuffer<f32>,
313    /// Gradient for final norm weight [hidden_size]
314    grad_final_norm_weight: GpuBuffer<f32>,
315    /// RMSNorm output buffer (reused each step) [seq_len * hidden_size]
316    norm_output: GpuBuffer<f32>,
317    /// Logits buffer (reused each step) [seq_len * vocab_size]
318    logits_buf: GpuBuffer<f32>,
319    /// LM head gradient buffer [seq_len * hidden_size] (grad w.r.t. normed hidden)
320    lm_head_grad_hidden: GpuBuffer<f32>,
321    /// Per-block optimizer states
322    optimizer_states: Vec<GpuBlockOptimizerState>,
323    /// Optimizer step counter
324    step: u32,
325}
326
327/// GPU-resident transformer trainer for pretraining.
328///
329/// Uses `CudaTransformerBlock` forward/backward/optimizer_step on GPU,
330/// keeping only embedding lookup and cross-entropy loss on CPU.
331///
332/// # Contract (C-GPUTRAIN-002)
333///
334/// - Exactly 3 PCIe transfers per training step
335/// - Graceful fallback to CPU `TransformerTrainer` on any CUDA failure
336/// - Weight sync via `sync_weights_to_cpu()` before save
337#[cfg(feature = "cuda")]
338pub struct CudaTransformerTrainer {
339    /// CPU model (for embedding, saving, fallback)
340    model: Transformer,
341    /// CUDA device context
342    cuda_trainer: CudaTrainer,
343    /// GPU-resident transformer blocks (fp32 or NF4 via CudaBlock enum)
344    cuda_blocks: Vec<CudaBlock>,
345    /// Shared gradient workspace (one set, reused across layers; fp32 path only)
346    cuda_grad_workspace: CudaGradWorkspace,
347    /// ENT-263: Shared scratch for NF4 blocks (C-SCRATCH-001). None when fp32.
348    nf4_shared_scratch: Option<CudaBlockScratch>,
349    /// ENT-263: Shared LoRA gradient workspace for NF4 backward. None when fp32.
350    nf4_lora_grad_workspace: Option<CudaLoraGradWorkspace>,
351    /// ENT-263: Per-block LoRA optimizer states for NF4. None when fp32.
352    nf4_lora_optimizer_states: Option<Vec<GpuLoraOptimizerState>>,
353    /// GPU training state (layer inputs, grad bufs, optimizer states)
354    gpu_training: GpuPretrainState,
355    /// LM head weight on GPU [vocab_size * hidden_size]
356    lm_head_weight_gpu: GpuBuffer<f32>,
357    /// LM head weight gradient on GPU [vocab_size * hidden_size]
358    lm_head_grad_gpu: GpuBuffer<f32>,
359    /// LM head AdamW first moment [vocab_size * hidden_size]
360    lm_head_m: GpuBuffer<f32>,
361    /// LM head AdamW second moment [vocab_size * hidden_size]
362    lm_head_v: GpuBuffer<f32>,
363    /// Final norm weight AdamW first moment [hidden_size]
364    final_norm_m: GpuBuffer<f32>,
365    /// Final norm weight AdamW second moment [hidden_size]
366    final_norm_v: GpuBuffer<f32>,
367    /// CPU optimizer for embedding weights only
368    embed_optimizer: AdamW,
369    /// Training configuration
370    config: TransformerTrainConfig,
371    /// Metrics tracker
372    pub metrics: MetricsTracker,
373    /// Current optimizer step
374    step: usize,
375    /// Accumulated loss (for gradient accumulation)
376    accumulated_loss: f32,
377    /// Accumulated batch count
378    accumulated_batches: usize,
379    /// R-004: Last observed LM head gradient L2 norm (proxy for global grad norm)
380    last_grad_norm: f32,
381    /// R-040: Last observed embedding activation gradient L2 norm
382    last_embed_grad_norm: f32,
383    /// R-038: Per-block gradient accumulation for true multi-step gradient accumulation.
384    /// Only allocated when accumulation_steps > 1. CPU-side buffers (~335 MB for 350M).
385    grad_accum: Option<super::grad_accumulator::PerBlockGradientAccumulator>,
386    /// ALB-091: GPU-resident gradient accumulation (replaces CPU accum when available).
387    /// Eliminates 24 × ga stream.synchronize() + D2H transfers per optimizer step.
388    gpu_grad_accum: Option<super::gpu_grad_accumulator::GpuGradientAccumulator>,
389    /// R-002: Gradient scaler for mixed-precision training.
390    /// For BF16: no-op (scale=1.0, dynamic=false).
391    /// For FP16: dynamic loss scaling to prevent gradient underflow.
392    grad_scaler: GradScaler,
393    /// KAIZEN-047: Per-step wall-clock profiler.
394    /// Reports timing breakdown for each training phase.
395    profiler: StepProfiler,
396    /// KAIZEN-053: Pre-allocated forward scratch buffers [max_seq_len * hidden_size].
397    /// Reused every step — eliminates 2 × cuMemAlloc/Free per training step.
398    fwd_scratch_a: GpuBuffer<f32>,
399    fwd_scratch_b: GpuBuffer<f32>,
400    /// KAIZEN-056: Pre-allocated CPU staging buffer for H2D hidden state upload.
401    /// Eliminates vec![0.0; max_seq_len * hidden_size] allocation per step.
402    h2d_staging: Vec<f32>,
403    /// KAIZEN-059: Pre-allocated CPU staging buffer for D2H gradient downloads
404    /// during gradient accumulation. Sized to max(h*intermediate, vocab*h).
405    /// Eliminates ~15GB of per-step heap churn (36 × vec![0.0; h*i] + vec![0.0; vocab*h]
406    /// per micro-batch × accumulation_steps).
407    d2h_staging: Vec<f32>,
408    /// ALB-078: Pre-allocated state for fused gradient clipping pipeline.
409    /// Eliminates 24 stream.synchronize() calls per step.
410    fused_clip: Option<FusedClipState>,
411    /// Pre-allocated host zero buffer for zeroing final norm grad [hidden_size].
412    /// BatchedRmsNormBackwardKernel accumulates grad_gamma via atomicAdd,
413    /// so the buffer must be zeroed before each rms_norm_backward call.
414    final_norm_zero_buf: Vec<f32>,
415}
416
417#[cfg(feature = "cuda")]
418impl CudaTransformerTrainer {
419    /// Create a new GPU-resident trainer.
420    ///
421    /// # Errors
422    ///
423    /// Returns `Err` if CUDA initialization, kernel pre-warming, or block upload fails.
424    /// Caller should fall back to CPU `TransformerTrainer` on error.
425    pub fn new(config: TransformerTrainConfig) -> crate::Result<Self> {
426        let model = Transformer::new(&config.model_config);
427        Self::with_model(model, config)
428    }
429
430    /// ALB-089: Load SafeTensors checkpoint for GPU inference (forward-only).
431    ///
432    /// Creates a `CudaTransformerTrainer` in inference mode. The optimizer
433    /// state is allocated (wasteful but simple), but `forward_logits()` only
434    /// uses the forward path. Call `forward_logits(&tokens)` to generate.
435    ///
436    /// # Arguments
437    /// * `checkpoint_dir` - Directory containing model.safetensors + config.json
438    /// * `model_config` - Transformer architecture config
439    ///
440    /// # Errors
441    ///
442    /// Returns `Err` if SafeTensors loading or CUDA initialization fails.
443    pub fn for_inference(
444        checkpoint_dir: impl AsRef<std::path::Path>,
445        model_config: crate::transformer::TransformerConfig,
446    ) -> crate::Result<Self> {
447        let dir = checkpoint_dir.as_ref();
448
449        // ALB-089: Try APR format first (our native checkpoint format), then SafeTensors
450        let model = if let Some((Some(m), _step)) =
451            crate::config::try_load_apr_for_inference(dir, &model_config)
452        {
453            m
454        } else {
455            Transformer::from_safetensors(dir, &model_config)?
456        };
457
458        let mut config = TransformerTrainConfig::new(model_config);
459        config.max_seq_len = config.model_config.max_position_embeddings;
460        Self::with_model(model, config)
461    }
462
463    /// Create a GPU-resident trainer from an existing model.
464    ///
465    /// # Errors
466    ///
467    /// Returns `Err` if CUDA initialization fails.
468    pub fn with_model(model: Transformer, config: TransformerTrainConfig) -> crate::Result<Self> {
469        if !cuda_training_available() {
470            return Err(crate::error::Error::ConfigError("CUDA not available".into()));
471        }
472
473        let mc = &config.model_config;
474        let max_seq_len = config.max_seq_len;
475        let hidden_size = mc.hidden_size;
476        let vocab_size = mc.vocab_size;
477        let num_layers = mc.num_hidden_layers;
478
479        // Step 1: Create CUDA trainer (initializes kernel caches)
480        let cuda_trainer = CudaTrainer::new().map_err(|e| {
481            crate::error::Error::ConfigError(format!("CUDA trainer init failed: {e:?}"))
482        })?;
483
484        println!(
485            "  GPU: {} ({:.1} GB)",
486            cuda_trainer.device_name(),
487            cuda_trainer.total_memory() as f64 / 1e9
488        );
489
490        let ctx = cuda_trainer.context().clone();
491        let stream = cuda_trainer.stream();
492
493        // Step 2: Pre-warm forward kernels (C-PREWARM-001)
494        // Must happen before block upload — JIT compilation needs free VRAM
495        pre_warm_forward_kernels(
496            hidden_size,
497            mc.intermediate_size,
498            mc.num_attention_heads,
499            mc.num_kv_heads,
500            mc.head_dim(),
501            max_seq_len,
502        )
503        .map_err(|e| crate::error::Error::ConfigError(format!("Kernel pre-warm failed: {e:?}")))?;
504
505        // Step 2a: Pre-warm backward kernels (trueno#200)
506        // MUST happen before any GPU work — Blackwell's cuModuleLoadData fails
507        // with ILLEGAL_ADDRESS when called during active GPU computation.
508        {
509            use crate::autograd::cuda_backward::pre_warm_lora_backward_kernels;
510            let head_dim = mc.head_dim();
511            pre_warm_lora_backward_kernels(
512                hidden_size,
513                mc.num_attention_heads * head_dim,
514                mc.num_kv_heads * head_dim,
515                max_seq_len,
516                config.lora_rank.unwrap_or(0),
517                mc.intermediate_size,
518                mc.num_attention_heads,
519                config.quantize_nf4 && config.is_lora(),
520            )
521            .map_err(|e| {
522                crate::error::Error::ConfigError(format!("Backward kernel pre-warm failed: {e:?}"))
523            })?;
524            eprintln!("  ✓ Backward kernels pre-warmed (silu_backward, rms_norm_backward, etc.)");
525        }
526
527        // Step 2b: Bind cuBLAS handles to training stream (ALB-075)
528        // Must happen after kernel cache init, before any GEMM calls.
529        if let Err(e) = crate::autograd::cuda_forward::set_forward_cublas_stream(stream) {
530            println!("[WARN] cuBLAS forward stream bind failed: {e:?} — falling back to PTX");
531        }
532        if let Err(e) = crate::autograd::cuda_backward::set_backward_cublas_stream(stream) {
533            println!("[WARN] cuBLAS backward stream bind failed: {e:?} — falling back to PTX");
534        }
535
536        // Step 3: Upload transformer blocks to GPU
537        let use_nf4 = config.quantize_nf4 && config.is_lora();
538        let cuda_blocks = Self::upload_blocks(
539            &model,
540            mc,
541            &config,
542            &ctx,
543            use_nf4,
544            num_layers,
545            hidden_size,
546            max_seq_len,
547        )?;
548
549        // Step 4: Allocate shared gradient workspace
550        let cuda_grad_workspace = CudaGradWorkspace::new(&ctx, mc).map_err(|e| {
551            crate::error::Error::ConfigError(format!("Grad workspace alloc failed: {e:?}"))
552        })?;
553
554        // Step 5: Allocate GPU training state
555        let buf_size = max_seq_len * hidden_size;
556        let logits_size = max_seq_len * vocab_size;
557
558        // Activation checkpointing: determine which layers save their inputs.
559        // Checkpoint boundary layers (every segment_size layers) are always saved.
560        // Non-boundary layers are recomputed from the nearest checkpoint during backward.
561        let checkpointing = config.checkpoint_config.enabled;
562        let segment_size = if checkpointing {
563            let ns = config.checkpoint_config.num_segments.max(1);
564            num_layers.div_ceil(ns)
565        } else {
566            1 // Every layer is a checkpoint (no recomputation)
567        };
568        let saved_layer_mask: Vec<bool> =
569            (0..num_layers).map(|i| !checkpointing || i % segment_size == 0).collect();
570
571        let mut layer_inputs = Vec::with_capacity(num_layers);
572        for _ in 0..num_layers {
573            layer_inputs.push(GpuBuffer::new(&ctx, buf_size).map_err(|e| {
574                crate::error::Error::ConfigError(format!("Layer input alloc failed: {e:?}"))
575            })?);
576        }
577
578        // Allocate recompute buffer if checkpointing is enabled
579        let recompute_buf = if checkpointing {
580            Some(GpuBuffer::new(&ctx, buf_size).map_err(|e| {
581                crate::error::Error::ConfigError(format!("Recompute buf alloc failed: {e:?}"))
582            })?)
583        } else {
584            None
585        };
586
587        if checkpointing {
588            let saved_count = saved_layer_mask.iter().filter(|&&x| x).count();
589            println!(
590                "  ✓ Activation checkpointing: {} segments, saving {}/{} layer inputs",
591                config.checkpoint_config.num_segments, saved_count, num_layers
592            );
593        }
594
595        // Upload final RMSNorm weight
596        let norm_slice = model.norm.weight.data().as_slice().expect("contiguous");
597        let final_norm_weight = GpuBuffer::from_host(&ctx, norm_slice).map_err(|e| {
598            crate::error::Error::ConfigError(format!("Norm weight upload failed: {e:?}"))
599        })?;
600
601        let blocks_output = GpuBuffer::new(&ctx, buf_size).map_err(|e| {
602            crate::error::Error::ConfigError(format!("Blocks output alloc failed: {e:?}"))
603        })?;
604        let grad_buf_a = GpuBuffer::new(&ctx, buf_size).map_err(|e| {
605            crate::error::Error::ConfigError(format!("Grad buf A alloc failed: {e:?}"))
606        })?;
607        let grad_buf_b = GpuBuffer::new(&ctx, buf_size).map_err(|e| {
608            crate::error::Error::ConfigError(format!("Grad buf B alloc failed: {e:?}"))
609        })?;
610        let grad_final_norm_weight = GpuBuffer::new(&ctx, hidden_size).map_err(|e| {
611            crate::error::Error::ConfigError(format!("Grad norm alloc failed: {e:?}"))
612        })?;
613        let norm_output = GpuBuffer::new(&ctx, buf_size).map_err(|e| {
614            crate::error::Error::ConfigError(format!("Norm output alloc failed: {e:?}"))
615        })?;
616        let logits_buf = GpuBuffer::new(&ctx, logits_size).map_err(|e| {
617            crate::error::Error::ConfigError(format!("Logits buf alloc failed: {e:?}"))
618        })?;
619        let lm_head_grad_hidden = GpuBuffer::new(&ctx, buf_size).map_err(|e| {
620            crate::error::Error::ConfigError(format!("LM head grad alloc failed: {e:?}"))
621        })?;
622
623        // Initialize per-block optimizer states (fp32 path only; NF4 uses LoRA states)
624        let mut optimizer_states = Vec::new();
625        if !use_nf4 {
626            optimizer_states.reserve(num_layers);
627            for (i, block) in cuda_blocks.iter().enumerate() {
628                optimizer_states.push(block.init_optimizer_state().map_err(|e| {
629                    crate::error::Error::ConfigError(format!("Block {i} opt state failed: {e:?}"))
630                })?);
631            }
632        }
633
634        let gpu_training = GpuPretrainState {
635            layer_inputs,
636            saved_layer_mask,
637            recompute_buf,
638            final_norm_weight,
639            blocks_output,
640            grad_buf_a,
641            grad_buf_b,
642            grad_final_norm_weight,
643            norm_output,
644            logits_buf,
645            lm_head_grad_hidden,
646            optimizer_states,
647            step: 0,
648        };
649
650        // Step 6: Upload LM head weight to GPU
651        // Use tied weights (embed_tokens.weight) or separate lm_head
652        let lm_head_data = model.lm_head.as_ref().unwrap_or(&model.embed_tokens.weight).data();
653        let lm_head_slice = lm_head_data.as_slice().expect("contiguous");
654        let lm_head_weight_gpu = GpuBuffer::from_host(&ctx, lm_head_slice).map_err(|e| {
655            crate::error::Error::ConfigError(format!("LM head upload failed: {e:?}"))
656        })?;
657        let lm_head_grad_gpu = GpuBuffer::new(&ctx, vocab_size * hidden_size).map_err(|e| {
658            crate::error::Error::ConfigError(format!("LM head grad alloc failed: {e:?}"))
659        })?;
660        // CRITICAL: Must zero-initialize m/v buffers. GpuBuffer::new() does NOT
661        // zero memory (cuMemAlloc returns uninitialized VRAM).
662        let lm_head_m = GpuBuffer::from_host(&ctx, &vec![0.0f32; vocab_size * hidden_size])
663            .map_err(|e| {
664                crate::error::Error::ConfigError(format!("LM head m alloc failed: {e:?}"))
665            })?;
666        let lm_head_v = GpuBuffer::from_host(&ctx, &vec![0.0f32; vocab_size * hidden_size])
667            .map_err(|e| {
668                crate::error::Error::ConfigError(format!("LM head v alloc failed: {e:?}"))
669            })?;
670
671        // Final norm optimizer states
672        let final_norm_m = GpuBuffer::from_host(&ctx, &vec![0.0f32; hidden_size]).map_err(|e| {
673            crate::error::Error::ConfigError(format!("Final norm m alloc failed: {e:?}"))
674        })?;
675        let final_norm_v = GpuBuffer::from_host(&ctx, &vec![0.0f32; hidden_size]).map_err(|e| {
676            crate::error::Error::ConfigError(format!("Final norm v alloc failed: {e:?}"))
677        })?;
678
679        // KAIZEN-053: Pre-allocate forward scratch buffers (reused every step)
680        let buf_size = max_seq_len * hidden_size;
681        let fwd_scratch_a = GpuBuffer::new(&ctx, buf_size).map_err(|e| {
682            crate::error::Error::ConfigError(format!("Fwd scratch A alloc failed: {e:?}"))
683        })?;
684        let fwd_scratch_b = GpuBuffer::new(&ctx, buf_size).map_err(|e| {
685            crate::error::Error::ConfigError(format!("Fwd scratch B alloc failed: {e:?}"))
686        })?;
687
688        // Sync to ensure all uploads completed
689        stream
690            .synchronize()
691            .map_err(|e| crate::error::Error::ConfigError(format!("Stream sync failed: {e:?}")))?;
692
693        println!(
694            "  ✓ GPU training state allocated (LM head: {:.1} MB)",
695            (vocab_size * hidden_size * 4) as f64 / 1e6
696        );
697
698        // ENT-263: Allocate NF4 infrastructure (shared scratch, LoRA grad workspace, optimizer states)
699        let (nf4_shared_scratch, nf4_lora_grad_workspace, nf4_lora_optimizer_states) = if use_nf4 {
700            let lora_rank = config.lora_rank.unwrap_or(16);
701
702            // C-SCRATCH-001: Shared scratch for NF4 blocks (reused across all layers)
703            let scratch = CudaBlockScratch::new(mc, max_seq_len, &ctx, lora_rank).map_err(|e| {
704                crate::error::Error::ConfigError(format!("NF4 shared scratch alloc failed: {e:?}"))
705            })?;
706
707            // LoRA gradient workspace (shared, reused per-block like CudaGradWorkspace)
708            let grad_ws = CudaLoraGradWorkspace::new(&ctx, mc, lora_rank).map_err(|e| {
709                crate::error::Error::ConfigError(format!(
710                    "NF4 LoRA grad workspace alloc failed: {e:?}"
711                ))
712            })?;
713
714            // Per-block LoRA optimizer states
715            let mut lora_opt_states = Vec::with_capacity(num_layers);
716            for (i, block) in cuda_blocks.iter().enumerate() {
717                lora_opt_states.push(block.init_lora_optimizer_state().map_err(|e| {
718                    crate::error::Error::ConfigError(format!(
719                        "Block {i} LoRA opt state failed: {e:?}"
720                    ))
721                })?);
722            }
723
724            println!(
725                "  ✓ NF4 training infrastructure allocated (shared scratch + LoRA optimizer × {num_layers})"
726            );
727            (Some(scratch), Some(grad_ws), Some(lora_opt_states))
728        } else {
729            (None, None, None)
730        };
731
732        // KAIZEN-050: loss_fn removed — cross-entropy computed by fused GPU kernel
733        // C-EMBED-GRAD-001: CPU optimizer must match YAML hyperparams (not defaults)
734        let embed_optimizer =
735            AdamW::new(config.lr, config.beta1, config.beta2, 1e-8, config.weight_decay);
736
737        // R-038: Allocate per-block gradient accumulation buffers (CPU-side)
738        // when accumulation_steps > 1 for true gradient accumulation.
739        let grad_accum = if config.accumulation_steps > 1 {
740            let kv_hidden = mc.num_kv_heads * mc.head_dim();
741            let block_sizes =
742                super::grad_accumulator::PerBlockGradientAccumulator::compute_block_sizes(
743                    hidden_size,
744                    kv_hidden,
745                    mc.intermediate_size,
746                );
747            let accum = super::grad_accumulator::PerBlockGradientAccumulator::new(
748                num_layers,
749                block_sizes,
750                vocab_size,
751                hidden_size,
752            );
753            println!(
754                "  ✓ Gradient accumulation: {} steps, CPU buffers ({:.1} MB)",
755                config.accumulation_steps,
756                (accum
757                    .block_grads
758                    .iter()
759                    .map(super::grad_accumulator::BlockGradientSet::total_elements)
760                    .sum::<usize>()
761                    + accum.lm_head_grad.len()
762                    + accum.final_norm_grad.len()
763                    + accum.embedding_grad.len()) as f64
764                    * 4.0
765                    / 1e6,
766            );
767            Some(accum)
768        } else {
769            None
770        };
771
772        // ALB-091: GPU-resident gradient accumulation (eliminates D2H bottleneck).
773        // Falls back to CPU accum if GPU allocation fails.
774        let gpu_grad_accum = if config.accumulation_steps > 1 {
775            match super::gpu_grad_accumulator::GpuGradientAccumulator::new(&ctx, mc) {
776                Ok(accum) => {
777                    println!("  ✓ GPU gradient accumulation enabled (ALB-091)");
778                    Some(accum)
779                }
780                Err(e) => {
781                    eprintln!(
782                        "  [WARN] GPU gradient accumulation failed ({e}), using CPU fallback"
783                    );
784                    None
785                }
786            }
787        } else {
788            None
789        };
790
791        // KAIZEN-059: Pre-allocate D2H staging buffer for gradient accumulation
792        // downloads. Only needed when GPU accum is unavailable (CPU fallback path).
793        let d2h_staging = if config.accumulation_steps > 1 && gpu_grad_accum.is_none() {
794            let ws_max = hidden_size * mc.intermediate_size;
795            let lm_max = vocab_size * hidden_size;
796            vec![0.0f32; ws_max.max(lm_max)]
797        } else {
798            Vec::new()
799        };
800
801        // ALB-078: Pre-allocate fused gradient clipping state.
802        // Eliminates 24 stream syncs per step by keeping norm+clip on GPU.
803        let kv_hidden = mc.num_kv_heads * mc.head_dim();
804        let fused_clip = Self::init_fused_clip(&ctx, &config, hidden_size, kv_hidden, mc);
805
806        // R-002: Initialize gradient scaler from precision config
807        let grad_scaler = GradScaler::from_config(&config.precision_config);
808        if config.precision_config.is_mixed() {
809            println!(
810                "  ✓ Mixed precision: {} (loss scale={}, dynamic={})",
811                config.precision_config.compute_precision,
812                grad_scaler.scale(),
813                grad_scaler.is_dynamic(),
814            );
815        }
816
817        Ok(Self {
818            model,
819            cuda_trainer,
820            cuda_blocks,
821            cuda_grad_workspace,
822            nf4_shared_scratch,
823            nf4_lora_grad_workspace,
824            nf4_lora_optimizer_states,
825            gpu_training,
826            lm_head_weight_gpu,
827            lm_head_grad_gpu,
828            lm_head_m,
829            lm_head_v,
830            final_norm_m,
831            final_norm_v,
832            embed_optimizer,
833            // KAIZEN-047: Read profile_interval before moving config into struct.
834            profiler: if config.profile_interval > 0 {
835                StepProfiler::new(true, config.profile_interval)
836            } else {
837                StepProfiler::disabled()
838            },
839            config,
840            metrics: MetricsTracker::new(),
841            step: 0,
842            accumulated_loss: 0.0,
843            accumulated_batches: 0,
844            last_grad_norm: 0.0,
845            last_embed_grad_norm: 0.0,
846            grad_accum,
847            gpu_grad_accum,
848            grad_scaler,
849            fwd_scratch_a,
850            fwd_scratch_b,
851            h2d_staging: vec![0.0f32; max_seq_len * hidden_size],
852            d2h_staging,
853            fused_clip,
854            final_norm_zero_buf: vec![0.0f32; hidden_size],
855        })
856    }
857
858    /// Upload transformer blocks to GPU (NF4 or fp32 path).
859    #[allow(clippy::too_many_arguments)]
860    fn upload_blocks(
861        model: &Transformer,
862        mc: &crate::transformer::TransformerConfig,
863        config: &TransformerTrainConfig,
864        ctx: &std::sync::Arc<trueno_gpu::driver::CudaContext>,
865        use_nf4: bool,
866        num_layers: usize,
867        hidden_size: usize,
868        max_seq_len: usize,
869    ) -> crate::Result<Vec<CudaBlock>> {
870        let mut cuda_blocks: Vec<CudaBlock> = Vec::with_capacity(num_layers);
871
872        if use_nf4 {
873            let lora_rank = config.lora_rank.unwrap_or(16);
874            let lora_alpha = config.lora_alpha.unwrap_or(2.0 * lora_rank as f32);
875            let lora_scale = lora_alpha / lora_rank as f32;
876            let head_dim = mc.head_dim();
877            let q_dim = mc.num_attention_heads * head_dim;
878            let kv_hidden = mc.num_kv_heads * head_dim;
879
880            for (i, layer) in model.layers.iter().enumerate() {
881                let lora_a_q: Vec<f32> = (0..hidden_size * lora_rank)
882                    .map(|j| ((j as f32 + i as f32 * 1000.0) * 0.1).sin() * 0.01)
883                    .collect();
884                let lora_b_q = vec![0.0f32; lora_rank * q_dim];
885                let lora_a_v: Vec<f32> = (0..hidden_size * lora_rank)
886                    .map(|j| ((j as f32 + i as f32 * 2000.0 + 500.0) * 0.1).sin() * 0.01)
887                    .collect();
888                let lora_b_v = vec![0.0f32; lora_rank * kv_hidden];
889
890                let q_norm_data = layer
891                    .self_attn
892                    .q_norm
893                    .as_ref()
894                    .map(|t| t.data().as_slice().expect("contiguous q_norm").to_vec());
895                let k_norm_data = layer
896                    .self_attn
897                    .k_norm
898                    .as_ref()
899                    .map(|t| t.data().as_slice().expect("contiguous k_norm").to_vec());
900
901                // FALSIFY-CUDA-NF4-TRAIN-LOSS-PARITY-001: thread Q/K/V biases
902                let b_q_data = layer
903                    .self_attn
904                    .b_q
905                    .as_ref()
906                    .map(|t| t.data().as_slice().expect("contiguous b_q").to_vec());
907                let b_k_data = layer
908                    .self_attn
909                    .b_k
910                    .as_ref()
911                    .map(|t| t.data().as_slice().expect("contiguous b_k").to_vec());
912                let b_v_data = layer
913                    .self_attn
914                    .b_v
915                    .as_ref()
916                    .map(|t| t.data().as_slice().expect("contiguous b_v").to_vec());
917
918                let block = crate::transformer::CudaNf4TransformerBlock::new(
919                    mc,
920                    i,
921                    ctx.clone(),
922                    layer.input_norm.weight.data().as_slice().expect("contiguous"),
923                    layer.post_attn_norm.weight.data().as_slice().expect("contiguous"),
924                    layer.self_attn.w_q.data().as_slice().expect("contiguous"),
925                    layer.self_attn.w_k.data().as_slice().expect("contiguous"),
926                    layer.self_attn.w_v.data().as_slice().expect("contiguous"),
927                    layer.self_attn.w_o.data().as_slice().expect("contiguous"),
928                    layer.ffn.w_gate.data().as_slice().expect("contiguous"),
929                    layer.ffn.w_up.data().as_slice().expect("contiguous"),
930                    layer.ffn.w_down.data().as_slice().expect("contiguous"),
931                    max_seq_len,
932                    Some((&lora_a_q, &lora_b_q)),
933                    Some((&lora_a_v, &lora_b_v)),
934                    lora_scale,
935                    lora_rank,
936                    q_norm_data.as_deref(),
937                    k_norm_data.as_deref(),
938                    b_q_data.as_deref(),
939                    b_k_data.as_deref(),
940                    b_v_data.as_deref(),
941                )
942                .map_err(|e| {
943                    crate::error::Error::ConfigError(format!("NF4 block {i} upload failed: {e:?}"))
944                })?;
945                cuda_blocks.push(CudaBlock::Nf4(block));
946            }
947            println!("  ✓ {num_layers} NF4 transformer blocks uploaded (LoRA rank={lora_rank}, alpha={lora_alpha})");
948        } else {
949            for (i, layer) in model.layers.iter().enumerate() {
950                // FALSIFY-CUDA-FORWARD-PARITY-002 thread Q/K/V biases
951                // through to the CudaTransformerBlock when present
952                // (Qwen2 family use_bias=true). Pre-fix these were
953                // silently dropped → val_loss > ln(vocab) on Qwen.
954                let b_q = layer
955                    .self_attn
956                    .b_q
957                    .as_ref()
958                    .map(|t| t.data().as_slice().expect("contiguous b_q").to_vec());
959                let b_k = layer
960                    .self_attn
961                    .b_k
962                    .as_ref()
963                    .map(|t| t.data().as_slice().expect("contiguous b_k").to_vec());
964                let b_v = layer
965                    .self_attn
966                    .b_v
967                    .as_ref()
968                    .map(|t| t.data().as_slice().expect("contiguous b_v").to_vec());
969                let block = CudaTransformerBlock::new(
970                    mc,
971                    i,
972                    ctx.clone(),
973                    layer.input_norm.weight.data().as_slice().expect("contiguous"),
974                    layer.post_attn_norm.weight.data().as_slice().expect("contiguous"),
975                    layer.self_attn.w_q.data().as_slice().expect("contiguous"),
976                    layer.self_attn.w_k.data().as_slice().expect("contiguous"),
977                    layer.self_attn.w_v.data().as_slice().expect("contiguous"),
978                    layer.self_attn.w_o.data().as_slice().expect("contiguous"),
979                    layer.ffn.w_gate.data().as_slice().expect("contiguous"),
980                    layer.ffn.w_up.data().as_slice().expect("contiguous"),
981                    layer.ffn.w_down.data().as_slice().expect("contiguous"),
982                    max_seq_len,
983                    b_q.as_deref(),
984                    b_k.as_deref(),
985                    b_v.as_deref(),
986                )
987                .map_err(|e| {
988                    crate::error::Error::ConfigError(format!("Block {i} upload failed: {e:?}"))
989                })?;
990                cuda_blocks.push(CudaBlock::Fp32(block));
991            }
992            println!("  ✓ {num_layers} transformer blocks uploaded to GPU");
993        }
994
995        Ok(cuda_blocks)
996    }
997
998    /// ALB-078: Initialize fused gradient clipping state (extracted for complexity).
999    fn init_fused_clip(
1000        ctx: &std::sync::Arc<trueno_gpu::driver::CudaContext>,
1001        config: &TransformerTrainConfig,
1002        hidden_size: usize,
1003        kv_hidden: usize,
1004        mc: &crate::transformer::TransformerConfig,
1005    ) -> Option<FusedClipState> {
1006        config.base.max_grad_norm?;
1007        let grad_sizes: [u32; 9] = [
1008            (hidden_size * hidden_size) as u32,
1009            (hidden_size * kv_hidden) as u32,
1010            (hidden_size * kv_hidden) as u32,
1011            (hidden_size * hidden_size) as u32,
1012            (hidden_size * mc.intermediate_size) as u32,
1013            (hidden_size * mc.intermediate_size) as u32,
1014            (mc.intermediate_size * hidden_size) as u32,
1015            hidden_size as u32,
1016            hidden_size as u32,
1017        ];
1018        match FusedClipState::new(ctx, &grad_sizes) {
1019            Ok(state) => {
1020                println!(
1021                    "  ✓ Fused gradient clipping: {} partials ({:.1} KB)",
1022                    state.total_partials,
1023                    f64::from(state.total_partials) * 4.0 / 1024.0,
1024                );
1025                Some(state)
1026            }
1027            Err(e) => {
1028                println!("  ⚠ Fused clip alloc failed ({e:?}), using sync fallback");
1029                None
1030            }
1031        }
1032    }
1033
1034    /// Run one forward+backward step for a single sequence.
1035    ///
1036    /// # Contract (C-GPUSTEP-001)
1037    ///
1038    /// - Precondition: `input_ids.len() == target_ids.len() <= max_seq_len`
1039    /// - Postcondition: If `accumulate_only` is false, all GPU weights updated.
1040    ///   If true, gradients accumulated into CPU buffers (no weight updates).
1041    /// - Transfer count: 1 PCIe H2D + ~1KB control (KAIZEN-050, + 24×9 D2H if accumulating)
1042    fn train_step_single(
1043        &mut self,
1044        input_ids: &[u32],
1045        target_ids: &[u32],
1046        accumulate_only: bool,
1047    ) -> Option<f32> {
1048        self.profiler.begin_step();
1049        let result = self.train_step_inner(input_ids, target_ids, accumulate_only);
1050        self.profiler.finish_step();
1051        result
1052    }
1053
1054    /// Inner training step — separated so profiler always records the step.
1055    fn train_step_inner(
1056        &mut self,
1057        input_ids: &[u32],
1058        target_ids: &[u32],
1059        accumulate_only: bool,
1060    ) -> Option<f32> {
1061        let hidden_size = self.config.model_config.hidden_size;
1062        let vocab_size = self.config.model_config.vocab_size;
1063
1064        // Truncate to max_seq_len — GPU buffers are pre-allocated for this size
1065        let max_sl = self.config.max_seq_len;
1066        let input_ids = if input_ids.len() > max_sl { &input_ids[..max_sl] } else { input_ids };
1067        let target_ids = if target_ids.len() > max_sl { &target_ids[..max_sl] } else { target_ids };
1068        let seq_len = input_ids.len();
1069
1070        // Steps 1-6: GPU forward pass — logits stay GPU-resident (KAIZEN-050)
1071        // (sub-phases embed, h2d, forward, norm_lm instrumented inside gpu_forward)
1072        if self.gpu_forward(input_ids, seq_len, hidden_size, vocab_size).is_none() {
1073            eprintln!(
1074                "[train_step_inner] gpu_forward returned None (seq_len={seq_len}, \
1075                 hidden={hidden_size}, vocab={vocab_size}) — CUDA context likely poisoned"
1076            );
1077            return None;
1078        }
1079
1080        // Step 7: Fused GPU cross-entropy loss + softmax backward (KAIZEN-050)
1081        // Eliminates: logits D2H (77.8MB) + CPU softmax (40ms) + grad H2D (77.8MB)
1082        self.profiler.begin(StepProfiler::LOSS);
1083        let stream = self.cuda_trainer.stream();
1084
1085        // Compute combined scale: (1/seq_len) * (1/accum_steps)
1086        //
1087        // ALB-072: Do NOT multiply by grad_scaler.scale() here. All backward
1088        // computation uses f32 GpuBuffers — there is no fp16 gradient underflow
1089        // risk. The 65536x loss scaling caused gradient overflow in early layers
1090        // (blocks 0-1 went NaN). The GradScaler remains active for the CPU
1091        // embedding path (unscale_and_check in optimizer_step) as a safety check,
1092        // but it operates with scale=1.0 effective for GPU gradients.
1093        let mut loss_scale = 1.0 / seq_len as f32;
1094        if self.config.accumulation_steps > 1 {
1095            loss_scale /= self.config.accumulation_steps as f32;
1096        }
1097
1098        // KAIZEN-052: In-place — gradient written directly to logits_buf.
1099        let loss_val = fused_cross_entropy_cuda(
1100            &mut self.gpu_training.logits_buf,
1101            target_ids,
1102            seq_len as u32,
1103            vocab_size as u32,
1104            loss_scale,
1105            stream,
1106        )
1107        .ok()?;
1108
1109        // NaN guard (replaces logits NaN check — NaN logits → NaN loss via kernel)
1110        if !loss_val.is_finite() {
1111            return None;
1112        }
1113        self.profiler.end(StepProfiler::LOSS);
1114
1115        // Steps 8-11: GPU backward pass (with or without optimizer)
1116        // (sub-phases lm_bwd, norm_bwd, blk_bwd instrumented inside gpu_backward)
1117        // KAIZEN-050: grad_logits on GPU. KAIZEN-052: grad lives in logits_buf (in-place).
1118        //
1119        // ENT-263 fix: Capture loss regardless of backward success. The NF4 backward
1120        // path may fail (e.g., gemm_nf4_backward_a stub) but the loss was already
1121        // computed by fused_cross_entropy_cuda. Dropping the loss silently causes
1122        // loss=0.0 reporting despite valid forward passes.
1123        if let Some(grad_output_is_a) =
1124            self.gpu_backward(seq_len, hidden_size, vocab_size, accumulate_only)
1125        {
1126            // Step 12: Embedding backward (CPU scatter-add always accumulates)
1127            self.profiler.begin(StepProfiler::EMBED_BWD);
1128            self.embed_backward(input_ids, seq_len, hidden_size, vocab_size, grad_output_is_a);
1129
1130            self.profiler.end(StepProfiler::EMBED_BWD);
1131        }
1132
1133        Some(loss_val)
1134    }
1135
1136    /// GPU forward pass: embed → blocks → norm → LM head.
1137    ///
1138    /// Logits stay GPU-resident in `self.gpu_training.logits_buf` (KAIZEN-050).
1139    /// Transfers: 1 H2D (hidden states). No D2H — logits consumed by fused kernel.
1140    #[allow(unsafe_code)]
1141    fn gpu_forward(
1142        &mut self,
1143        input_ids: &[u32],
1144        seq_len: usize,
1145        hidden_size: usize,
1146        vocab_size: usize,
1147    ) -> Option<()> {
1148        contract_pre_gpu_forward!();
1149        let stream = self.cuda_trainer.stream();
1150
1151        // Embedding lookup (CPU)
1152        self.profiler.begin(StepProfiler::EMBED);
1153        let hidden = self.model.embed_tokens.forward(input_ids);
1154        let hidden_slice = hidden.data().as_slice()?;
1155        self.profiler.end(StepProfiler::EMBED);
1156
1157        // Upload hidden states to GPU (Transfer 1: H2D)
1158        // Pad to max_seq_len so D2D copies to pre-allocated layer_inputs match.
1159        // KAIZEN-053: Reuse pre-allocated scratch buffers instead of cuMemAlloc per step.
1160        // KAIZEN-056: Reuse pre-allocated h2d_staging instead of alloc per step.
1161        self.profiler.begin(StepProfiler::H2D);
1162        self.h2d_staging[..hidden_slice.len()].copy_from_slice(hidden_slice);
1163        self.h2d_staging[hidden_slice.len()..].fill(0.0);
1164        if let Err(e) = self.fwd_scratch_a.copy_from_host(&self.h2d_staging) {
1165            eprintln!("[gpu_forward] H2D copy failed: {e:?} — CUDA context may be poisoned");
1166            return None;
1167        }
1168        self.profiler.end(StepProfiler::H2D);
1169
1170        // Forward through CUDA blocks using pre-allocated ping-pong buffers.
1171        // KAIZEN-053: fwd_scratch_a/b are top-level fields (not in gpu_training)
1172        // so borrowing them doesn't conflict with gpu_training.layer_inputs.
1173        self.profiler.begin(StepProfiler::FORWARD);
1174        let mut input_is_a = true; // Track which scratch buffer is "input"
1175        for (i, block) in self.cuda_blocks.iter_mut().enumerate() {
1176            // Use raw pointers for the ping-pong to avoid borrow conflicts
1177            // with self.gpu_training.layer_inputs
1178            let (input_ptr, output_ptr): (*const GpuBuffer<f32>, *mut GpuBuffer<f32>) =
1179                if input_is_a {
1180                    (
1181                        std::ptr::from_ref(&self.fwd_scratch_a),
1182                        std::ptr::from_mut(&mut self.fwd_scratch_b),
1183                    )
1184                } else {
1185                    (
1186                        std::ptr::from_ref(&self.fwd_scratch_b),
1187                        std::ptr::from_mut(&mut self.fwd_scratch_a),
1188                    )
1189                };
1190            if self.gpu_training.saved_layer_mask[i] {
1191                // SAFETY: Both buffers are valid GPU allocations with matching max_seq_len size.
1192                // Copy completes before block.forward() reads from input (same stream ordering).
1193                unsafe {
1194                    self.gpu_training.layer_inputs[i]
1195                        .copy_from_buffer_async(&*input_ptr, stream)
1196                        .ok()?;
1197                }
1198            }
1199            // SAFETY: input_ptr and output_ptr point to disjoint fwd_scratch_{a,b}.
1200            // ENT-263: Pass shared scratch for NF4 blocks (C-SCRATCH-001).
1201            self.profiler.begin_layer();
1202            // SAFETY: `input_ptr` and `output_ptr` point to disjoint scratch buffers (distinct `fwd_scratch_{a,b}` / `layer_inputs` slots), so reborrowing one as `&` and the other as `&mut` for the GPU forward never creates aliasing references.
1203            unsafe {
1204                block
1205                    .forward(
1206                        &*input_ptr,
1207                        &mut *output_ptr,
1208                        seq_len,
1209                        stream,
1210                        self.nf4_shared_scratch.as_mut(),
1211                    )
1212                    .ok()?;
1213            }
1214            self.profiler.end_layer_fwd(i);
1215            input_is_a = !input_is_a;
1216        }
1217        self.profiler.end(StepProfiler::FORWARD);
1218
1219        // After the loop, input_is_a tells us which buffer has the final output
1220        let final_output: &GpuBuffer<f32> =
1221            if input_is_a { &self.fwd_scratch_a } else { &self.fwd_scratch_b };
1222
1223        // Save blocks output for final norm backward
1224        // SAFETY: Disjoint GPU buffers with matching max_seq_len sizes.
1225        self.profiler.begin(StepProfiler::NORM_LM);
1226        // SAFETY: stream-ordered device-to-device copy between two distinct `GpuBuffer`s of matching element length on the same context; both allocations outlive the async copy on `stream`.
1227        unsafe {
1228            self.gpu_training.blocks_output.copy_from_buffer_async(final_output, stream).ok()?;
1229        }
1230
1231        // Final RMSNorm forward (GPU)
1232        // FALSIFY-CUDA-RMSNORM-EPS-PARITY-001: thread `config.rms_norm_eps`
1233        // through so Qwen2 (1e-6) gets the right epsilon. Pre-fix the
1234        // legacy `rms_norm_forward` hardcoded 1e-5 (Llama default).
1235        rms_norm_forward_with_eps(
1236            final_output,
1237            &self.gpu_training.final_norm_weight,
1238            &mut self.gpu_training.norm_output,
1239            seq_len as u32,
1240            hidden_size as u32,
1241            self.config.model_config.rms_norm_eps,
1242            stream,
1243        )
1244        .ok()?;
1245
1246        // LM head GEMM forward (GPU)
1247        // gemm_forward treats flat (V,H) memory as (H,V) row-major, which
1248        // implicitly transposes — matching the CPU matmul's tied-weight behavior.
1249        gemm_forward(
1250            &self.gpu_training.norm_output,
1251            &self.lm_head_weight_gpu,
1252            &mut self.gpu_training.logits_buf,
1253            seq_len as u32,
1254            hidden_size as u32,
1255            vocab_size as u32,
1256            stream,
1257        )
1258        .ok()?;
1259
1260        // KAIZEN-050: Logits stay GPU-resident — no D2H transfer.
1261        // Fused cross-entropy kernel reads logits_buf directly on GPU.
1262        self.profiler.end(StepProfiler::NORM_LM);
1263
1264        Some(())
1265    }
1266
1267    /// ALB-089: Forward-only pass that returns last-position logits on CPU.
1268    ///
1269    /// Runs the same GPU forward as training but downloads only the last
1270    /// SPEC-DISTILL-001 Phase 2d (PMAT-697): forward + caller-supplied
1271    /// logit-gradient backward + optimizer step.
1272    ///
1273    /// Unlike `forward_backward_batch` (which computes the gradient from
1274    /// CE loss internally), this method takes a precomputed last-position
1275    /// logit gradient — useful for knowledge distillation where the
1276    /// gradient is computed externally as the KD logit gradient
1277    /// `α·(softmax(s) - one_hot(label)) + (1-α)·T·(softmax(s/T) - softmax(t/T))`
1278    /// (per `aprender-train-distill::kd_step::kd_logit_gradient`).
1279    ///
1280    /// Flow:
1281    /// 1. `gpu_forward(input_ids)` — produces last-position logits in
1282    ///    `gpu_training.logits_buf`.
1283    /// 2. Upload `logit_gradient` into the last-position slice of
1284    ///    `logits_buf`, OVERWRITING what gpu_forward produced (matching
1285    ///    the in-place gradient convention `fused_cross_entropy_cuda`
1286    ///    uses for the CE path).
1287    /// 3. `gpu_backward` — back-props from the uploaded gradient through
1288    ///    the transformer stack, accumulating weight gradients.
1289    /// 4. `embed_backward` — embedding-table scatter-add.
1290    ///
1291    /// **Limitations** (Phase 2d):
1292    /// - The gradient applies to the LAST POSITION only (this is the KD
1293    ///   training objective for next-token-prediction). Sequence-wise KD
1294    ///   (every position) is a Phase 2e enhancement.
1295    /// - Returns `Some(())` on success, `None` on CUDA failure. Loss is
1296    ///   not computed (caller computes from kd_loss separately).
1297    ///
1298    /// # Errors
1299    ///
1300    /// Returns `None` if `gpu_forward`, the gradient upload, or
1301    /// `gpu_backward` fails. The CUDA stream may be in a poisoned state
1302    /// after such a failure; subsequent training steps should be
1303    /// considered unreliable.
1304    pub fn forward_backward_with_grad(
1305        &mut self,
1306        input_ids: &[u32],
1307        logit_gradient: &[f32],
1308    ) -> Option<()> {
1309        let seq_len = input_ids.len();
1310        let hidden_size = self.config.model_config.hidden_size;
1311        let vocab_size = self.config.model_config.vocab_size;
1312
1313        if seq_len == 0 || seq_len > self.config.max_seq_len {
1314            return None;
1315        }
1316        if logit_gradient.len() != vocab_size {
1317            eprintln!(
1318                "[forward_backward_with_grad] gradient len {} != vocab_size {}",
1319                logit_gradient.len(),
1320                vocab_size
1321            );
1322            return None;
1323        }
1324
1325        self.gpu_forward(input_ids, seq_len, hidden_size, vocab_size)?;
1326
1327        // Upload the KD gradient into the last-position slice of logits_buf,
1328        // replacing whatever gpu_forward wrote there. This matches the
1329        // KAIZEN-052 in-place gradient convention that gpu_backward expects.
1330        let offset = (seq_len - 1) * vocab_size;
1331        self.gpu_training.logits_buf.copy_from_host_at(logit_gradient, offset).ok()?;
1332        let stream = self.cuda_trainer.stream();
1333        stream.synchronize().ok()?;
1334
1335        // Back-prop from the uploaded gradient through the transformer.
1336        // accumulate_only=false → run the optimizer step at the end.
1337        let grad_output_is_a = self.gpu_backward(seq_len, hidden_size, vocab_size, false)?;
1338        // Embedding backward (CPU scatter-add). Pre-condition: grad_output_is_a
1339        // is the buffer-flip flag from gpu_backward (per existing
1340        // `train_step_inner` pattern at line ~1108).
1341        self.embed_backward(input_ids, seq_len, hidden_size, vocab_size, grad_output_is_a);
1342
1343        Some(())
1344    }
1345
1346    /// position's logits (vocab_size floats) for token sampling. No backward
1347    /// pass, no loss computation.
1348    ///
1349    /// # Contract (C-CUDA-INF-001)
1350    ///
1351    /// - Same forward path as `gpu_forward()` — identical logits
1352    /// - Only downloads `logits[seq_len-1, :]` (128 KB for 32K vocab)
1353    /// - stream.synchronize() before D2H (C-STREAMSYNC-001)
1354    pub fn forward_logits(&mut self, input_ids: &[u32]) -> Option<Vec<f32>> {
1355        let seq_len = input_ids.len();
1356        let hidden_size = self.config.model_config.hidden_size;
1357        let vocab_size = self.config.model_config.vocab_size;
1358
1359        if seq_len == 0 || seq_len > self.config.max_seq_len {
1360            return None;
1361        }
1362
1363        // Reuse gpu_forward for the actual computation
1364        self.gpu_forward(input_ids, seq_len, hidden_size, vocab_size)?;
1365
1366        // C-STREAMSYNC-001: synchronize before D2H
1367        let stream = self.cuda_trainer.stream();
1368        stream.synchronize().ok()?;
1369
1370        // Download last position logits only: logits_buf[seq_len-1, :]
1371        let offset = (seq_len - 1) * vocab_size;
1372        let mut logits = vec![0.0f32; vocab_size];
1373        self.gpu_training.logits_buf.copy_to_host_at(&mut logits, offset).ok()?;
1374
1375        Some(logits)
1376    }
1377
1378    /// GPU backward pass with interleaved per-block optimizer step.
1379    ///
1380    /// Each block's backward writes weight gradients to the shared `CudaGradWorkspace`.
1381    /// Recompute layer inputs for a segment during backward (activation checkpointing).
1382    ///
1383    /// When checkpointing is enabled, non-checkpoint layers don't save their inputs
1384    /// during forward. Before their backward pass, we recompute from the nearest
1385    /// checkpoint by re-running forward through intermediate blocks.
1386    ///
1387    /// This recomputes the entire segment [checkpoint..=target_layer], storing
1388    /// intermediate layer_inputs so subsequent layers in the same segment don't
1389    /// need redundant recomputation.
1390    ///
1391    /// # Contract (R-021)
1392    ///
1393    /// After this call, `layer_inputs[i]` is valid for all i in [checkpoint..=target_layer].
1394    #[allow(unsafe_code)]
1395    fn recompute_segment(
1396        gpu_training: &mut GpuPretrainState,
1397        cuda_blocks: &mut [CudaBlock],
1398        nf4_shared_scratch: &mut Option<CudaBlockScratch>,
1399        target_layer: usize,
1400        seq_len: usize,
1401        stream: &CudaStream,
1402    ) -> Option<()> {
1403        // Find nearest saved checkpoint at or before target
1404        let seg_start = (0..=target_layer).rev().find(|&i| gpu_training.saved_layer_mask[i])?;
1405
1406        if seg_start == target_layer {
1407            return Some(()); // Already saved
1408        }
1409
1410        // Copy checkpoint input to recompute_buf as starting point.
1411        // SAFETY: recompute_buf and layer_inputs are disjoint allocations.
1412        let recompute_buf = gpu_training.recompute_buf.as_mut()?;
1413        // SAFETY: stream-ordered device-to-device copy between two distinct `GpuBuffer`s of matching element length on the same context; both allocations outlive the async copy on `stream`.
1414        unsafe {
1415            recompute_buf
1416                .copy_from_buffer_async(&gpu_training.layer_inputs[seg_start], stream)
1417                .ok()?;
1418        }
1419
1420        // Forward through blocks [seg_start..target_layer], saving intermediate inputs.
1421        // For block i, input → block i → output becomes input for block i+1.
1422        // We save output (= input to block i+1) in layer_inputs[i+1].
1423        //
1424        // Buffer pattern:
1425        //   i == seg_start: input = recompute_buf, output = layer_inputs[seg_start+1]
1426        //   i > seg_start:  input = layer_inputs[i], output = layer_inputs[i+1]
1427        //
1428        // SAFETY: split_at_mut ensures non-overlapping borrows of layer_inputs.
1429        // recompute_buf is separate from layer_inputs.
1430        for i in seg_start..target_layer {
1431            if i == seg_start {
1432                // Input is in recompute_buf, output goes to layer_inputs[i+1]
1433                let recompute_ptr: *const GpuBuffer<f32> = recompute_buf;
1434                let li = &mut gpu_training.layer_inputs;
1435                // SAFETY: `input_ptr` and `output_ptr` point to disjoint scratch buffers (distinct `fwd_scratch_{a,b}` / `layer_inputs` slots), so reborrowing one as `&` and the other as `&mut` for the GPU forward never creates aliasing references.
1436                unsafe {
1437                    cuda_blocks[i]
1438                        .forward(
1439                            &*recompute_ptr,
1440                            &mut li[i + 1],
1441                            seq_len,
1442                            stream,
1443                            nf4_shared_scratch.as_mut(),
1444                        )
1445                        .ok()?;
1446                }
1447            } else {
1448                // Input = layer_inputs[i], output = layer_inputs[i+1]
1449                let li = &mut gpu_training.layer_inputs;
1450                let (left, right) = li.split_at_mut(i + 1);
1451                cuda_blocks[i]
1452                    .forward(&left[i], &mut right[0], seq_len, stream, nf4_shared_scratch.as_mut())
1453                    .ok()?;
1454            }
1455        }
1456
1457        Some(())
1458    }
1459
1460    /// Since `gemm_backward_b` overwrites (not accumulates), we must run each block's
1461    /// optimizer step immediately after its backward, before the next block overwrites
1462    /// the workspace. This also enables per-block gradient clipping.
1463    ///
1464    /// When `accumulate_only` is true (R-038 gradient accumulation), the per-block
1465    /// optimizer steps are skipped and workspace gradients are downloaded to CPU-side
1466    /// `PerBlockGradientAccumulator` instead. LM head and final norm gradients are
1467    /// also downloaded and accumulated. The optimizer step is deferred until
1468    /// `gpu_optimizer_from_accum()` is called.
1469    ///
1470    /// Returns `grad_output_is_a` flag for embedding backward.
1471    /// Transfer: 0 H2D (KAIZEN-050/052: grad in logits_buf) + 24×9 D2H if accumulating.
1472    #[allow(unsafe_code)]
1473    fn gpu_backward(
1474        &mut self,
1475        seq_len: usize,
1476        hidden_size: usize,
1477        vocab_size: usize,
1478        accumulate_only: bool,
1479    ) -> Option<bool> {
1480        let stream = self.cuda_trainer.stream();
1481        let max_grad_norm = self.config.base.max_grad_norm;
1482        let lr = self.current_lr();
1483        // ALB-072: No inv_scale needed — loss_scale no longer includes grad_scaler.
1484        let beta1 = self.config.beta1;
1485        let beta2 = self.config.beta2;
1486        let weight_decay = self.config.weight_decay;
1487
1488        // KAIZEN-050: grad_logits GPU-resident. KAIZEN-052: grad lives in logits_buf (in-place).
1489        // No separate grad buffer. No GRAD_H2D transfer.
1490
1491        // LM head GEMM backward
1492        self.profiler.begin(StepProfiler::LM_BWD);
1493        gemm_backward_a(
1494            &self.gpu_training.logits_buf,
1495            &self.lm_head_weight_gpu,
1496            &mut self.gpu_training.lm_head_grad_hidden,
1497            seq_len as u32,
1498            hidden_size as u32,
1499            vocab_size as u32,
1500            stream,
1501        )
1502        .ok()?;
1503
1504        gemm_backward_b(
1505            &self.gpu_training.norm_output,
1506            &self.gpu_training.logits_buf,
1507            &mut self.lm_head_grad_gpu,
1508            seq_len as u32,
1509            hidden_size as u32,
1510            vocab_size as u32,
1511            stream,
1512        )
1513        .ok()?;
1514
1515        // Clip LM head weight gradient
1516        // KAIZEN-049: GPU norm reduction.
1517        // KAIZEN-051: No explicit sync needed — same stream ordering.
1518        // ALB-071: Always compute LM head grad norm for observability (R-004).
1519        // C-CLIP-001: squared_sum_cuda returns ||g||². Take sqrt for L2 norm (entrenar#311).
1520        let lm_sq_norm =
1521            squared_sum_cuda(&self.lm_head_grad_gpu, self.lm_head_grad_gpu.len() as u32, stream)
1522                .unwrap_or(0.0);
1523        let lm_norm = lm_sq_norm.sqrt(); // L2 norm, NOT squared
1524        self.last_grad_norm = lm_norm; // R-004: capture for observability
1525                                       // C-BACKPARITY-001: LM head gradient norm tracing (pre-clip).
1526        if std::env::var("ENTRENAR_TRACE_GRADIENTS").is_ok() {
1527            eprintln!("[grad-trace] lm_head gnorm={lm_norm:.6}");
1528            // Also trace the grad_hidden flowing to blocks
1529            let gh_sq = squared_sum_cuda(
1530                &self.gpu_training.lm_head_grad_hidden,
1531                self.gpu_training.lm_head_grad_hidden.len() as u32,
1532                stream,
1533            )
1534            .unwrap_or(0.0);
1535            eprintln!("[grad-trace] lm_head_grad_hidden gnorm={:.6}", gh_sq.sqrt());
1536        }
1537        if let Some(max_norm) = max_grad_norm {
1538            let clip_scale = if lm_norm > max_norm { max_norm / lm_norm } else { 1.0 };
1539            let n = self.lm_head_grad_gpu.len() as u32;
1540            let _ = gradient_clip_cuda(&mut self.lm_head_grad_gpu, clip_scale, n, stream);
1541        }
1542        self.profiler.end(StepProfiler::LM_BWD);
1543
1544        // Final RMSNorm backward
1545        self.profiler.begin(StepProfiler::NORM_BWD);
1546        // Zero grad_final_norm_weight before backward — kernel accumulates via atomicAdd
1547        self.gpu_training.grad_final_norm_weight.copy_from_host(&self.final_norm_zero_buf).ok()?;
1548        rms_norm_backward(
1549            &self.gpu_training.blocks_output,
1550            &self.gpu_training.final_norm_weight,
1551            &self.gpu_training.lm_head_grad_hidden,
1552            &mut self.gpu_training.grad_buf_a,
1553            &mut self.gpu_training.grad_final_norm_weight,
1554            seq_len as u32,
1555            hidden_size as u32,
1556            1e-5_f32,
1557            stream,
1558        )
1559        .ok()?;
1560
1561        // Clip final norm weight gradient
1562        // KAIZEN-051: No explicit sync needed — same stream ordering as LM head clip.
1563        if let Some(max_norm) = max_grad_norm {
1564            let (scale, _) = Self::compute_clip_scale_with_norm(
1565                &self.gpu_training.grad_final_norm_weight,
1566                max_norm,
1567                stream,
1568            );
1569            let n = self.gpu_training.grad_final_norm_weight.len() as u32;
1570            let _ =
1571                gradient_clip_cuda(&mut self.gpu_training.grad_final_norm_weight, scale, n, stream);
1572        }
1573        self.profiler.end(StepProfiler::NORM_BWD);
1574
1575        // R-038: Either accumulate non-block grads or run non-block optimizer.
1576        if accumulate_only {
1577            // ALB-091: GPU-resident accumulation (no sync, no D2H) or CPU fallback.
1578            if let Some(ref mut gpu_accum) = self.gpu_grad_accum {
1579                let _ = gpu_accum.accumulate_nonblock(
1580                    &self.lm_head_grad_gpu,
1581                    &self.gpu_training.grad_final_norm_weight,
1582                    stream,
1583                );
1584            } else {
1585                stream.synchronize().ok()?;
1586                Self::download_nonblock_grads_to_accum(
1587                    &self.lm_head_grad_gpu,
1588                    &self.gpu_training.grad_final_norm_weight,
1589                    &mut self.grad_accum,
1590                    &mut self.d2h_staging,
1591                )?;
1592            }
1593        } else {
1594            Self::run_nonblock_optimizer_step(
1595                &mut self.gpu_training,
1596                Some(&mut self.lm_head_weight_gpu),
1597                &self.lm_head_grad_gpu,
1598                &mut self.lm_head_m,
1599                &mut self.lm_head_v,
1600                &mut self.final_norm_m,
1601                &mut self.final_norm_v,
1602                lr,
1603                beta1,
1604                beta2,
1605                weight_decay,
1606                stream,
1607            );
1608        }
1609
1610        // Backward through blocks in reverse, with interleaved clip + optimizer.
1611        // Each block's backward writes weight gradients to shared CudaGradWorkspace.
1612        //
1613        // SAFETY: grad_buf_a and grad_buf_b are disjoint fields. Raw pointers
1614        // allow alternating read/write without violating aliasing rules.
1615        self.profiler.begin(StepProfiler::BLK_BWD);
1616        let grad_a_ptr: *mut GpuBuffer<f32> = &raw mut self.gpu_training.grad_buf_a;
1617        let grad_b_ptr: *mut GpuBuffer<f32> = &raw mut self.gpu_training.grad_buf_b;
1618        let mut grad_output_is_a = true;
1619        let use_nf4 = self.config.quantize_nf4 && self.config.is_lora();
1620
1621        for layer_idx in (0..self.cuda_blocks.len()).rev() {
1622            // Activation checkpointing: if this layer's input wasn't saved during
1623            // forward, recompute the segment from the nearest checkpoint.
1624            if !self.gpu_training.saved_layer_mask[layer_idx] {
1625                Self::recompute_segment(
1626                    &mut self.gpu_training,
1627                    &mut self.cuda_blocks,
1628                    &mut self.nf4_shared_scratch,
1629                    layer_idx,
1630                    seq_len,
1631                    stream,
1632                )?;
1633            }
1634
1635            // SAFETY: ping-pong double-buffering. The two raw pointers reference distinct, non-overlapping device buffers (the `_a`/`_b` scratch pair); the boolean flag picks one as `&` input and the other as `&mut` output, so the resulting references never alias the same allocation.
1636            let (grad_output, grad_input) = unsafe {
1637                if grad_output_is_a {
1638                    (&*grad_a_ptr, &mut *grad_b_ptr)
1639                } else {
1640                    (&*grad_b_ptr, &mut *grad_a_ptr)
1641                }
1642            };
1643
1644            self.profiler.begin_layer();
1645            if use_nf4 {
1646                // ENT-263: NF4 backward — LoRA gradient computation
1647                // Uses backward_nf4() which computes gradients for LoRA weights and norms only.
1648                // output_scratch reuses grad_buf_a/b as temporary storage for recomputed forward.
1649                let _output_scratch_ptr: *mut GpuBuffer<f32> = if grad_output_is_a {
1650                    grad_b_ptr // grad_input is in b, use as output_scratch too (will be overwritten)
1651                } else {
1652                    grad_a_ptr
1653                };
1654                // We need a separate output_scratch. Reuse blocks_output as scratch since
1655                // it was already consumed for norm backward above.
1656                match self.cuda_blocks[layer_idx].backward_nf4(
1657                    &self.gpu_training.layer_inputs[layer_idx],
1658                    grad_output,
1659                    grad_input,
1660                    &mut self.gpu_training.blocks_output, // reuse as output_scratch
1661                    seq_len,
1662                    stream,
1663                    self.nf4_shared_scratch.as_mut().expect("NF4 requires shared scratch"),
1664                    self.nf4_lora_grad_workspace
1665                        .as_mut()
1666                        .expect("NF4 requires LoRA grad workspace"),
1667                ) {
1668                    Ok(()) => {}
1669                    Err(e) => {
1670                        eprintln!(
1671                            "[backward_nf4] Layer {} FAILED: {:?} (seq_len={}, hidden={})",
1672                            layer_idx, e, seq_len, self.config.model_config.hidden_size
1673                        );
1674                        return None;
1675                    }
1676                }
1677
1678                // ENT-265: Clip LoRA gradients before optimizer step.
1679                // Without this, NF4 LoRA grads are unbounded — causes weight
1680                // divergence and embedding grad explosion (Run 7c: 26M at step 225).
1681                if let Some(max_norm) = max_grad_norm {
1682                    self.nf4_lora_grad_workspace
1683                        .as_mut()
1684                        .expect("NF4 requires LoRA grad ws")
1685                        .clip_gradients(max_norm, stream);
1686                }
1687
1688                // NF4 LoRA optimizer step — always runs, even during accumulation.
1689                //
1690                // BUG FIX (entrenar#264): Previously gated by `if !accumulate_only`.
1691                // Design: NF4 LoRA has ~6M params, so we scale lr by 1/accum_steps
1692                // for micro-batches instead of accumulating gradients.
1693                {
1694                    let step = self.gpu_training.step;
1695                    let effective_lr = if accumulate_only {
1696                        lr / self.config.accumulation_steps as f32
1697                    } else {
1698                        lr
1699                    };
1700                    if let Some(ref mut opt_states) = self.nf4_lora_optimizer_states {
1701                        let _ = self.cuda_blocks[layer_idx].lora_optimizer_step(
1702                            &mut opt_states[layer_idx],
1703                            step,
1704                            effective_lr,
1705                            beta1,
1706                            beta2,
1707                            1e-8,
1708                            weight_decay,
1709                            stream,
1710                            self.nf4_lora_grad_workspace
1711                                .as_ref()
1712                                .expect("NF4 requires LoRA grad ws"),
1713                        );
1714                    }
1715                }
1716            } else {
1717                // Standard fp32 backward path
1718                self.cuda_blocks[layer_idx]
1719                    .backward(
1720                        &self.gpu_training.layer_inputs[layer_idx],
1721                        grad_output,
1722                        grad_input,
1723                        seq_len,
1724                        stream,
1725                        &mut self.cuda_grad_workspace,
1726                    )
1727                    .ok()?;
1728
1729                // C-CLIP-001 / entrenar#312: DISABLED per-block gradient clipping.
1730                // Per-block clipping distorts gradient flow across layers.
1731
1732                // C-BACKPARITY-001: Per-block gradient norm tracing for parity testing.
1733                // Only runs when ENTRENAR_TRACE_GRADIENTS=1 — zero overhead in production.
1734                if std::env::var("ENTRENAR_TRACE_GRADIENTS").is_ok() {
1735                    let (_, block_gnorm) = compute_workspace_clip_scale_gpu(
1736                        &self.cuda_grad_workspace,
1737                        f32::MAX,
1738                        stream,
1739                    );
1740                    // Also trace the activation gradient (flows between blocks)
1741                    let act_sq = squared_sum_cuda(grad_input, grad_input.len() as u32, stream)
1742                        .unwrap_or(0.0);
1743                    let act_gnorm = act_sq.sqrt();
1744                    eprintln!(
1745                        "[grad-trace] block={layer_idx} weight_gnorm={block_gnorm:.6} act_gnorm={act_gnorm:.6}"
1746                    );
1747                }
1748
1749                // R-038: Either accumulate workspace grads or run optimizer per-block.
1750                if accumulate_only {
1751                    // ALB-091: GPU-resident accumulation (no sync, no D2H) or CPU fallback.
1752                    if let Some(ref mut gpu_accum) = self.gpu_grad_accum {
1753                        let _ = gpu_accum.accumulate_block(
1754                            &self.cuda_grad_workspace,
1755                            layer_idx,
1756                            stream,
1757                        );
1758                    } else {
1759                        // CPU fallback: SYNC + D2H (ALB-065 / Rule 6).
1760                        stream.synchronize().ok()?;
1761                        if let Some(accum) = &mut self.grad_accum {
1762                            Self::download_workspace_to_accum(
1763                                &self.cuda_grad_workspace,
1764                                accum,
1765                                layer_idx,
1766                                &mut self.d2h_staging,
1767                            )?;
1768                        }
1769                    }
1770                } else {
1771                    // Per-block optimizer step: consume workspace gradients before next block overwrites
1772                    let step = self.gpu_training.step;
1773                    let _ = self.cuda_blocks[layer_idx].optimizer_step(
1774                        &mut self.gpu_training.optimizer_states[layer_idx],
1775                        step,
1776                        lr,
1777                        beta1,
1778                        beta2,
1779                        1e-8,
1780                        weight_decay,
1781                        stream,
1782                        &self.cuda_grad_workspace,
1783                    );
1784                }
1785            }
1786
1787            self.profiler.end_layer_bwd(layer_idx);
1788            grad_output_is_a = !grad_output_is_a;
1789        }
1790
1791        stream.synchronize().ok()?;
1792        self.profiler.end(StepProfiler::BLK_BWD);
1793
1794        Some(grad_output_is_a)
1795    }
1796
1797    /// R-038: Download non-block (LM head + final norm) gradients to CPU accumulator.
1798    /// Static method to avoid borrow conflicts.
1799    // KAIZEN-044: Pre-allocate single buffer for LM head + norm D2H downloads.
1800    // lm_head_grad is vocab×hidden (389M elements = 1.5 GB for Qwen3-4B).
1801    // KAIZEN-059: Host buffer now passed in (d2h_staging) — zero per-call allocations.
1802    fn download_nonblock_grads_to_accum(
1803        lm_head_grad: &GpuBuffer<f32>,
1804        final_norm_grad: &GpuBuffer<f32>,
1805        grad_accum: &mut Option<super::grad_accumulator::PerBlockGradientAccumulator>,
1806        host: &mut [f32],
1807    ) -> Option<()> {
1808        let accum = grad_accum.as_mut()?;
1809
1810        let lm_slice = &mut host[..lm_head_grad.len()];
1811        lm_head_grad.copy_to_host_at(lm_slice, 0).ok()?;
1812        for (d, s) in accum.lm_head_grad.iter_mut().zip(lm_slice.iter()) {
1813            *d += s;
1814        }
1815
1816        let norm_slice = &mut host[..final_norm_grad.len()];
1817        final_norm_grad.copy_to_host_at(norm_slice, 0).ok()?;
1818        for (d, s) in accum.final_norm_grad.iter_mut().zip(norm_slice.iter()) {
1819            *d += s;
1820        }
1821        Some(())
1822    }
1823
1824    /// Run LM head + final norm optimizer step (non-accumulating path).
1825    /// Static method to avoid borrow conflicts with `stream`.
1826    #[allow(clippy::too_many_arguments)]
1827    fn run_nonblock_optimizer_step(
1828        gpu_training: &mut GpuPretrainState,
1829        lm_head_weight_gpu: Option<&mut GpuBuffer<f32>>,
1830        lm_head_grad_gpu: &GpuBuffer<f32>,
1831        lm_head_m: &mut GpuBuffer<f32>,
1832        lm_head_v: &mut GpuBuffer<f32>,
1833        final_norm_m: &mut GpuBuffer<f32>,
1834        final_norm_v: &mut GpuBuffer<f32>,
1835        lr: f32,
1836        beta1: f32,
1837        beta2: f32,
1838        weight_decay: f32,
1839        stream: &CudaStream,
1840    ) {
1841        gpu_training.step += 1;
1842        let step = gpu_training.step;
1843
1844        if let Some(lm_head_weight) = lm_head_weight_gpu {
1845            let n_lm = lm_head_weight.len() as u32;
1846            let _ = adamw_step_cuda(
1847                lm_head_weight,
1848                lm_head_grad_gpu,
1849                lm_head_m,
1850                lm_head_v,
1851                lr,
1852                beta1,
1853                beta2,
1854                1e-8,
1855                weight_decay,
1856                step,
1857                n_lm,
1858                stream,
1859            );
1860        }
1861
1862        let n_norm = gpu_training.final_norm_weight.len() as u32;
1863        let _ = adamw_step_cuda(
1864            &mut gpu_training.final_norm_weight,
1865            &gpu_training.grad_final_norm_weight,
1866            final_norm_m,
1867            final_norm_v,
1868            lr,
1869            beta1,
1870            beta2,
1871            1e-8,
1872            weight_decay,
1873            step,
1874            n_norm,
1875            stream,
1876        );
1877    }
1878
1879    /// R-038: Download shared CudaGradWorkspace to CPU per-block accumulation buffers.
1880    ///
1881    /// Static method to avoid borrow conflicts with `stream` (same pattern as
1882    /// `recompute_segment`). Must be called after stream.synchronize() (ALB-065 / Rule 6).
1883    // KAIZEN-044: Pre-allocate a single host buffer for all D2H downloads
1884    // in download_workspace_to_accum. Was allocating vec![0.0f32; len] × 9 buffers.
1885    // KAIZEN-059: Host buffer now passed in (d2h_staging) — zero per-call allocations.
1886    fn download_workspace_to_accum(
1887        ws: &CudaGradWorkspace,
1888        accum: &mut super::grad_accumulator::PerBlockGradientAccumulator,
1889        layer_idx: usize,
1890        host: &mut [f32],
1891    ) -> Option<()> {
1892        let bg = &mut accum.block_grads[layer_idx];
1893
1894        use super::grad_accumulator::component;
1895        let bufs_and_components: [(&GpuBuffer<f32>, usize); 9] = [
1896            (&ws.grad_w_q, component::W_Q),
1897            (&ws.grad_w_k, component::W_K),
1898            (&ws.grad_w_v, component::W_V),
1899            (&ws.grad_w_o, component::W_O),
1900            (&ws.grad_gate, component::GATE),
1901            (&ws.grad_up, component::UP),
1902            (&ws.grad_down, component::DOWN),
1903            (&ws.grad_input_norm, component::INPUT_NORM),
1904            (&ws.grad_post_attn_norm, component::POST_ATTN_NORM),
1905        ];
1906
1907        for (gpu_buf, comp_idx) in &bufs_and_components {
1908            let slice = &mut host[..gpu_buf.len()];
1909            gpu_buf.copy_to_host_at(slice, 0).ok()?;
1910            for (d, s) in bg.components[*comp_idx].iter_mut().zip(slice.iter()) {
1911                *d += s;
1912            }
1913        }
1914        Some(())
1915    }
1916
1917    /// R-038: Upload averaged CPU accumulation buffers to GPU workspace and run
1918    /// optimizer step for all blocks + LM head + final norm.
1919    ///
1920    /// Called once after `accumulation_steps` micro-batches have been accumulated.
1921    /// ALB-091: Run optimizer step from GPU-resident accumulated gradients.
1922    /// D2D copy accum → workspace, then run per-block optimizer. Zero accum after.
1923    fn gpu_optimizer_from_gpu_accum(&mut self) -> Option<()> {
1924        let stream = self.cuda_trainer.stream();
1925        let lr = self.current_lr();
1926        let beta1 = self.config.beta1;
1927        let beta2 = self.config.beta2;
1928        let weight_decay = self.config.weight_decay;
1929
1930        // Sync once to ensure all accumulation kernels complete
1931        stream.synchronize().ok()?;
1932
1933        self.gpu_training.step += 1;
1934        let step = self.gpu_training.step;
1935
1936        // Upload GPU accum → workspace (D2D) and run optimizer per block
1937        let gpu_accum = self.gpu_grad_accum.as_ref()?;
1938        for layer_idx in 0..self.cuda_blocks.len() {
1939            gpu_accum.upload_to_workspace(&mut self.cuda_grad_workspace, layer_idx).ok()?;
1940
1941            let _ = self.cuda_blocks[layer_idx].optimizer_step(
1942                &mut self.gpu_training.optimizer_states[layer_idx],
1943                step,
1944                lr,
1945                beta1,
1946                beta2,
1947                1e-8,
1948                weight_decay,
1949                stream,
1950                &self.cuda_grad_workspace,
1951            );
1952        }
1953
1954        // LM head: D2D copy accum → grad buffer, then optimizer step
1955        gpu_accum
1956            .upload_nonblock(
1957                &mut self.lm_head_grad_gpu,
1958                &mut self.gpu_training.grad_final_norm_weight,
1959            )
1960            .ok()?;
1961
1962        let n_lm = self.lm_head_weight_gpu.len() as u32;
1963        let _ = adamw_step_cuda(
1964            &mut self.lm_head_weight_gpu,
1965            &self.lm_head_grad_gpu,
1966            &mut self.lm_head_m,
1967            &mut self.lm_head_v,
1968            lr,
1969            beta1,
1970            beta2,
1971            1e-8,
1972            weight_decay,
1973            step,
1974            n_lm,
1975            stream,
1976        );
1977
1978        // Final norm optimizer step
1979        let n_norm = self.gpu_training.final_norm_weight.len() as u32;
1980        let _ = adamw_step_cuda(
1981            &mut self.gpu_training.final_norm_weight,
1982            &self.gpu_training.grad_final_norm_weight,
1983            &mut self.final_norm_m,
1984            &mut self.final_norm_v,
1985            lr,
1986            beta1,
1987            beta2,
1988            1e-8,
1989            weight_decay,
1990            step,
1991            n_norm,
1992            stream,
1993        );
1994
1995        stream.synchronize().ok()?;
1996
1997        // Zero accum for next window
1998        if let Some(ref mut gpu_accum) = self.gpu_grad_accum {
1999            let _ = gpu_accum.zero_all();
2000        }
2001
2002        Some(())
2003    }
2004
2005    #[allow(unsafe_code)]
2006    fn gpu_optimizer_from_accum(&mut self) -> Option<()> {
2007        let stream = self.cuda_trainer.stream();
2008        let lr = self.current_lr();
2009        let beta1 = self.config.beta1;
2010        let beta2 = self.config.beta2;
2011        let weight_decay = self.config.weight_decay;
2012
2013        // Average accumulated gradients
2014        let accum = self.grad_accum.as_mut()?;
2015        accum.average();
2016
2017        // Jidoka: check for NaN/Inf before applying
2018        if accum.has_non_finite() {
2019            println!("[WARN] R-038: NaN/Inf in accumulated gradients, skipping optimizer step");
2020            accum.zero_all();
2021            return Some(());
2022        }
2023
2024        self.gpu_training.step += 1;
2025        let step = self.gpu_training.step;
2026
2027        // Upload accumulated gradients and run optimizer for each block
2028        use super::grad_accumulator::component;
2029        for layer_idx in 0..self.cuda_blocks.len() {
2030            let bg = &accum.block_grads[layer_idx];
2031
2032            // Upload accumulated gradients to shared workspace
2033            // SAFETY: async host-to-device copies within the training stream; host buffers
2034            // (bg.components) are stable for the duration of the stream operations.
2035            unsafe {
2036                self.cuda_grad_workspace
2037                    .grad_w_q
2038                    .copy_from_host_async(&bg.components[component::W_Q], stream)
2039                    .ok()?;
2040                self.cuda_grad_workspace
2041                    .grad_w_k
2042                    .copy_from_host_async(&bg.components[component::W_K], stream)
2043                    .ok()?;
2044                self.cuda_grad_workspace
2045                    .grad_w_v
2046                    .copy_from_host_async(&bg.components[component::W_V], stream)
2047                    .ok()?;
2048                self.cuda_grad_workspace
2049                    .grad_w_o
2050                    .copy_from_host_async(&bg.components[component::W_O], stream)
2051                    .ok()?;
2052                self.cuda_grad_workspace
2053                    .grad_gate
2054                    .copy_from_host_async(&bg.components[component::GATE], stream)
2055                    .ok()?;
2056                self.cuda_grad_workspace
2057                    .grad_up
2058                    .copy_from_host_async(&bg.components[component::UP], stream)
2059                    .ok()?;
2060                self.cuda_grad_workspace
2061                    .grad_down
2062                    .copy_from_host_async(&bg.components[component::DOWN], stream)
2063                    .ok()?;
2064                self.cuda_grad_workspace
2065                    .grad_input_norm
2066                    .copy_from_host_async(&bg.components[component::INPUT_NORM], stream)
2067                    .ok()?;
2068                self.cuda_grad_workspace
2069                    .grad_post_attn_norm
2070                    .copy_from_host_async(&bg.components[component::POST_ATTN_NORM], stream)
2071                    .ok()?;
2072            }
2073
2074            // Run optimizer step with uploaded averaged gradients
2075            let _ = self.cuda_blocks[layer_idx].optimizer_step(
2076                &mut self.gpu_training.optimizer_states[layer_idx],
2077                step,
2078                lr,
2079                beta1,
2080                beta2,
2081                1e-8,
2082                weight_decay,
2083                stream,
2084                &self.cuda_grad_workspace,
2085            );
2086        }
2087
2088        // Upload accumulated LM head gradients and run AdamW step
2089        // entrenar#314: Skip GPU LM head optimizer for tied weights.
2090        // SAFETY: async host-to-device copy; host buffer (accum.lm_head_grad) is stable.
2091        unsafe {
2092            self.lm_head_grad_gpu.copy_from_host_async(&accum.lm_head_grad, stream).ok()?;
2093        }
2094        let n_lm = self.lm_head_weight_gpu.len() as u32;
2095        let _ = adamw_step_cuda(
2096            &mut self.lm_head_weight_gpu,
2097            &self.lm_head_grad_gpu,
2098            &mut self.lm_head_m,
2099            &mut self.lm_head_v,
2100            lr,
2101            beta1,
2102            beta2,
2103            1e-8,
2104            weight_decay,
2105            step,
2106            n_lm,
2107            stream,
2108        );
2109
2110        // Upload accumulated final norm gradients and run AdamW step
2111        // SAFETY: async host-to-device copy; host buffer (accum.final_norm_grad) is stable.
2112        unsafe {
2113            self.gpu_training
2114                .grad_final_norm_weight
2115                .copy_from_host_async(&accum.final_norm_grad, stream)
2116                .ok()?;
2117        }
2118        let n_norm = self.gpu_training.final_norm_weight.len() as u32;
2119        let _ = adamw_step_cuda(
2120            &mut self.gpu_training.final_norm_weight,
2121            &self.gpu_training.grad_final_norm_weight,
2122            &mut self.final_norm_m,
2123            &mut self.final_norm_v,
2124            lr,
2125            beta1,
2126            beta2,
2127            1e-8,
2128            weight_decay,
2129            step,
2130            n_norm,
2131            stream,
2132        );
2133
2134        stream.synchronize().ok()?;
2135
2136        // Zero accum for next window
2137        accum.zero_all();
2138        Some(())
2139    }
2140
2141    /// Compute gradient L2 norm via GPU reduction kernel (KAIZEN-049).
2142    ///
2143    /// Runs `SquaredSumKernel` on GPU, downloads only `num_blocks` partial sums (~1KB)
2144    /// instead of the full buffer (128MB for lm_head). Falls back to CPU download on error.
2145    ///
2146    /// # Contract (C-CLIPNORM-GPU-001)
2147    ///
2148    /// - **Precondition**: `buf.len() > 0`, stream is synchronized with prior kernel
2149    /// - **Postcondition**: `grad_norm ≈ sqrt(sum(buf[i]^2))`, `scale = min(1, max_norm/norm)`
2150    /// - **Transfer**: ~1KB D2H (num_blocks × 4B) vs n×4B (128MB for 32M elements)
2151    ///
2152    /// R-004: Returns `(clip_scale, grad_norm)` for observability.
2153    fn compute_clip_scale_with_norm(
2154        buf: &GpuBuffer<f32>,
2155        max_norm: f32,
2156        stream: &CudaStream,
2157    ) -> (f32, f32) {
2158        let n = buf.len() as u32;
2159        // Try GPU reduction first — ~1KB D2H instead of n×4 bytes
2160        let grad_norm = match squared_sum_cuda(buf, n, stream) {
2161            Ok(norm) => norm,
2162            Err(_) => {
2163                // Fallback: full D2H (original path)
2164                let mut host = vec![0.0f32; buf.len()];
2165                if buf.copy_to_host_at(&mut host, 0).is_err() {
2166                    return (1.0, 0.0);
2167                }
2168                let sq_sum: f64 = host.iter().map(|&x| f64::from(x) * f64::from(x)).sum();
2169                sq_sum.sqrt() as f32
2170            }
2171        };
2172        let scale = if grad_norm > max_norm { max_norm / grad_norm } else { 1.0 };
2173        (scale, grad_norm)
2174    }
2175
2176    /// Download embedding gradient from GPU, clip, and scatter-add into CPU weight.
2177    ///
2178    /// # Contract (C-EMBED-GRAD-001)
2179    ///
2180    /// The activation gradient from block[0]'s backward is unclipped (per-block clipping
2181    /// only applies to weight gradients in the shared workspace). For deep networks with
2182    /// random init, this gradient can overflow f32, producing NaN in the CPU AdamW.
2183    /// We clip the activation gradient to max_grad_norm before scatter-adding.
2184    #[allow(unsafe_code)]
2185    fn embed_backward(
2186        &mut self,
2187        input_ids: &[u32],
2188        _seq_len: usize,
2189        hidden_size: usize,
2190        vocab_size: usize,
2191        grad_output_is_a: bool,
2192    ) -> Option<()> {
2193        // The final backward output is in whichever buffer was last written
2194        let grad_a_ptr: *const GpuBuffer<f32> = &raw const self.gpu_training.grad_buf_a;
2195        let grad_b_ptr: *const GpuBuffer<f32> = &raw const self.gpu_training.grad_buf_b;
2196        // SAFETY: ping-pong double-buffering. The two raw pointers reference distinct, non-overlapping device buffers (the `_a`/`_b` scratch pair); the boolean flag picks one as `&` input and the other as `&mut` output, so the resulting references never alias the same allocation.
2197        let embed_grad_buf = unsafe {
2198            if grad_output_is_a {
2199                &*grad_a_ptr
2200            } else {
2201                &*grad_b_ptr
2202            }
2203        };
2204        let mut embed_grad_data = self.cuda_trainer.download(embed_grad_buf).ok()?;
2205
2206        // C-EMBED-GRAD-001: ALWAYS clip activation gradient before scatter-add.
2207        // Without this, 24-layer random-init backward amplifies gradients to ~1e35,
2208        // which overflows the CPU AdamW's second moment buffer.
2209        //
2210        // ALB-071: Decoupled from general grad_clip config. Embed activation gradient
2211        // clipping is a SAFETY constraint (prevents NaN), not a training hyperparameter.
2212        // Uses dedicated max_embed_grad_norm (default 1.0) independent of weight grad_clip.
2213        let embed_clip_norm = self.config.base.max_grad_norm.unwrap_or(1.0);
2214        {
2215            let sq_sum: f64 = embed_grad_data.iter().map(|&x| f64::from(x) * f64::from(x)).sum();
2216            let grad_norm = sq_sum.sqrt() as f32;
2217            self.last_embed_grad_norm = grad_norm; // R-040: per-parameter-group tracking
2218            if grad_norm > embed_clip_norm {
2219                let scale = embed_clip_norm / grad_norm;
2220                for g in &mut embed_grad_data {
2221                    *g *= scale;
2222                }
2223            }
2224        }
2225
2226        // KAIZEN-048: In-place scatter-add via grad_cell().borrow_mut().
2227        // Before: 3 × 128MB clones per step (grad() deep-copies Array1).
2228        // After: zero clones — mutate existing gradient buffer directly.
2229        let embed_weight = &mut self.model.embed_tokens.weight;
2230        let grad_cell = embed_weight.grad_cell();
2231        let mut grad_ref = grad_cell.borrow_mut();
2232        if grad_ref.is_none() {
2233            *grad_ref = Some(ndarray::Array1::zeros(embed_weight.len()));
2234        }
2235        if let Some(grad) = grad_ref.as_mut() {
2236            for (pos, &token_id) in input_ids.iter().enumerate() {
2237                let tid = token_id as usize;
2238                if tid < vocab_size {
2239                    let src = pos * hidden_size;
2240                    let dst = tid * hidden_size;
2241                    for h in 0..hidden_size {
2242                        grad[dst + h] += embed_grad_data[src + h];
2243                    }
2244                }
2245            }
2246        }
2247        Some(())
2248    }
2249
2250    /// Apply optimizer step to CPU embedding and update metrics.
2251    ///
2252    /// GPU block optimizer steps now run interleaved with backward in `gpu_backward()`.
2253    /// LM head and final norm optimizer steps also run in `gpu_backward()`.
2254    /// This method handles only CPU embedding and bookkeeping.
2255    fn optimizer_step(&mut self) {
2256        // ALB-072: Gradients are no longer scaled by grad_scaler (loss_scale excludes
2257        // grad_scaler.scale()). All backward computation uses f32 — no fp16 underflow
2258        // risk. Skip unscaling; just update scaler as successful.
2259        self.grad_scaler.update(true);
2260
2261        // ALB-079: Sync CPU embedding optimizer lr with cosine schedule
2262        self.embed_optimizer.set_lr(self.current_lr());
2263        // CPU optimizer step for embedding weight
2264        let mut embed_params = vec![&mut self.model.embed_tokens.weight];
2265        self.embed_optimizer.step_refs(&mut embed_params);
2266
2267        self.step += 1;
2268        self.metrics.losses.push(self.accumulated_loss);
2269        self.metrics.increment_step();
2270
2271        self.accumulated_loss = 0.0;
2272        self.accumulated_batches = 0;
2273    }
2274
2275    /// Process a batch (forward + backward + optimizer step with accumulation).
2276    ///
2277    /// R-038: When `accumulation_steps > 1`, runs forward+backward without optimizer
2278    /// for each micro-batch, downloading per-block weight gradients to CPU-side
2279    /// `PerBlockGradientAccumulator`. After `accumulation_steps` batches, averages
2280    /// the accumulated gradients, uploads them to GPU, and runs a single optimizer step.
2281    ///
2282    /// When `accumulation_steps == 1` (default), runs forward+backward+optimizer
2283    /// immediately per sequence (original behavior).
2284    ///
2285    /// Returns average loss for the batch.
2286    pub fn train_batch(&mut self, batch: &LMBatch) -> f32 {
2287        if batch.batch_size == 0 {
2288            return 0.0;
2289        }
2290
2291        let accumulating = self.grad_accum.is_some() || self.gpu_grad_accum.is_some();
2292
2293        if self.accumulated_batches == 0 {
2294            // Zero embedding gradients at start of accumulation window
2295            self.embed_optimizer.zero_grad_refs(&mut vec![&mut self.model.embed_tokens.weight]);
2296        }
2297
2298        let mut total_loss = 0.0;
2299        let mut valid_count = 0;
2300
2301        for i in 0..batch.batch_size {
2302            let Some(input_ids) = batch.get_input(i) else {
2303                continue;
2304            };
2305            let Some(target_ids) = batch.get_target(i) else {
2306                continue;
2307            };
2308
2309            // R-038: When accumulating, run backward without optimizer (accumulate_only=true).
2310            // Gradients are downloaded to CPU per-block accum buffers. Embedding grads are
2311            // scatter-added normally (they're already on CPU).
2312            if let Some(loss) = self.train_step_single(input_ids, target_ids, accumulating) {
2313                total_loss += loss;
2314                valid_count += 1;
2315                if accumulating {
2316                    if let Some(accum) = &mut self.gpu_grad_accum {
2317                        accum.accumulated_count += 1;
2318                    } else if let Some(accum) = &mut self.grad_accum {
2319                        accum.accumulated_count += 1;
2320                    }
2321                }
2322            }
2323        }
2324
2325        let avg_loss = if valid_count > 0 { total_loss / valid_count as f32 } else { 0.0 };
2326
2327        // Debug: help diagnose loss=0.0 when gradients are non-zero
2328        if avg_loss == 0.0 && valid_count > 0 {
2329            eprintln!(
2330                "[train_batch DEBUG] avg_loss=0.0 but valid_count={}, total_loss={}, batch_size={}",
2331                valid_count, total_loss, batch.batch_size
2332            );
2333        }
2334
2335        self.accumulated_loss += avg_loss / self.config.accumulation_steps as f32;
2336        self.accumulated_batches += 1;
2337
2338        if self.accumulated_batches >= self.config.accumulation_steps {
2339            if accumulating {
2340                // ALB-091: Prefer GPU-resident accum path (zero D2H), fall back to CPU.
2341                if self.gpu_grad_accum.is_some() {
2342                    self.gpu_optimizer_from_gpu_accum();
2343                } else {
2344                    self.gpu_optimizer_from_accum();
2345                }
2346            }
2347            self.optimizer_step();
2348        }
2349
2350        avg_loss
2351    }
2352
2353    /// R-005: Evaluate a batch without backward pass or weight updates.
2354    /// Returns average cross-entropy loss, or 0.0 if no valid items.
2355    /// KAIZEN-050: Uses fused GPU cross-entropy (no logits D2H).
2356    pub fn eval_batch(&mut self, batch: &LMBatch) -> f32 {
2357        let hidden_size = self.config.model_config.hidden_size;
2358        let vocab_size = self.config.model_config.vocab_size;
2359        let max_sl = self.config.max_seq_len;
2360        let mut total_loss = 0.0;
2361        let mut valid_count = 0;
2362        for i in 0..batch.batch_size {
2363            if let Some(loss) = self.eval_single_sequence(batch, i, max_sl, hidden_size, vocab_size)
2364            {
2365                total_loss += loss;
2366                valid_count += 1;
2367            }
2368        }
2369        if valid_count > 0 {
2370            total_loss / valid_count as f32
2371        } else {
2372            0.0
2373        }
2374    }
2375
2376    /// Evaluate a single sequence from a batch. Returns None if invalid.
2377    fn eval_single_sequence(
2378        &mut self,
2379        batch: &LMBatch,
2380        i: usize,
2381        max_sl: usize,
2382        hidden_size: usize,
2383        vocab_size: usize,
2384    ) -> Option<f32> {
2385        let input_ids = batch.get_input(i)?;
2386        let target_ids = batch.get_target(i)?;
2387        // Truncate to max_seq_len — GPU buffers are pre-allocated for this size
2388        let input_ids = if input_ids.len() > max_sl { &input_ids[..max_sl] } else { input_ids };
2389        let target_ids = if target_ids.len() > max_sl { &target_ids[..max_sl] } else { target_ids };
2390        let seq_len = input_ids.len();
2391        self.gpu_forward(input_ids, seq_len, hidden_size, vocab_size)?;
2392        let stream = self.cuda_trainer.stream();
2393        let scale = 1.0 / seq_len as f32;
2394        let loss = fused_cross_entropy_cuda(
2395            &mut self.gpu_training.logits_buf,
2396            target_ids,
2397            seq_len as u32,
2398            vocab_size as u32,
2399            scale,
2400            stream,
2401        )
2402        .ok()?;
2403        if loss.is_finite() {
2404            Some(loss)
2405        } else {
2406            None
2407        }
2408    }
2409
2410    /// Train for one epoch over batches.
2411    pub fn train_epoch(&mut self, batches: &[LMBatch]) -> f32 {
2412        self.train_epoch_with_callback(batches, |_, _, _| {})
2413    }
2414
2415    /// Train for one epoch with a per-step callback.
2416    ///
2417    /// Stops early if `max_steps` is set and reached.
2418    pub fn train_epoch_with_callback<F>(&mut self, batches: &[LMBatch], mut on_batch: F) -> f32
2419    where
2420        F: FnMut(usize, f32, &Self),
2421    {
2422        if batches.is_empty() {
2423            return 0.0;
2424        }
2425
2426        let mut total_loss = 0.0;
2427        let mut batches_processed = 0;
2428
2429        for (i, batch) in batches.iter().enumerate() {
2430            if let Some(max) = self.config.max_steps {
2431                if self.step >= max {
2432                    break;
2433                }
2434            }
2435
2436            let batch_loss = self.train_batch(batch);
2437            total_loss += batch_loss;
2438            batches_processed += 1;
2439            on_batch(i, batch_loss, self);
2440        }
2441
2442        // KAIZEN-047: Print profiler summary at end of epoch
2443        if self.profiler.is_enabled() && self.profiler.step_count() > 0 {
2444            self.profiler.print_report();
2445        }
2446
2447        total_loss / batches_processed.max(1) as f32
2448    }
2449
2450    // --- DDP (data-parallel) support methods ---
2451
2452    /// Ensure the per-block gradient accumulator exists.
2453    ///
2454    /// For DDP, we always need accumulation buffers (even with accumulation_steps=1)
2455    /// because gradients must be downloaded to CPU for AllReduce before optimizer step.
2456    pub(crate) fn ensure_grad_accum(&mut self) {
2457        if self.grad_accum.is_some() {
2458            return;
2459        }
2460        let mc = &self.config.model_config;
2461        let hidden_size = mc.hidden_size;
2462        let kv_hidden = mc.num_kv_heads * mc.head_dim();
2463        let block_sizes = super::grad_accumulator::PerBlockGradientAccumulator::compute_block_sizes(
2464            hidden_size,
2465            kv_hidden,
2466            mc.intermediate_size,
2467        );
2468        self.grad_accum = Some(super::grad_accumulator::PerBlockGradientAccumulator::new(
2469            self.cuda_blocks.len(),
2470            block_sizes,
2471            mc.vocab_size,
2472            hidden_size,
2473        ));
2474    }
2475
2476    /// Forward + backward for one batch, always accumulating (no optimizer step).
2477    ///
2478    /// Used by `DistributedCudaTrainer` to compute local gradients before AllReduce.
2479    /// Returns average loss for the batch.
2480    pub(crate) fn forward_backward_batch(&mut self, batch: &LMBatch) -> f32 {
2481        if batch.batch_size == 0 {
2482            return 0.0;
2483        }
2484
2485        if self.accumulated_batches == 0 {
2486            self.embed_optimizer.zero_grad_refs(&mut vec![&mut self.model.embed_tokens.weight]);
2487        }
2488
2489        let mut total_loss = 0.0;
2490        let mut valid_count = 0;
2491
2492        for i in 0..batch.batch_size {
2493            let Some(input_ids) = batch.get_input(i) else { continue };
2494            let Some(target_ids) = batch.get_target(i) else { continue };
2495
2496            // Always accumulate_only=true: gradients go to CPU accum buffers
2497            if let Some(loss) = self.train_step_single(input_ids, target_ids, true) {
2498                total_loss += loss;
2499                valid_count += 1;
2500                if let Some(accum) = &mut self.grad_accum {
2501                    accum.accumulated_count += 1;
2502                }
2503            }
2504        }
2505
2506        if valid_count > 0 {
2507            total_loss / valid_count as f32
2508        } else {
2509            0.0
2510        }
2511    }
2512
2513    /// Apply DDP-averaged gradients: upload to GPU and run optimizer step.
2514    ///
2515    /// Called after AllReduce has written averaged gradients into the grad_accum.
2516    /// Runs gpu_optimizer_from_accum() for blocks + LM head + final norm,
2517    /// then optimizer_step() for embedding.
2518    pub(crate) fn apply_ddp_gradients(&mut self) {
2519        self.accumulated_loss = 0.0;
2520        self.accumulated_batches = 0;
2521        self.gpu_optimizer_from_accum();
2522        self.optimizer_step();
2523    }
2524
2525    /// Get a reference to the gradient accumulator (for DDP AllReduce).
2526    pub(crate) fn grad_accum_ref(
2527        &self,
2528    ) -> Option<&super::grad_accumulator::PerBlockGradientAccumulator> {
2529        self.grad_accum.as_ref()
2530    }
2531
2532    /// Get a mutable reference to the gradient accumulator (for DDP AllReduce).
2533    pub(crate) fn grad_accum_mut(
2534        &mut self,
2535    ) -> Option<&mut super::grad_accumulator::PerBlockGradientAccumulator> {
2536        self.grad_accum.as_mut()
2537    }
2538
2539    /// Get the training config.
2540    pub(crate) fn config(&self) -> &TransformerTrainConfig {
2541        &self.config
2542    }
2543
2544    /// Get CPU embedding gradient as flat Vec for AllReduce.
2545    pub(crate) fn embed_grad_vec(&self) -> Option<Vec<f32>> {
2546        self.model.embed_tokens.weight.grad().map(|g| g.to_vec())
2547    }
2548
2549    /// Set CPU embedding gradient from AllReduced flat Vec.
2550    pub(crate) fn set_embed_grad(&mut self, grad: Vec<f32>) {
2551        self.model.embed_tokens.weight.set_grad(ndarray::Array1::from(grad));
2552    }
2553
2554    /// Returns true if max_steps has been reached.
2555    pub fn reached_max_steps(&self) -> bool {
2556        self.config.max_steps.is_some_and(|max| self.step >= max)
2557    }
2558
2559    /// Get current step count.
2560    pub fn step(&self) -> usize {
2561        self.step
2562    }
2563
2564    /// Set initial step for resume from checkpoint.
2565    ///
2566    /// Updates both the outer step counter (LR schedule, logging) and the
2567    /// GPU-side AdamW step counter (bias correction). Must be called before
2568    /// any `train_batch()` calls.
2569    pub fn set_initial_step(&mut self, step: usize) {
2570        self.step = step;
2571        self.gpu_training.step = step as u32;
2572    }
2573
2574    /// Set max_steps for cosine LR scheduler (ENT-275).
2575    ///
2576    /// Called by `train_loop_cuda` when `max_steps` is not explicitly set in
2577    /// the YAML config — auto-computes `epochs × batches_per_epoch` so cosine
2578    /// decay activates instead of falling back to constant lr.
2579    pub fn set_max_steps(&mut self, max_steps: usize) {
2580        self.config.max_steps = Some(max_steps);
2581    }
2582
2583    /// Get current learning rate (warmup + cosine decay).
2584    ///
2585    /// ALB-079: Phase 1 = linear warmup (0 → lr_max), Phase 2 = cosine decay
2586    /// (lr_max → 0) over remaining steps. Requires `max_steps` for decay;
2587    /// without it, falls back to constant lr after warmup.
2588    pub fn current_lr(&self) -> f32 {
2589        let base_lr = self.config.lr;
2590        if self.step < self.config.warmup_steps {
2591            // Phase 1: Linear warmup
2592            base_lr * (self.step as f32 / self.config.warmup_steps.max(1) as f32)
2593        } else if let Some(max_steps) = self.config.max_steps {
2594            // Phase 2: Cosine decay from lr_max to 0
2595            let decay_steps = max_steps.saturating_sub(self.config.warmup_steps);
2596            if decay_steps == 0 {
2597                return base_lr;
2598            }
2599            let decay_step = self.step - self.config.warmup_steps;
2600            let progress = (decay_step as f32 / decay_steps as f32).min(1.0);
2601            0.5 * base_lr * (1.0 + (std::f32::consts::PI * progress).cos())
2602        } else {
2603            // No max_steps: constant lr (legacy behavior)
2604            base_lr
2605        }
2606    }
2607
2608    /// KAIZEN-047: Enable step profiling with a report every `interval` steps.
2609    ///
2610    /// When enabled, prints a table of wall-clock timings per training phase
2611    /// every `interval` training steps. Use interval=0 for manual-only reporting.
2612    ///
2613    /// # Contract (C-STEPPROF-001)
2614    ///
2615    /// - No additional GPU synchronization points (relies on existing syncs)
2616    /// - Overhead: ~11 `Instant::now()` calls per step (~1µs total on Linux)
2617    /// - Timings include async dispatch overhead (not pure kernel time)
2618    pub fn enable_profiler(&mut self, interval: usize) {
2619        self.profiler = StepProfiler::new(true, interval);
2620    }
2621
2622    /// Print the profiler report (if profiling is enabled).
2623    pub fn print_profiler_report(&self) {
2624        self.profiler.print_report();
2625    }
2626
2627    /// R-004: Get last observed gradient L2 norm (LM head proxy).
2628    pub fn last_grad_norm(&self) -> f32 {
2629        self.last_grad_norm
2630    }
2631
2632    /// R-040: Get per-parameter-group gradient norms.
2633    /// Returns (lm_head_grad_norm, embed_grad_norm).
2634    pub fn param_grad_norms(&self) -> (f32, f32) {
2635        (self.last_grad_norm, self.last_embed_grad_norm)
2636    }
2637
2638    /// R-012: Get total trainable parameter count for MFU calculation.
2639    pub fn num_params(&self) -> usize {
2640        self.model.parameters().iter().map(|t| t.len()).sum()
2641    }
2642
2643    /// R-013: Query GPU memory usage (used_mb, total_mb).
2644    pub fn gpu_memory_mb(&self) -> (u64, u64) {
2645        match self.cuda_trainer.context().memory_info() {
2646            Ok((free, total)) => {
2647                let total_mb = (total / (1024 * 1024)) as u64;
2648                let used_mb = ((total - free) / (1024 * 1024)) as u64;
2649                (used_mb, total_mb)
2650            }
2651            Err(_) => (0, 0),
2652        }
2653    }
2654
2655    /// Sync all GPU weights back to CPU model.
2656    ///
2657    /// # Contract (C-SYNCWT-001)
2658    ///
2659    /// Must be called before save or any CPU model access after training.
2660    pub fn sync_weights_to_cpu(&mut self) {
2661        let use_nf4 = self.config.quantize_nf4 && self.config.is_lora();
2662
2663        if use_nf4 {
2664            // ENT-263: NF4 blocks are frozen — base weights don't change.
2665            // Only download LoRA adapter weights for checkpoint saving.
2666            // The base model on CPU stays as-is (original pretrained weights).
2667            // LoRA weights are saved separately (adapter_config.json + adapter.safetensors).
2668            // For now, skip per-layer sync — base weights are unchanged.
2669        } else {
2670            for (layer_idx, block) in self.cuda_blocks.iter().enumerate() {
2671                if let Ok(weights) = block.download_weights() {
2672                    let layer = &mut self.model.layers[layer_idx];
2673
2674                    layer.self_attn.w_q = Tensor::from_vec(weights.w_q, false);
2675                    layer.self_attn.w_k = Tensor::from_vec(weights.w_k, false);
2676                    layer.self_attn.w_v = Tensor::from_vec(weights.w_v, false);
2677                    layer.self_attn.w_o = Tensor::from_vec(weights.w_o, false);
2678
2679                    layer.ffn.w_gate = Tensor::from_vec(weights.w_gate, false);
2680                    layer.ffn.w_up = Tensor::from_vec(weights.w_up, false);
2681                    layer.ffn.w_down = Tensor::from_vec(weights.w_down, false);
2682
2683                    layer.input_norm.weight = Tensor::from_vec(weights.input_norm_weight, false);
2684                    layer.post_attn_norm.weight =
2685                        Tensor::from_vec(weights.post_attn_norm_weight, false);
2686                }
2687            }
2688        }
2689
2690        // Sync final norm weight
2691        if let Ok(norm_data) = self.cuda_trainer.download(&self.gpu_training.final_norm_weight) {
2692            self.model.norm.weight = Tensor::from_vec(norm_data, false);
2693        }
2694
2695        // Sync LM head weight
2696        // ALB-097: ALWAYS save GPU-trained LM head, even for tied-weight models.
2697        // During GPU training, lm_head diverges from embed_tokens because they have
2698        // separate optimizers (GPU AdamW vs CPU AdamW). If we skip the sync for tied
2699        // weights, the checkpoint loses 500+ steps of GPU LM head training → random-init
2700        // loss on resume (Five Whys root cause of ALB-097).
2701        if let Ok(lm_data) = self.cuda_trainer.download(&self.lm_head_weight_gpu) {
2702            self.model.lm_head = Some(Tensor::from_vec(lm_data, false));
2703        }
2704    }
2705
2706    /// Get reference to model (syncs weights first).
2707    pub fn model(&self) -> &Transformer {
2708        &self.model
2709    }
2710
2711    /// Get mutable reference to model.
2712    pub fn model_mut(&mut self) -> &mut Transformer {
2713        &mut self.model
2714    }
2715
2716    /// Check if using mixed precision.
2717    pub fn is_mixed_precision(&self) -> bool {
2718        self.config.precision_config.is_mixed()
2719    }
2720
2721    /// Get the gradient scaler (R-002: loss scaling for mixed precision).
2722    pub fn grad_scaler(&self) -> &GradScaler {
2723        &self.grad_scaler
2724    }
2725
2726    /// Check if using gradient checkpointing.
2727    pub fn is_checkpointing(&self) -> bool {
2728        self.config.checkpoint_config.enabled
2729    }
2730
2731    /// Save model weights (syncs GPU→CPU first).
2732    pub fn save(
2733        &mut self,
2734        path: impl AsRef<std::path::Path>,
2735        name: &str,
2736        architecture: &str,
2737    ) -> crate::Result<()> {
2738        self.sync_weights_to_cpu();
2739
2740        // Use named_parameters() for correct name mapping (handles attention biases etc.)
2741        let params: Vec<(String, Tensor)> = self
2742            .model
2743            .named_parameters()
2744            .into_iter()
2745            .map(|(name, tensor)| (name, tensor.clone()))
2746            .collect();
2747
2748        let metadata = ModelMetadata::new(name, architecture);
2749        let model = Model::new(metadata, params);
2750        let config = SaveConfig::new(ModelFormat::SafeTensors);
2751
2752        save_model(&model, path, &config)
2753    }
2754
2755    /// R-011: Prepare checkpoint data for async save.
2756    /// Syncs GPU weights to CPU and snapshots tensor data as Send-able Vec<f32>.
2757    /// Returns a closure that writes the checkpoint file from another thread.
2758    pub fn prepare_async_save(
2759        &mut self,
2760        name: &str,
2761        architecture: &str,
2762    ) -> Box<dyn FnOnce(&std::path::Path) -> crate::Result<()> + Send> {
2763        self.sync_weights_to_cpu();
2764
2765        // Use named_parameters() for correct name mapping (handles attention biases etc.)
2766        let param_data: Vec<(String, Vec<f32>)> = self
2767            .model
2768            .named_parameters()
2769            .into_iter()
2770            .map(|(n, t)| (n, t.data().to_vec()))
2771            .collect();
2772
2773        let name = name.to_string();
2774        let architecture = architecture.to_string();
2775
2776        Box::new(move |path: &std::path::Path| {
2777            let params: Vec<(String, Tensor)> =
2778                param_data.into_iter().map(|(n, d)| (n, Tensor::from_vec(d, false))).collect();
2779            let metadata = ModelMetadata::new(&name, &architecture);
2780            let model = Model::new(metadata, params);
2781            let config = SaveConfig::new(ModelFormat::SafeTensors);
2782            save_model(&model, path, &config)
2783        })
2784    }
2785
2786    /// ALB-096: Save model weights as APR checkpoint (syncs GPU→CPU first).
2787    ///
2788    /// Single atomic file containing all model weights. Use `save_apr_checkpoint()`
2789    /// to include optimizer state and training metadata in the same file.
2790    pub fn save_apr(
2791        &mut self,
2792        path: impl AsRef<std::path::Path>,
2793        name: &str,
2794        architecture: &str,
2795    ) -> crate::Result<()> {
2796        self.save_apr_with_tokenizer(path, name, architecture, None)
2797    }
2798
2799    /// SPEC-SHIP-TWO-001 §81 P0-D + P0-E: save APR checkpoint with arch
2800    /// metadata keys AND optionally embed the source tokenizer.json.
2801    ///
2802    /// When `tokenizer_dir` is `Some`, reads `<dir>/tokenizer.json` and
2803    /// embeds the vocabulary + merges + BOS/EOS IDs as well-known
2804    /// metadata keys. This makes the resulting .apr file standalone for
2805    /// `apr qa`, `apr run`, etc. — no `--tokenizer` flag required at
2806    /// downstream tool dispatch.
2807    pub fn save_apr_with_tokenizer(
2808        &mut self,
2809        path: impl AsRef<std::path::Path>,
2810        name: &str,
2811        architecture: &str,
2812        tokenizer_dir: Option<&std::path::Path>,
2813    ) -> crate::Result<()> {
2814        self.sync_weights_to_cpu();
2815
2816        let params: Vec<(String, Tensor)> = self
2817            .model
2818            .named_parameters()
2819            .into_iter()
2820            .map(|(name, tensor)| (name, tensor.clone()))
2821            .collect();
2822
2823        // SPEC-SHIP-TWO-001 §81 P0-E: write individual arch metadata keys
2824        // so downstream tools (apr qa C-03, apr bench, realizar) can read them
2825        // via AprV2Metadata's typed fields. The legacy save_model() path only
2826        // carries `name + architecture + format + version` which fails C-03.
2827        use crate::io::save::infer_all_tensor_shapes;
2828        use aprender::serialization::apr::AprWriter;
2829        use serde_json::Value as Jv;
2830
2831        let mc = &self.config.model_config;
2832        let mut writer = AprWriter::new();
2833
2834        // Identity / version metadata (preserves save_model behavior)
2835        writer.set_metadata("model_name", Jv::String(name.to_string()));
2836        writer.set_metadata("architecture", Jv::String(architecture.to_string()));
2837        writer.set_metadata("version", Jv::String("0.1.0".into()));
2838        writer.set_metadata("format", Jv::String("entrenar-checkpoint".into()));
2839
2840        // Arch dim keys (well-known to AprWriter::build_v2_metadata,
2841        // map to AprV2Metadata typed fields).
2842        writer.set_metadata(
2843            "hidden_size",
2844            Jv::Number(serde_json::Number::from(mc.hidden_size as u64)),
2845        );
2846        writer.set_metadata(
2847            "num_hidden_layers",
2848            Jv::Number(serde_json::Number::from(mc.num_hidden_layers as u64)),
2849        );
2850        writer.set_metadata(
2851            "num_attention_heads",
2852            Jv::Number(serde_json::Number::from(mc.num_attention_heads as u64)),
2853        );
2854        writer.set_metadata(
2855            "num_kv_heads",
2856            Jv::Number(serde_json::Number::from(mc.num_kv_heads as u64)),
2857        );
2858        writer.set_metadata(
2859            "intermediate_size",
2860            Jv::Number(serde_json::Number::from(mc.intermediate_size as u64)),
2861        );
2862        writer
2863            .set_metadata("vocab_size", Jv::Number(serde_json::Number::from(mc.vocab_size as u64)));
2864        writer.set_metadata(
2865            "max_position_embeddings",
2866            Jv::Number(serde_json::Number::from(mc.max_position_embeddings as u64)),
2867        );
2868        if let Some(rope) = serde_json::Number::from_f64(mc.rope_theta as f64) {
2869            writer.set_metadata("rope_theta", Jv::Number(rope));
2870        }
2871        if let Some(eps) = serde_json::Number::from_f64(mc.rms_norm_eps as f64) {
2872            writer.set_metadata("rms_norm_eps", Jv::Number(eps));
2873        }
2874
2875        // SPEC-SHIP-TWO-001 §81 P0-D: embed tokenizer.json from
2876        // `tokenizer_dir/tokenizer.json` so `apr qa` (which requires
2877        // an embedded tokenizer) accepts the resulting .apr file.
2878        // ALB-130 style: parse vocab + merges + special token IDs and
2879        // set as well-known metadata keys.
2880        if let Some(dir) = tokenizer_dir {
2881            let tok_path = dir.join("tokenizer.json");
2882            if let Ok(json_bytes) = std::fs::read(&tok_path) {
2883                if let Ok(tok) = serde_json::from_slice::<Jv>(&json_bytes) {
2884                    if let Some(model) = tok.get("model") {
2885                        if let Some(vocab_obj) = model.get("vocab").and_then(|v| v.as_object()) {
2886                            let mut vocab_pairs: Vec<(String, u64)> = vocab_obj
2887                                .iter()
2888                                .filter_map(|(k, v)| Some((k.clone(), v.as_u64()?)))
2889                                .collect();
2890                            vocab_pairs.sort_by_key(|(_, id)| *id);
2891                            let vocab: Vec<Jv> =
2892                                vocab_pairs.into_iter().map(|(k, _)| Jv::String(k)).collect();
2893                            writer.set_metadata("tokenizer.vocabulary", Jv::Array(vocab));
2894                        }
2895                        if let Some(merges_arr) = model.get("merges").and_then(|m| m.as_array()) {
2896                            let merges: Vec<Jv> = merges_arr
2897                                .iter()
2898                                .filter_map(|v| v.as_str().map(|s| Jv::String(s.to_string())))
2899                                .collect();
2900                            writer.set_metadata("tokenizer.merges", Jv::Array(merges));
2901                        }
2902                    }
2903                    // BOS / EOS from added_tokens (HF format).
2904                    if let Some(added) = tok.get("added_tokens").and_then(|a| a.as_array()) {
2905                        for entry in added {
2906                            let content =
2907                                entry.get("content").and_then(|c| c.as_str()).unwrap_or("");
2908                            let id = entry.get("id").and_then(|i| i.as_u64());
2909                            if let Some(id) = id {
2910                                match content {
2911                                    "<s>" | "<|im_start|>" | "<|begin_of_text|>" => {
2912                                        writer.set_metadata(
2913                                            "tokenizer.bos_token_id",
2914                                            Jv::Number(serde_json::Number::from(id)),
2915                                        );
2916                                    }
2917                                    "</s>" | "<|im_end|>" | "<|end_of_text|>" | "<|endoftext|>" => {
2918                                        writer.set_metadata(
2919                                            "tokenizer.eos_token_id",
2920                                            Jv::Number(serde_json::Number::from(id)),
2921                                        );
2922                                    }
2923                                    _ => {}
2924                                }
2925                            }
2926                        }
2927                    }
2928                }
2929            }
2930        }
2931
2932        // Tensors — reuse io::save's shape inference for 2D weight handling.
2933        let shapes = infer_all_tensor_shapes(&params);
2934        for (tname, tensor) in &params {
2935            let data = tensor.data();
2936            let slice = data.as_slice().expect("tensor data must be contiguous");
2937            let shape = shapes.get(tname).cloned().unwrap_or_else(|| vec![tensor.len()]);
2938            writer.add_tensor_f32(tname, shape, slice);
2939        }
2940
2941        writer
2942            .write(path)
2943            .map_err(|e| crate::error::Error::Serialization(format!("APR write failed: {e}")))
2944    }
2945
2946    /// ALB-096: Prepare APR checkpoint data for async save.
2947    ///
2948    /// Syncs GPU weights to CPU and snapshots tensor data + optimizer state as
2949    /// Send-able `Vec<f32>`. Returns a closure that writes a single atomic APR
2950    /// file from another thread. Includes model weights + CPU embedding optimizer
2951    /// state + training metadata — all in one file.
2952    fn snapshot_param_data(&self) -> Vec<(String, Vec<f32>)> {
2953        let use_nf4 = self.config.quantize_nf4 && self.config.is_lora();
2954        if use_nf4 {
2955            let frozen_suffixes = [
2956                "q_proj.weight",
2957                "k_proj.weight",
2958                "v_proj.weight",
2959                "o_proj.weight",
2960                "gate_proj.weight",
2961                "up_proj.weight",
2962                "down_proj.weight",
2963            ];
2964            self.model
2965                .named_parameters()
2966                .into_iter()
2967                .filter(|(n, _)| !frozen_suffixes.iter().any(|s| n.ends_with(s)))
2968                .map(|(n, t)| (n, t.data().to_vec()))
2969                .collect()
2970        } else {
2971            self.model.named_parameters().into_iter().map(|(n, t)| (n, t.data().to_vec())).collect()
2972        }
2973    }
2974
2975    fn snapshot_lora_data(&self) -> Vec<(usize, Vec<f32>, Vec<f32>, Vec<f32>, Vec<f32>)> {
2976        if self.config.quantize_nf4 && self.config.is_lora() {
2977            self.cuda_blocks
2978                .iter()
2979                .enumerate()
2980                .filter_map(|(i, block)| {
2981                    block
2982                        .download_lora_weights()
2983                        .ok()
2984                        .map(|(a_q, b_q, a_v, b_v)| (i, a_q, b_q, a_v, b_v))
2985                })
2986                .collect()
2987        } else {
2988            Vec::new()
2989        }
2990    }
2991
2992    pub fn prepare_async_apr_save(
2993        &mut self,
2994        name: &str,
2995        architecture: &str,
2996        step: usize,
2997        loss: f64,
2998        lr: f64,
2999    ) -> Box<dyn FnOnce(&std::path::Path) -> crate::Result<()> + Send> {
3000        self.prepare_async_apr_save_with_tokenizer(name, architecture, step, loss, lr, None)
3001    }
3002
3003    /// ALB-130: Prepare APR checkpoint with embedded tokenizer for inference.
3004    ///
3005    /// Training checkpoints must be self-contained for eval (`apr eval --task humaneval`).
3006    /// Without embedded tokenizer, inference falls back to structural validation (fake 100%).
3007    /// The tokenizer path comes from `spec.data.tokenizer` in the training YAML.
3008    pub fn prepare_async_apr_save_with_tokenizer(
3009        &mut self,
3010        name: &str,
3011        architecture: &str,
3012        step: usize,
3013        loss: f64,
3014        lr: f64,
3015        tokenizer_path: Option<&std::path::Path>,
3016    ) -> Box<dyn FnOnce(&std::path::Path) -> crate::Result<()> + Send> {
3017        self.sync_weights_to_cpu();
3018
3019        let param_data = self.snapshot_param_data();
3020        let lora_data = self.snapshot_lora_data();
3021
3022        // Snapshot CPU embedding optimizer state
3023        let embed_m: Vec<Vec<f32>> = self
3024            .embed_optimizer
3025            .first_moments()
3026            .iter()
3027            .filter_map(|opt| opt.as_ref().map(ndarray::ArrayBase::to_vec))
3028            .collect();
3029        let embed_v: Vec<Vec<f32>> = self
3030            .embed_optimizer
3031            .second_moments()
3032            .iter()
3033            .filter_map(|opt| opt.as_ref().map(ndarray::ArrayBase::to_vec))
3034            .collect();
3035        let embed_step = self.embed_optimizer.step_count();
3036
3037        // ALB-118: Download GPU block optimizer states (m/v moments) for checkpointing.
3038        // Without this, resume re-initializes all 24 blocks' AdamW state to zero,
3039        // causing loss spikes and convergence failure (v10/v11/v12 post-mortems).
3040        // Transfer cost: ~2.3 GB D2H, <6ms on PCIe4/5.
3041        let block_optim_data: Vec<Vec<(String, Vec<f32>)>> = self
3042            .gpu_training
3043            .optimizer_states
3044            .iter()
3045            .map(|state| state.download_to_host().unwrap_or_default())
3046            .collect();
3047
3048        // ALB-118: Download LM head and final norm optimizer states
3049        let lm_head_m_host = {
3050            let mut buf = vec![0.0f32; self.lm_head_m.len()];
3051            let _ = self.lm_head_m.copy_to_host(&mut buf);
3052            buf
3053        };
3054        let lm_head_v_host = {
3055            let mut buf = vec![0.0f32; self.lm_head_v.len()];
3056            let _ = self.lm_head_v.copy_to_host(&mut buf);
3057            buf
3058        };
3059        let final_norm_m_host = {
3060            let mut buf = vec![0.0f32; self.final_norm_m.len()];
3061            let _ = self.final_norm_m.copy_to_host(&mut buf);
3062            buf
3063        };
3064        let final_norm_v_host = {
3065            let mut buf = vec![0.0f32; self.final_norm_v.len()];
3066            let _ = self.final_norm_v.copy_to_host(&mut buf);
3067            buf
3068        };
3069
3070        let name = name.to_string();
3071        let architecture = architecture.to_string();
3072        let model_config_json = serde_json::to_string(&self.config.model_config).ok();
3073        let is_delta_checkpoint = self.config.quantize_nf4 && self.config.is_lora();
3074
3075        // SPEC-SHIP-TWO-001 §81 P0-E: extract individual arch metadata keys
3076        // so downstream tools (apr qa, apr bench, apr export) can read them
3077        // via AprV2Metadata's typed fields. The `model_config` JSON blob is
3078        // unrecognized by AprWriter::build_v2_metadata and goes into the
3079        // `custom` map — which `realizar::gguf::config::from_apr` does NOT
3080        // read (it requires `apr.metadata.hidden_size` etc. to be Some).
3081        let arch_hidden_size = self.config.model_config.hidden_size;
3082        let arch_num_layers = self.config.model_config.num_hidden_layers;
3083        let arch_num_heads = self.config.model_config.num_attention_heads;
3084        let arch_num_kv_heads = self.config.model_config.num_kv_heads;
3085        let arch_intermediate_size = self.config.model_config.intermediate_size;
3086        let arch_vocab_size = self.config.model_config.vocab_size;
3087        let arch_max_position_embeddings = self.config.model_config.max_position_embeddings;
3088        let arch_rope_theta = self.config.model_config.rope_theta;
3089        let arch_rms_norm_eps = self.config.model_config.rms_norm_eps;
3090
3091        // ALB-130: Pre-read tokenizer.json for embedding in checkpoint.
3092        // Parse HuggingFace tokenizer format → extract vocab + merges + special token IDs.
3093        let tokenizer_data: Option<(Vec<String>, Vec<String>, Option<u64>, Option<u64>)> =
3094            tokenizer_path.and_then(|p| {
3095                let json_bytes = std::fs::read(p).ok()?;
3096                let tok: serde_json::Value = serde_json::from_slice(&json_bytes).ok()?;
3097                let model = tok.get("model")?;
3098                let vocab_obj = model.get("vocab")?.as_object()?;
3099                // Build sorted-by-id vocab list
3100                let mut vocab_pairs: Vec<(String, u64)> =
3101                    vocab_obj.iter().filter_map(|(k, v)| Some((k.clone(), v.as_u64()?))).collect();
3102                vocab_pairs.sort_by_key(|(_, id)| *id);
3103                let vocab: Vec<String> = vocab_pairs.into_iter().map(|(k, _)| k).collect();
3104                // Merges as "token1 token2" strings
3105                let merges: Vec<String> = model
3106                    .get("merges")?
3107                    .as_array()?
3108                    .iter()
3109                    .filter_map(|v| v.as_str().map(String::from))
3110                    .collect();
3111                // Special tokens: BOS=<s>=1, EOS=</s>=2 (from added_tokens)
3112                let added = tok.get("added_tokens").and_then(|a| a.as_array());
3113                let bos_id = added.and_then(|arr| {
3114                    arr.iter()
3115                        .find(|t| t.get("content").and_then(|c| c.as_str()) == Some("<s>"))
3116                        .and_then(|t| t.get("id")?.as_u64())
3117                });
3118                let eos_id = added.and_then(|arr| {
3119                    arr.iter()
3120                        .find(|t| t.get("content").and_then(|c| c.as_str()) == Some("</s>"))
3121                        .and_then(|t| t.get("id")?.as_u64())
3122                });
3123                if vocab.is_empty() {
3124                    return None;
3125                }
3126                println!(
3127                    "  [ALB-130] Embedding tokenizer: {} vocab, {} merges",
3128                    vocab.len(),
3129                    merges.len()
3130                );
3131                Some((vocab, merges, bos_id, eos_id))
3132            });
3133
3134        Box::new(move |path: &std::path::Path| {
3135            use aprender::serialization::apr::AprWriter;
3136            use serde_json::Value as Jv;
3137
3138            let mut writer = AprWriter::new();
3139
3140            // Metadata
3141            writer.set_metadata("model_name", Jv::String(name));
3142            writer.set_metadata("architecture", Jv::String(architecture));
3143            writer.set_metadata(
3144                "format",
3145                Jv::String(if is_delta_checkpoint {
3146                    "entrenar-delta-checkpoint".into()
3147                } else {
3148                    "entrenar-checkpoint".into()
3149                }),
3150            );
3151            writer.set_metadata("checkpoint_step", Jv::String(step.to_string()));
3152            writer.set_metadata("loss", Jv::String(format!("{loss:.6}")));
3153            writer.set_metadata("learning_rate", Jv::String(format!("{lr:.6e}")));
3154            writer.set_metadata("optimizer_step", Jv::String(embed_step.to_string()));
3155            if let Some(cfg) = model_config_json {
3156                writer.set_metadata("model_config", Jv::String(cfg));
3157            }
3158
3159            // SPEC-SHIP-TWO-001 §81 P0-E: write individual arch metadata keys
3160            // so realizar's `from_apr` (C-03 gate) accepts the checkpoint.
3161            // `serde_json::Number::from(u as u64)` converts usize losslessly.
3162            writer.set_metadata(
3163                "hidden_size",
3164                Jv::Number(serde_json::Number::from(arch_hidden_size as u64)),
3165            );
3166            writer.set_metadata(
3167                "num_hidden_layers",
3168                Jv::Number(serde_json::Number::from(arch_num_layers as u64)),
3169            );
3170            writer.set_metadata(
3171                "num_attention_heads",
3172                Jv::Number(serde_json::Number::from(arch_num_heads as u64)),
3173            );
3174            writer.set_metadata(
3175                "num_kv_heads",
3176                Jv::Number(serde_json::Number::from(arch_num_kv_heads as u64)),
3177            );
3178            writer.set_metadata(
3179                "intermediate_size",
3180                Jv::Number(serde_json::Number::from(arch_intermediate_size as u64)),
3181            );
3182            writer.set_metadata(
3183                "vocab_size",
3184                Jv::Number(serde_json::Number::from(arch_vocab_size as u64)),
3185            );
3186            writer.set_metadata(
3187                "max_position_embeddings",
3188                Jv::Number(serde_json::Number::from(arch_max_position_embeddings as u64)),
3189            );
3190            if let Some(rope) = serde_json::Number::from_f64(arch_rope_theta as f64) {
3191                writer.set_metadata("rope_theta", Jv::Number(rope));
3192            }
3193            if let Some(eps) = serde_json::Number::from_f64(arch_rms_norm_eps as f64) {
3194                writer.set_metadata("rms_norm_eps", Jv::Number(eps));
3195            }
3196
3197            // ALB-130: Embed tokenizer vocab + merges for standalone inference
3198            if let Some((vocab, merges, bos_id, eos_id)) = tokenizer_data {
3199                writer.set_metadata(
3200                    "tokenizer.vocabulary",
3201                    Jv::Array(vocab.into_iter().map(Jv::String).collect()),
3202                );
3203                writer.set_metadata(
3204                    "tokenizer.merges",
3205                    Jv::Array(merges.into_iter().map(Jv::String).collect()),
3206                );
3207                if let Some(bos) = bos_id {
3208                    writer.set_metadata("tokenizer.bos_token_id", Jv::Number(bos.into()));
3209                }
3210                if let Some(eos) = eos_id {
3211                    writer.set_metadata("tokenizer.eos_token_id", Jv::Number(eos.into()));
3212                }
3213            }
3214
3215            // Find hidden_size from norm weights for shape inference
3216            let hidden_size = param_data
3217                .iter()
3218                .find(|(n, _)| n.ends_with("layernorm.weight") || n == "model.norm.weight")
3219                .map_or(0, |(_, d)| d.len());
3220
3221            // Model weight tensors
3222            for (tensor_name, data) in &param_data {
3223                let shape = infer_tensor_shape(tensor_name, data.len(), hidden_size);
3224                writer.add_tensor_f32(tensor_name.clone(), shape, data);
3225            }
3226
3227            // Optimizer state tensors
3228            for (i, m_data) in embed_m.iter().enumerate() {
3229                let len = m_data.len();
3230                writer.add_tensor_f32(
3231                    format!("__training__.embed_optimizer.m.{i}"),
3232                    vec![len],
3233                    m_data,
3234                );
3235            }
3236            for (i, v_data) in embed_v.iter().enumerate() {
3237                let len = v_data.len();
3238                writer.add_tensor_f32(
3239                    format!("__training__.embed_optimizer.v.{i}"),
3240                    vec![len],
3241                    v_data,
3242                );
3243            }
3244
3245            // ALB-118: Save GPU block optimizer states (m/v moments for all 24 blocks)
3246            for (layer_idx, buffers) in block_optim_data.iter().enumerate() {
3247                for (suffix, data) in buffers {
3248                    let len = data.len();
3249                    writer.add_tensor_f32(
3250                        format!("__training__.block_optimizer.{layer_idx}.{suffix}"),
3251                        vec![len],
3252                        data,
3253                    );
3254                }
3255            }
3256
3257            // ALB-118: Save LM head and final norm optimizer states
3258            if !lm_head_m_host.is_empty() {
3259                let len = lm_head_m_host.len();
3260                writer.add_tensor_f32(
3261                    "__training__.lm_head_optimizer.m".to_string(),
3262                    vec![len],
3263                    &lm_head_m_host,
3264                );
3265                let len = lm_head_v_host.len();
3266                writer.add_tensor_f32(
3267                    "__training__.lm_head_optimizer.v".to_string(),
3268                    vec![len],
3269                    &lm_head_v_host,
3270                );
3271            }
3272            if !final_norm_m_host.is_empty() {
3273                let len = final_norm_m_host.len();
3274                writer.add_tensor_f32(
3275                    "__training__.final_norm_optimizer.m".to_string(),
3276                    vec![len],
3277                    &final_norm_m_host,
3278                );
3279                let len = final_norm_v_host.len();
3280                writer.add_tensor_f32(
3281                    "__training__.final_norm_optimizer.v".to_string(),
3282                    vec![len],
3283                    &final_norm_v_host,
3284                );
3285            }
3286
3287            // ENT-276: Save LoRA adapter weights (QLoRA checkpoint resume)
3288            for (layer_idx, a_q, b_q, a_v, b_v) in &lora_data {
3289                if !a_q.is_empty() {
3290                    writer.add_tensor_f32(
3291                        format!("lora.{layer_idx}.q_proj.lora_a"),
3292                        vec![a_q.len()],
3293                        a_q,
3294                    );
3295                    writer.add_tensor_f32(
3296                        format!("lora.{layer_idx}.q_proj.lora_b"),
3297                        vec![b_q.len()],
3298                        b_q,
3299                    );
3300                }
3301                if !a_v.is_empty() {
3302                    writer.add_tensor_f32(
3303                        format!("lora.{layer_idx}.v_proj.lora_a"),
3304                        vec![a_v.len()],
3305                        a_v,
3306                    );
3307                    writer.add_tensor_f32(
3308                        format!("lora.{layer_idx}.v_proj.lora_b"),
3309                        vec![b_v.len()],
3310                        b_v,
3311                    );
3312                }
3313            }
3314
3315            // Write APR checkpoint to file
3316            writer
3317                .write(path)
3318                .map_err(|e| crate::error::Error::Serialization(format!("APR save failed: {e}")))?;
3319
3320            Ok(())
3321        })
3322    }
3323
3324    /// GPU device name.
3325    pub fn gpu_name(&self) -> String {
3326        self.cuda_trainer.device_name()
3327    }
3328
3329    /// ENT-269: Save LoRA adapter weights as PEFT-compatible files.
3330    ///
3331    /// Downloads LoRA A/B matrices from GPU, un-scales B (divide by lora_scale),
3332    /// transposes to PEFT convention (A=[rank, d_in], B=[d_out, rank]),
3333    /// and writes `adapter_model.safetensors` + `adapter_config.json`.
3334    ///
3335    /// # Contract: C-QLORA-SAVE-001
3336    ///
3337    /// NF4 QLoRA training MUST produce `adapter_model.safetensors` in output_dir.
3338    pub fn save_cuda_lora_adapter(
3339        &self,
3340        output_dir: &std::path::Path,
3341        base_model_name: Option<&str>,
3342    ) -> crate::Result<()> {
3343        if !self.config.quantize_nf4 || !self.config.is_lora() {
3344            return Ok(()); // Not a QLoRA run, nothing to save
3345        }
3346
3347        let lora_rank = self.config.lora_rank.unwrap_or(16);
3348        let lora_alpha = self.config.lora_alpha.unwrap_or(2.0 * lora_rank as f32);
3349        let lora_scale = lora_alpha / lora_rank as f32;
3350        let hidden_size = self.config.model_config.hidden_size;
3351        let head_dim = self.config.model_config.head_dim();
3352        let q_dim = self.config.model_config.num_attention_heads * head_dim;
3353        let kv_hidden = self.config.model_config.num_kv_heads * head_dim;
3354
3355        let lora_config =
3356            crate::lora::LoRAConfig::new(lora_rank, lora_alpha).target_qv_projections();
3357
3358        let mut adapters: Vec<(String, crate::lora::LoRALayer)> = Vec::new();
3359
3360        for (i, block) in self.cuda_blocks.iter().enumerate() {
3361            let (a_q, b_q_scaled, a_v, b_v_scaled) = match block.download_lora_weights() {
3362                Ok(weights) => weights,
3363                Err(_) => continue, // Skip non-NF4 blocks
3364            };
3365
3366            if a_q.is_empty() && a_v.is_empty() {
3367                continue;
3368            }
3369
3370            // Q projection LoRA
3371            if !a_q.is_empty() {
3372                // GPU stores A_q as [hidden, rank] row-major, PEFT expects [rank, hidden]
3373                let mut a_transposed = vec![0.0f32; lora_rank * hidden_size];
3374                for r in 0..hidden_size {
3375                    for c in 0..lora_rank {
3376                        a_transposed[c * hidden_size + r] = a_q[r * lora_rank + c];
3377                    }
3378                }
3379
3380                // GPU stores B_q as [rank, q_dim] pre-scaled by lora_scale
3381                // PEFT expects [q_dim, rank] un-scaled
3382                let inv_scale = if lora_scale.abs() > 1e-10 { 1.0 / lora_scale } else { 1.0 };
3383                let mut b_transposed = vec![0.0f32; q_dim * lora_rank];
3384                for r in 0..lora_rank {
3385                    for c in 0..q_dim {
3386                        b_transposed[c * lora_rank + r] = b_q_scaled[r * q_dim + c] * inv_scale;
3387                    }
3388                }
3389
3390                let base_weight = crate::autograd::Tensor::zeros(q_dim * hidden_size, false);
3391                let mut layer = crate::lora::LoRALayer::new(
3392                    base_weight,
3393                    q_dim,
3394                    hidden_size,
3395                    lora_rank,
3396                    lora_alpha,
3397                );
3398                // Overwrite the A and B data with trained weights
3399                layer.lora_a_mut().data_mut().assign(&ndarray::Array1::from(a_transposed));
3400                layer.lora_b_mut().data_mut().assign(&ndarray::Array1::from(b_transposed));
3401
3402                adapters.push((format!("model.layers.{i}.self_attn.q_proj"), layer));
3403            }
3404
3405            // V projection LoRA
3406            if !a_v.is_empty() {
3407                let mut a_transposed = vec![0.0f32; lora_rank * hidden_size];
3408                for r in 0..hidden_size {
3409                    for c in 0..lora_rank {
3410                        a_transposed[c * hidden_size + r] = a_v[r * lora_rank + c];
3411                    }
3412                }
3413
3414                let inv_scale = if lora_scale.abs() > 1e-10 { 1.0 / lora_scale } else { 1.0 };
3415                let mut b_transposed = vec![0.0f32; kv_hidden * lora_rank];
3416                for r in 0..lora_rank {
3417                    for c in 0..kv_hidden {
3418                        b_transposed[c * lora_rank + r] = b_v_scaled[r * kv_hidden + c] * inv_scale;
3419                    }
3420                }
3421
3422                let base_weight = crate::autograd::Tensor::zeros(kv_hidden * hidden_size, false);
3423                let mut layer = crate::lora::LoRALayer::new(
3424                    base_weight,
3425                    kv_hidden,
3426                    hidden_size,
3427                    lora_rank,
3428                    lora_alpha,
3429                );
3430                layer.lora_a_mut().data_mut().assign(&ndarray::Array1::from(a_transposed));
3431                layer.lora_b_mut().data_mut().assign(&ndarray::Array1::from(b_transposed));
3432
3433                adapters.push((format!("model.layers.{i}.self_attn.v_proj"), layer));
3434            }
3435        }
3436
3437        if adapters.is_empty() {
3438            println!("  [WARN] No LoRA adapters found to save");
3439            return Ok(());
3440        }
3441
3442        let adapter_refs: Vec<(&str, &crate::lora::LoRALayer)> =
3443            adapters.iter().map(|(name, layer)| (name.as_str(), layer)).collect();
3444
3445        std::fs::create_dir_all(output_dir).ok();
3446        crate::lora::save_adapter_peft(&adapter_refs, &lora_config, base_model_name, output_dir)
3447            .map_err(|e| crate::error::Error::Io(format!("Failed to save PEFT adapter: {e}")))?;
3448
3449        let adapter_path = output_dir.join("adapter_model.safetensors");
3450        let size_mb =
3451            std::fs::metadata(&adapter_path).map(|m| m.len()).unwrap_or(0) / (1024 * 1024);
3452        println!(
3453            "✓ LoRA adapter saved ({} layers, {} MB) to {}",
3454            adapters.len(),
3455            size_mb,
3456            output_dir.display()
3457        );
3458
3459        Ok(())
3460    }
3461
3462    /// R-001: Save CPU embedding optimizer state (m/v buffers + step counter).
3463    ///
3464    /// Writes `optimizer_state.json` to the given directory. GPU block optimizer
3465    /// states remain on-device (D2H for 20 buffers × N blocks is deferred).
3466    pub fn save_optimizer_state(&self, dir: &std::path::Path) -> crate::Result<()> {
3467        let path = dir.join("optimizer_state.json");
3468        let m_data: Vec<Option<Vec<f32>>> = self
3469            .embed_optimizer
3470            .first_moments()
3471            .iter()
3472            .map(|opt| opt.as_ref().map(ndarray::ArrayBase::to_vec))
3473            .collect();
3474        let v_data: Vec<Option<Vec<f32>>> = self
3475            .embed_optimizer
3476            .second_moments()
3477            .iter()
3478            .map(|opt| opt.as_ref().map(ndarray::ArrayBase::to_vec))
3479            .collect();
3480        let state = serde_json::json!({
3481            "type": "adamw_cpu_embed",
3482            "step": self.embed_optimizer.step_count(),
3483            "m": m_data,
3484            "v": v_data,
3485        });
3486        let json_str = serde_json::to_string(&state).map_err(|e| {
3487            crate::error::Error::ConfigError(format!("serialize optimizer state: {e}"))
3488        })?;
3489        std::fs::write(&path, json_str)
3490            .map_err(|e| crate::error::Error::ConfigError(format!("write optimizer state: {e}")))?;
3491        Ok(())
3492    }
3493
3494    /// ENT-276: Restore LoRA adapter weights from APR checkpoint.
3495    ///
3496    /// Reads `lora.{layer}.{q,v}_proj.lora_{a,b}` tensors from the APR file
3497    /// and uploads them to the NF4 CUDA blocks, replacing the fresh random init.
3498    /// Returns (layers_restored, layers_total).
3499    pub fn restore_lora_from_apr(&mut self, apr_path: &std::path::Path) -> (usize, usize) {
3500        let reader = match aprender::serialization::apr::AprReader::open(apr_path) {
3501            Ok(r) => r,
3502            Err(_) => return (0, self.cuda_blocks.len()),
3503        };
3504
3505        let mut restored = 0usize;
3506        for (i, block) in self.cuda_blocks.iter_mut().enumerate() {
3507            let a_q =
3508                reader.read_tensor_f32(&format!("lora.{i}.q_proj.lora_a")).unwrap_or_default();
3509            let b_q =
3510                reader.read_tensor_f32(&format!("lora.{i}.q_proj.lora_b")).unwrap_or_default();
3511            let a_v =
3512                reader.read_tensor_f32(&format!("lora.{i}.v_proj.lora_a")).unwrap_or_default();
3513            let b_v =
3514                reader.read_tensor_f32(&format!("lora.{i}.v_proj.lora_b")).unwrap_or_default();
3515
3516            if a_q.is_empty() {
3517                continue; // No LoRA data for this layer in checkpoint
3518            }
3519
3520            if let Err(e) = block.upload_lora_weights(&a_q, &b_q, &a_v, &b_v) {
3521                eprintln!("Warning: failed to restore LoRA for layer {i}: {e}");
3522                continue;
3523            }
3524            restored += 1;
3525        }
3526
3527        (restored, self.cuda_blocks.len())
3528    }
3529
3530    /// ALB-096: Load CPU embedding optimizer state from APR checkpoint.
3531    ///
3532    /// Reads `__training__.embed_optimizer.{m,v}.*` tensors from the APR file.
3533    /// Returns true if state was loaded.
3534    pub fn load_optimizer_state_apr(&mut self, apr_path: &std::path::Path) -> bool {
3535        let reader = match aprender::serialization::apr::AprReader::open(apr_path) {
3536            Ok(r) => r,
3537            Err(_) => return false,
3538        };
3539
3540        // Restore step count from metadata
3541        if let Some(step_val) = reader.get_metadata("optimizer_step") {
3542            if let Some(step_str) = step_val.as_str() {
3543                if let Ok(step) = step_str.parse::<u64>() {
3544                    self.embed_optimizer.set_step_count(step);
3545                }
3546            }
3547        }
3548
3549        // Restore first moments (m)
3550        for i in 0..128 {
3551            let name = format!("__training__.embed_optimizer.m.{i}");
3552            match reader.read_tensor_f32(&name) {
3553                Ok(data) if !data.is_empty() => {
3554                    self.embed_optimizer.set_first_moment(i, ndarray::Array1::from_vec(data));
3555                }
3556                _ => break,
3557            }
3558        }
3559
3560        // Restore second moments (v)
3561        for i in 0..128 {
3562            let name = format!("__training__.embed_optimizer.v.{i}");
3563            match reader.read_tensor_f32(&name) {
3564                Ok(data) if !data.is_empty() => {
3565                    self.embed_optimizer.set_second_moment(i, ndarray::Array1::from_vec(data));
3566                }
3567                _ => break,
3568            }
3569        }
3570
3571        // ALB-118: Restore GPU block optimizer states (m/v moments for all blocks)
3572        let suffixes = [
3573            "m.w_q",
3574            "v.w_q",
3575            "m.w_k",
3576            "v.w_k",
3577            "m.w_v",
3578            "v.w_v",
3579            "m.w_o",
3580            "v.w_o",
3581            "m.w_gate",
3582            "v.w_gate",
3583            "m.w_up",
3584            "v.w_up",
3585            "m.w_down",
3586            "v.w_down",
3587            "m.input_norm",
3588            "v.input_norm",
3589            "m.post_attn_norm",
3590            "v.post_attn_norm",
3591        ];
3592        let mut blocks_restored = 0usize;
3593        for (layer_idx, state) in self.gpu_training.optimizer_states.iter_mut().enumerate() {
3594            let mut data = std::collections::HashMap::new();
3595            for suffix in &suffixes {
3596                let name = format!("__training__.block_optimizer.{layer_idx}.{suffix}");
3597                if let Ok(tensor_data) = reader.read_tensor_f32(&name) {
3598                    if !tensor_data.is_empty() {
3599                        data.insert(suffix.to_string(), tensor_data);
3600                    }
3601                }
3602            }
3603            if !data.is_empty() {
3604                let _ = state.restore_from_host(&data);
3605                blocks_restored += 1;
3606            }
3607        }
3608
3609        // ALB-118: Restore LM head optimizer state
3610        if let Ok(m_data) = reader.read_tensor_f32("__training__.lm_head_optimizer.m") {
3611            if m_data.len() == self.lm_head_m.len() {
3612                let _ = self.lm_head_m.copy_from_host(&m_data);
3613            }
3614        }
3615        if let Ok(v_data) = reader.read_tensor_f32("__training__.lm_head_optimizer.v") {
3616            if v_data.len() == self.lm_head_v.len() {
3617                let _ = self.lm_head_v.copy_from_host(&v_data);
3618            }
3619        }
3620
3621        // ALB-118: Restore final norm optimizer state
3622        if let Ok(m_data) = reader.read_tensor_f32("__training__.final_norm_optimizer.m") {
3623            if m_data.len() == self.final_norm_m.len() {
3624                let _ = self.final_norm_m.copy_from_host(&m_data);
3625            }
3626        }
3627        if let Ok(v_data) = reader.read_tensor_f32("__training__.final_norm_optimizer.v") {
3628            if v_data.len() == self.final_norm_v.len() {
3629                let _ = self.final_norm_v.copy_from_host(&v_data);
3630            }
3631        }
3632
3633        // ALB-132: Report restore results — don't silently swallow failures
3634        if blocks_restored > 0 {
3635            println!(
3636                "  ✓ GPU block optimizer states restored ({blocks_restored}/{} blocks)",
3637                self.gpu_training.optimizer_states.len()
3638            );
3639        } else if !self.gpu_training.optimizer_states.is_empty() {
3640            println!(
3641                "  [WARN] GPU block optimizer states NOT restored (0/{} blocks — zeroed m/v)",
3642                self.gpu_training.optimizer_states.len()
3643            );
3644        }
3645
3646        true
3647    }
3648
3649    /// R-001: Load CPU embedding optimizer state from `optimizer_state.json`.
3650    ///
3651    /// Returns true if state was loaded, false if file doesn't exist.
3652    pub fn load_optimizer_state(&mut self, dir: &std::path::Path) -> bool {
3653        let path = dir.join("optimizer_state.json");
3654        let data = match std::fs::read_to_string(&path) {
3655            Ok(d) => d,
3656            Err(_) => return false,
3657        };
3658        let state: serde_json::Value = match serde_json::from_str(&data) {
3659            Ok(v) => v,
3660            Err(_) => return false,
3661        };
3662        if let Some(step) = state["step"].as_u64() {
3663            self.embed_optimizer.set_step_count(step);
3664        }
3665        restore_moment_buffers(&state["m"], |idx, arr| {
3666            self.embed_optimizer.set_first_moment(idx, arr);
3667        });
3668        restore_moment_buffers(&state["v"], |idx, arr| {
3669            self.embed_optimizer.set_second_moment(idx, arr);
3670        });
3671        true
3672    }
3673}
3674
3675/// ALB-096: Infer 2D tensor shape from name and element count.
3676///
3677/// Same logic as `infer_all_tensor_shapes` in `io/save.rs` but for a single tensor.
3678#[cfg(feature = "cuda")]
3679fn infer_tensor_shape(name: &str, numel: usize, hidden_size: usize) -> Vec<usize> {
3680    if name.ends_with("layernorm.weight") || name == "model.norm.weight" {
3681        vec![numel]
3682    } else if hidden_size > 0 && numel.is_multiple_of(hidden_size) {
3683        let other_dim = numel / hidden_size;
3684        if name.ends_with("down_proj.weight") {
3685            vec![hidden_size, other_dim]
3686        } else {
3687            vec![other_dim, hidden_size]
3688        }
3689    } else {
3690        vec![numel]
3691    }
3692}
3693
3694/// Parse a JSON array of moment buffers and apply each via callback.
3695#[cfg(feature = "cuda")]
3696fn restore_moment_buffers(
3697    json_arr: &serde_json::Value,
3698    mut set_fn: impl FnMut(usize, ndarray::Array1<f32>),
3699) {
3700    let Some(arr) = json_arr.as_array() else { return };
3701    for (idx, val) in arr.iter().enumerate() {
3702        let Some(inner) = val.as_array() else { continue };
3703        let floats: Vec<f32> = inner.iter().filter_map(|v| v.as_f64().map(|f| f as f32)).collect();
3704        if !floats.is_empty() {
3705            set_fn(idx, ndarray::Array1::from_vec(floats));
3706        }
3707    }
3708}
3709
3710// ── Non-CUDA stub ──
3711
3712#[cfg(not(feature = "cuda"))]
3713pub struct CudaTransformerTrainer;
3714
3715#[cfg(not(feature = "cuda"))]
3716impl CudaTransformerTrainer {
3717    pub fn new(_config: super::config::TransformerTrainConfig) -> crate::Result<Self> {
3718        Err(crate::error::Error::ConfigError(
3719            "CUDA not available (compiled without cuda feature)".into(),
3720        ))
3721    }
3722
3723    pub fn with_model(
3724        _model: crate::transformer::Transformer,
3725        _config: super::config::TransformerTrainConfig,
3726    ) -> crate::Result<Self> {
3727        Err(crate::error::Error::ConfigError(
3728            "CUDA not available (compiled without cuda feature)".into(),
3729        ))
3730    }
3731
3732    pub fn gpu_name(&self) -> String {
3733        unreachable!("CudaTransformerTrainer stub should never be instantiated")
3734    }
3735}
3736
3737#[cfg(test)]
3738mod tests {
3739    #[test]
3740    #[cfg(not(feature = "cuda"))]
3741    fn test_cuda_trainer_stub_returns_error() {
3742        use super::super::config::TransformerTrainConfig;
3743        use crate::transformer::TransformerConfig;
3744
3745        let mc = TransformerConfig::tiny();
3746        let config = TransformerTrainConfig::new(mc);
3747        let result = super::CudaTransformerTrainer::new(config);
3748        assert!(result.is_err());
3749    }
3750}