1#[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#[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 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 if stream.synchronize().is_err() {
116 return (1.0, 0.0);
117 }
118
119 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); }
127 }
128
129 let grad_norm = total_sq.sqrt() as f32; let scale = if grad_norm > max_norm { max_norm / grad_norm } else { 1.0 };
131 (scale, grad_norm)
132}
133
134#[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#[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 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 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 let scale_ptr = state.scale_buf.as_ptr(); 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#[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#[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#[cfg(feature = "cuda")]
294struct GpuPretrainState {
295 layer_inputs: Vec<GpuBuffer<f32>>,
297 saved_layer_mask: Vec<bool>,
301 recompute_buf: Option<GpuBuffer<f32>>,
305 final_norm_weight: GpuBuffer<f32>,
307 blocks_output: GpuBuffer<f32>,
309 grad_buf_a: GpuBuffer<f32>,
311 grad_buf_b: GpuBuffer<f32>,
313 grad_final_norm_weight: GpuBuffer<f32>,
315 norm_output: GpuBuffer<f32>,
317 logits_buf: GpuBuffer<f32>,
319 lm_head_grad_hidden: GpuBuffer<f32>,
321 optimizer_states: Vec<GpuBlockOptimizerState>,
323 step: u32,
325}
326
327#[cfg(feature = "cuda")]
338pub struct CudaTransformerTrainer {
339 model: Transformer,
341 cuda_trainer: CudaTrainer,
343 cuda_blocks: Vec<CudaBlock>,
345 cuda_grad_workspace: CudaGradWorkspace,
347 nf4_shared_scratch: Option<CudaBlockScratch>,
349 nf4_lora_grad_workspace: Option<CudaLoraGradWorkspace>,
351 nf4_lora_optimizer_states: Option<Vec<GpuLoraOptimizerState>>,
353 gpu_training: GpuPretrainState,
355 lm_head_weight_gpu: GpuBuffer<f32>,
357 lm_head_grad_gpu: GpuBuffer<f32>,
359 lm_head_m: GpuBuffer<f32>,
361 lm_head_v: GpuBuffer<f32>,
363 final_norm_m: GpuBuffer<f32>,
365 final_norm_v: GpuBuffer<f32>,
367 embed_optimizer: AdamW,
369 config: TransformerTrainConfig,
371 pub metrics: MetricsTracker,
373 step: usize,
375 accumulated_loss: f32,
377 accumulated_batches: usize,
379 last_grad_norm: f32,
381 last_embed_grad_norm: f32,
383 grad_accum: Option<super::grad_accumulator::PerBlockGradientAccumulator>,
386 gpu_grad_accum: Option<super::gpu_grad_accumulator::GpuGradientAccumulator>,
389 grad_scaler: GradScaler,
393 profiler: StepProfiler,
396 fwd_scratch_a: GpuBuffer<f32>,
399 fwd_scratch_b: GpuBuffer<f32>,
400 h2d_staging: Vec<f32>,
403 d2h_staging: Vec<f32>,
408 fused_clip: Option<FusedClipState>,
411 final_norm_zero_buf: Vec<f32>,
415}
416
417#[cfg(feature = "cuda")]
418impl CudaTransformerTrainer {
419 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 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 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 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 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 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 {
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 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 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 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 let buf_size = max_seq_len * hidden_size;
556 let logits_size = max_seq_len * vocab_size;
557
558 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 };
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 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 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 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 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 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 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 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 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 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 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 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 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 let embed_optimizer =
735 AdamW::new(config.lr, config.beta1, config.beta2, 1e-8, config.weight_decay);
736
737 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 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 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 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 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 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 #[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 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 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 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 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 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 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 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 self.profiler.begin(StepProfiler::LOSS);
1083 let stream = self.cuda_trainer.stream();
1084
1085 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 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 if !loss_val.is_finite() {
1111 return None;
1112 }
1113 self.profiler.end(StepProfiler::LOSS);
1114
1115 if let Some(grad_output_is_a) =
1124 self.gpu_backward(seq_len, hidden_size, vocab_size, accumulate_only)
1125 {
1126 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 #[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 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 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 self.profiler.begin(StepProfiler::FORWARD);
1174 let mut input_is_a = true; for (i, block) in self.cuda_blocks.iter_mut().enumerate() {
1176 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 unsafe {
1194 self.gpu_training.layer_inputs[i]
1195 .copy_from_buffer_async(&*input_ptr, stream)
1196 .ok()?;
1197 }
1198 }
1199 self.profiler.begin_layer();
1202 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 let final_output: &GpuBuffer<f32> =
1221 if input_is_a { &self.fwd_scratch_a } else { &self.fwd_scratch_b };
1222
1223 self.profiler.begin(StepProfiler::NORM_LM);
1226 unsafe {
1228 self.gpu_training.blocks_output.copy_from_buffer_async(final_output, stream).ok()?;
1229 }
1230
1231 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 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 self.profiler.end(StepProfiler::NORM_LM);
1263
1264 Some(())
1265 }
1266
1267 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 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 let grad_output_is_a = self.gpu_backward(seq_len, hidden_size, vocab_size, false)?;
1338 self.embed_backward(input_ids, seq_len, hidden_size, vocab_size, grad_output_is_a);
1342
1343 Some(())
1344 }
1345
1346 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 self.gpu_forward(input_ids, seq_len, hidden_size, vocab_size)?;
1365
1366 let stream = self.cuda_trainer.stream();
1368 stream.synchronize().ok()?;
1369
1370 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 #[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 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(()); }
1409
1410 let recompute_buf = gpu_training.recompute_buf.as_mut()?;
1413 unsafe {
1415 recompute_buf
1416 .copy_from_buffer_async(&gpu_training.layer_inputs[seg_start], stream)
1417 .ok()?;
1418 }
1419
1420 for i in seg_start..target_layer {
1431 if i == seg_start {
1432 let recompute_ptr: *const GpuBuffer<f32> = recompute_buf;
1434 let li = &mut gpu_training.layer_inputs;
1435 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 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 #[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 let beta1 = self.config.beta1;
1485 let beta2 = self.config.beta2;
1486 let weight_decay = self.config.weight_decay;
1487
1488 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 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(); self.last_grad_norm = lm_norm; if std::env::var("ENTRENAR_TRACE_GRADIENTS").is_ok() {
1527 eprintln!("[grad-trace] lm_head gnorm={lm_norm:.6}");
1528 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 self.profiler.begin(StepProfiler::NORM_BWD);
1546 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 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 if accumulate_only {
1577 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 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 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 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 let _output_scratch_ptr: *mut GpuBuffer<f32> = if grad_output_is_a {
1650 grad_b_ptr } else {
1652 grad_a_ptr
1653 };
1654 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, 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 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 {
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 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 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 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 if accumulate_only {
1751 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 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 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 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 #[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 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 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 stream.synchronize().ok()?;
1932
1933 self.gpu_training.step += 1;
1934 let step = self.gpu_training.step;
1935
1936 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 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 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 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 let accum = self.grad_accum.as_mut()?;
2015 accum.average();
2016
2017 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 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 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 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 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 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 accum.zero_all();
2138 Some(())
2139 }
2140
2141 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 let grad_norm = match squared_sum_cuda(buf, n, stream) {
2161 Ok(norm) => norm,
2162 Err(_) => {
2163 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 #[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 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 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 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; 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 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 fn optimizer_step(&mut self) {
2256 self.grad_scaler.update(true);
2260
2261 self.embed_optimizer.set_lr(self.current_lr());
2263 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 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 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 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 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 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 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 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 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 pub fn train_epoch(&mut self, batches: &[LMBatch]) -> f32 {
2412 self.train_epoch_with_callback(batches, |_, _, _| {})
2413 }
2414
2415 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 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 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 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 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 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 pub(crate) fn grad_accum_ref(
2527 &self,
2528 ) -> Option<&super::grad_accumulator::PerBlockGradientAccumulator> {
2529 self.grad_accum.as_ref()
2530 }
2531
2532 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 pub(crate) fn config(&self) -> &TransformerTrainConfig {
2541 &self.config
2542 }
2543
2544 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 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 pub fn reached_max_steps(&self) -> bool {
2556 self.config.max_steps.is_some_and(|max| self.step >= max)
2557 }
2558
2559 pub fn step(&self) -> usize {
2561 self.step
2562 }
2563
2564 pub fn set_initial_step(&mut self, step: usize) {
2570 self.step = step;
2571 self.gpu_training.step = step as u32;
2572 }
2573
2574 pub fn set_max_steps(&mut self, max_steps: usize) {
2580 self.config.max_steps = Some(max_steps);
2581 }
2582
2583 pub fn current_lr(&self) -> f32 {
2589 let base_lr = self.config.lr;
2590 if self.step < self.config.warmup_steps {
2591 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 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 base_lr
2605 }
2606 }
2607
2608 pub fn enable_profiler(&mut self, interval: usize) {
2619 self.profiler = StepProfiler::new(true, interval);
2620 }
2621
2622 pub fn print_profiler_report(&self) {
2624 self.profiler.print_report();
2625 }
2626
2627 pub fn last_grad_norm(&self) -> f32 {
2629 self.last_grad_norm
2630 }
2631
2632 pub fn param_grad_norms(&self) -> (f32, f32) {
2635 (self.last_grad_norm, self.last_embed_grad_norm)
2636 }
2637
2638 pub fn num_params(&self) -> usize {
2640 self.model.parameters().iter().map(|t| t.len()).sum()
2641 }
2642
2643 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 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 } 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 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 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 pub fn model(&self) -> &Transformer {
2708 &self.model
2709 }
2710
2711 pub fn model_mut(&mut self) -> &mut Transformer {
2713 &mut self.model
2714 }
2715
2716 pub fn is_mixed_precision(&self) -> bool {
2718 self.config.precision_config.is_mixed()
2719 }
2720
2721 pub fn grad_scaler(&self) -> &GradScaler {
2723 &self.grad_scaler
2724 }
2725
2726 pub fn is_checkpointing(&self) -> bool {
2728 self.config.checkpoint_config.enabled
2729 }
2730
2731 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 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 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 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 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 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 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 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 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 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 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 let shapes = infer_all_tensor_shapes(¶ms);
2934 for (tname, tensor) in ¶ms {
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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 for (tensor_name, data) in ¶m_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 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 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 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 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 writer
3317 .write(path)
3318 .map_err(|e| crate::error::Error::Serialization(format!("APR save failed: {e}")))?;
3319
3320 Ok(())
3321 })
3322 }
3323
3324 pub fn gpu_name(&self) -> String {
3326 self.cuda_trainer.device_name()
3327 }
3328
3329 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(()); }
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, };
3365
3366 if a_q.is_empty() && a_v.is_empty() {
3367 continue;
3368 }
3369
3370 if !a_q.is_empty() {
3372 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 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 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 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 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 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; }
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 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 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 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 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 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 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 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 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 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#[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#[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#[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}