SIGN IN SIGN UP

[Operator Mechanism] Fix p_norm_grad kernel precision compatibility for GPU (#79481)

* Use double precision for the norm_p value in norm reduce GPU kernel

* Add torch-compatible precision mode to p_norm_grad GPU kernel

Introduce a FLAGS_use_accuracy_compatible_kernel guarded code path
that reproduces PyTorch's intermediate rounding behavior (round back
to storage dtype between fused ops) for fp16/bf16 p_norm backward.
This ensures bitwise alignment with torch when the flag is enabled,
while preserving the existing higher-precision fused path as default.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>

* Add a comment and fix codestyle

* fix a bug

* Fix precision in compute_pow_like_kernel_torch_compat for fp16

Align pow special cases with torch's Half operator* behavior:
- exp=-1: compute reciprocal in float32 (MT) instead of relying on
  float16 operator/
- exp=-2: compute val*val in float32 with RoundToStorage to match
  torch's Half*Half (float32 promotion + truncate), then divide in
  float32
- exp=2: compute in float32 to avoid __hmul on CUDA_ARCH >= 530
- exp=3: compute in float32 with intermediate RoundToStorage to match
  torch's chained Half*Half*Half behavior

Also move RoundToStorage definition above compute_pow_like_kernel_torch_compat
so it can be used within that function, and remove dead commented-out code.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>

* Add unit tests for fused p_norm_grad CUDA kernels

Add TestPnormGradFusedKernels class covering all dispatch branches:
- p=0 (SetConstant), p=1 (P1Kernel), p=2 (P2Kernel)
- 0<p<1 (PLessThan1Kernel), 1<p<2 (PBetween1And2Kernel), p>2 (PGreaterThan2Kernel)
- p=inf/-inf (ReduceAMaxGrad)
- reduce_all, norm==0 masked_fill, keepdim, 3D broadcasting
- fp16/bf16 dtype on GPU

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>

* Add Compat-path GPU tests for p_norm_grad with fp16/bf16

Replace TestPnormGradFusedKernels (CPU float64 numerical-diff tests)
with TestPnormGradCompatKernel that enables
FLAGS_use_accuracy_compatible_kernel and verifies the fused kernel's
torch-compatible rounding behavior against a numpy reference
simulating per-op storage-precision truncation.

Covers: p=2, p=3, p=4, p=1.5 with float16 and bfloat16, axis
variants, and reduce_all path. Assertions are bitwise (atol=0).

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>

* Fix bf16 reference in TestPnormGradCompatKernel to use RNE rounding

The bf16 _round_to_storage was using convert_float_to_uint16 which does
truncation (>> 16), but CUDA __float2bfloat16 uses round-to-nearest-even.
This caused ~65% mismatch in bf16 tests. Replace with vectorized RNE
matching the kernel behavior. Also remove ml_dtypes dependency, add proper
bf16 skip conditions, and unify rounding calls through _round_to_storage.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>

* Add p<1 path at test

* Use double precision for porder in UnsignedPowFunctor to to keep consistent with ReduceGpuKernel in the non-Windows path

* Revert "Use double precision for porder in UnsignedPowFunctor to to keep consistent with ReduceGpuKernel in the non-Windows path"

This reverts commit 68a5c60b26203e5b68f81993eaa19465b6eae048.

* Revert "Use double precision for the norm_p value in norm reduce GPU kernel"

This reverts commit dc0c83b77db3f33be1d5d5709d1ef5686e56ceec.

* Skip bf16 compat grad tests on ROCm/DCU

---------

Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com>
O
omoYang committed
1dfe2f2b9c53ee1a52bfec33d567d753d0694d3d
Parent: 1587050
Committed by GitHub <noreply@github.com> on 7/22/2026, 6:36:40 AM