pub fn ce_bwd( probs: &[f32], ids: &[u32], rows: usize, cols: usize, scale: f32, ) -> Vec<f32>
Cross-entropy backward: dlogits = (probs - onehot(ids)) * scale.