Skip to content

Commit ff98a05

Browse files
authored
fix: eliminate 3rd forward pass in distillation (hidden state MSE) (#74)
The distillation loop was running 3 forward passes per chunk: 1. forward_gpu for logits (removed in PR #73) 2. forward_train for activations (kept) 3. forward_with_hidden_states for MSE loss ← REMOVED NOW Now forward_train collects per-layer hidden states in TrainForwardOutput, and the MSE loss uses those directly. One forward pass per chunk. Impact: training step time drops by ~33% (was 2 forwards, now 1). Estimated: ~100s/step → ~67s/step.
1 parent cae68a0 commit ff98a05

2 files changed

Lines changed: 14 additions & 9 deletions

File tree

src/main.rs

Lines changed: 6 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -1829,14 +1829,13 @@ async fn cmd_distill(
18291829
);
18301830

18311831
// Hidden state matching loss: MSE between teacher and student per-layer states
1832-
let (hidden_mse_loss, _student_states) = if chunk_idx < frozen_states_per_chunk.len() {
1832+
// Uses per_layer_hidden collected during forward_train (no extra forward pass)
1833+
let hidden_mse_loss = if chunk_idx < frozen_states_per_chunk.len() && !train_output.per_layer_hidden.is_empty() {
18331834
let teacher_states = &frozen_states_per_chunk[chunk_idx];
1834-
let student_states = student.forward_with_hidden_states(batch_tokens);
1835-
let mse = if teacher_states.len() == student_states.len() && !teacher_states.is_empty() {
1836-
let _hd = config.hidden_dim;
1835+
let student_states = &train_output.per_layer_hidden;
1836+
if teacher_states.len() == student_states.len() && !teacher_states.is_empty() {
18371837
let mut total_mse = 0.0f32;
18381838
let mut num_layers_compared = 0usize;
1839-
// Compare at every 5th layer to save compute
18401839
for layer_idx in (0..teacher_states.len()).step_by(5) {
18411840
let t_state = &teacher_states[layer_idx];
18421841
let s_state = &student_states[layer_idx];
@@ -1851,10 +1850,9 @@ async fn cmd_distill(
18511850
if num_layers_compared > 0 { total_mse / num_layers_compared as f32 } else { 0.0 }
18521851
} else {
18531852
0.0
1854-
};
1855-
(mse, Some(student_states))
1853+
}
18561854
} else {
1857-
(0.0, None)
1855+
0.0
18581856
};
18591857

18601858
// MoE Load balance loss: num_experts × Σ(f_i × P_i)

src/training/backward.rs

Lines changed: 8 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,8 @@ pub struct TrainForwardOutput {
1616
pub activations: Vec<LayerActivations>,
1717
/// Post-final-norm hidden states [seq, hidden_dim] — input to lm_head
1818
pub final_hidden: Vec<f32>,
19+
/// Per-layer post-residual hidden states (for hidden state MSE loss)
20+
pub per_layer_hidden: Vec<Vec<f32>>,
1921
}
2022

2123
impl CpuBlockAttnResModel {
@@ -62,6 +64,8 @@ impl CpuBlockAttnResModel {
6264
let mut activations = Vec::with_capacity(self.num_layers);
6365
let lora_m = self.lora_manager.as_ref();
6466

67+
let mut per_layer_hidden: Vec<Vec<f32>> = Vec::with_capacity(self.num_layers);
68+
6569
for (layer_idx, layer) in self.layers.iter().enumerate() {
6670
let ple_slice = ple_precomputed.as_ref().map(|pre| {
6771
let mut slice = vec![0.0f32; seq * ple_dim];
@@ -218,6 +222,9 @@ impl CpuBlockAttnResModel {
218222
expert_activations: layer_expert_act.unwrap_or_default(),
219223
});
220224

225+
// Store per-layer hidden state for MSE loss
226+
per_layer_hidden.push(hidden.clone());
227+
221228
for t in 0..seq { for d in 0..hd { partial_sum[d] += hidden[t * hd + d]; } }
222229
if self.is_block_boundary(layer_idx) {
223230
for d in 0..hd { partial_sum[d] /= ((seq) * (self.block_config.layers_per_block)) as f32; }
@@ -236,7 +243,7 @@ impl CpuBlockAttnResModel {
236243
for l in logits.iter_mut() { *l = (*l / cap).tanh() * cap; }
237244
}
238245

239-
TrainForwardOutput { logits, routing_data, activations, final_hidden }
246+
TrainForwardOutput { logits, routing_data, activations, final_hidden, per_layer_hidden }
240247
}
241248
}
242249

0 commit comments

Comments
 (0)