1616extern " C" {
1717void LaunchBesselCorrection (float * data, float factor, int count,
1818 musaStream_t stream);
19+ void LaunchVarianceToInvStd (const float * variance, float * inv_std,
20+ float epsilon, int count, musaStream_t stream);
1921}
2022
2123namespace tensorflow {
@@ -260,7 +262,19 @@ class MusaFusedBatchNormGradOp : public MusaOpKernel {
260262
261263 mTensor mt_scale = CreateMTensor (scale, mFormat ::NCHW );
262264 mTensor mt_saved_mean = CreateMTensor (saved_mean, mFormat ::NCHW );
263- mTensor mt_saved_var = CreateMTensor (saved_var, mFormat ::NCHW );
265+ // TensorFlow's reserve_space_2 contains population variance. muDNN's
266+ // training backward kernels interpret their `v` input as reciprocal
267+ // standard deviation and use it directly (including inv_std^3 terms).
268+ // Passing variance happens to look plausible when variance is near one,
269+ // but makes dx scale catastrophically for high-variance inputs.
270+ Tensor saved_inv_std;
271+ OP_REQUIRES_OK (ctx,
272+ ctx->allocate_temp (DT_FLOAT , saved_var.shape (),
273+ &saved_inv_std));
274+ LaunchVarianceToInvStd (
275+ saved_var.flat <float >().data (), saved_inv_std.flat <float >().data (),
276+ epsilon_, static_cast <int >(saved_var.NumElements ()), stream);
277+ mTensor mt_saved_inv_std = CreateMTensor (saved_inv_std, mFormat ::NCHW );
264278
265279 mTensor mt_d_scale = CreateMTensor (*d_scale, mFormat ::NCHW );
266280 mTensor mt_d_offset = CreateMTensor (*d_offset, mFormat ::NCHW );
@@ -274,7 +288,7 @@ class MusaFusedBatchNormGradOp : public MusaOpKernel {
274288
275289 mStatus status = bn_op.RunBwd (
276290 handle, mt_dx, mt_d_mean, mt_d_var, mt_d_scale, mt_d_offset, mt_x,
277- mt_dy, mt_saved_mean, mt_saved_var , mt_scale, maintainer);
291+ mt_dy, mt_saved_mean, mt_saved_inv_std , mt_scale, maintainer);
278292
279293 OP_REQUIRES (ctx, status == mStatus ::SUCCESS ,
280294 errors::Internal (" MUSA BN Backward failed." ));
0 commit comments