# Compute weight gradients with accumulation original_grad = weight_grad.clone() deep_gemm.wgrad_gemm_fp8_fp8_fp32_nt(grad_input, act_input, weight_grad) # Verify: grad_output^T @ activations + original_grad reference_update = grad_output.float().t() @ activations.float() expected_grad = original_grad + reference_update error = torch.abs(weight_grad - expected_grad).max().item() relative_error = error / torch.abs(expected_grad).max().item() print(f"Weight gradient error: {error:.6f}") print(f"Relative error: {relative_error*100:.3f}%") print("✓ Weight gradient computation successful!")