diff --git a/tensorflow/compiler/mlir/lite/BUILD b/tensorflow/compiler/mlir/lite/BUILD index cdb76a3f91a6aa..021ef2f823af4c 100644 --- a/tensorflow/compiler/mlir/lite/BUILD +++ b/tensorflow/compiler/mlir/lite/BUILD @@ -1554,6 +1554,7 @@ cc_library( "transforms/decompose_hybrid_quantization.cc", "transforms/default_quant_params.cc", "transforms/fold_stablehlo_constant_transforms_pass.cc", + "transforms/fuse_a4w2_drq_fully_connected_pass.cc", "transforms/generated_post_quantize.inc", "transforms/generated_quantize.inc", "transforms/lower_quant_annotations_helper.cc", diff --git a/tensorflow/compiler/mlir/lite/python/stablehlo_tfl_pipeline.cc b/tensorflow/compiler/mlir/lite/python/stablehlo_tfl_pipeline.cc index 9518f575636555..898bb83666c83a 100644 --- a/tensorflow/compiler/mlir/lite/python/stablehlo_tfl_pipeline.cc +++ b/tensorflow/compiler/mlir/lite/python/stablehlo_tfl_pipeline.cc @@ -301,6 +301,11 @@ void AddPipelinePasses(mlir::OpPassManager& pass_manager, mlir::createCanonicalizerPass()); pass_manager.addNestedPass(mlir::createCSEPass()); pass_manager.addPass(CreatePruneDeadResourcesPass()); + pass_manager.addNestedPass( + mlir::TFL::CreateFuseA4W2DRQFullyConnectedPass()); + pass_manager.addNestedPass( + mlir::createCanonicalizerPass()); + pass_manager.addPass(mlir::createSymbolDCEPass()); pass_manager.addPass(mlir::TFL::CreateCleanupOptimizationBarrierPass()); pass_manager.addPass(mlir::odml::createLegalizeStablehloToVhloPass()); pass_manager.addPass(mlir::createReconcileUnrealizedCastsPass()); diff --git a/tensorflow/compiler/mlir/lite/tests/fuse-a4w2-drq-fully-connected.mlir b/tensorflow/compiler/mlir/lite/tests/fuse-a4w2-drq-fully-connected.mlir new file mode 100644 index 00000000000000..9bac2271e6a6a8 --- /dev/null +++ b/tensorflow/compiler/mlir/lite/tests/fuse-a4w2-drq-fully-connected.mlir @@ -0,0 +1,241 @@ +// Copyright 2026 The TensorFlow Authors. All Rights Reserved. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +// ============================================================================== + +// RUN: litert-opt %s -tfl-fuse-a4w2-drq-fully-connected | FileCheck %s + +// Every property the `cint2_fp32_int4_e8m0_drq` contract fixes is baked into the +// runtime kernel rather than carried in the flatbuffer, so each negative test +// below pins a case where matching would produce a model that runs but +// computes something other than what the graph says. + +// ----------------------------------------------------------------------------- +// Positive case. +// ----------------------------------------------------------------------------- + +// CHECK-LABEL: @FuseA4W2Drq +func.func @FuseA4W2Drq(%arg0: tensor<4x64xf32>) -> tensor<4x32xf32> { + %bias = "tfl.no_value"() {value} : () -> none + %w = "tfl.pseudo_const"() {value = dense<1> : tensor<32x64xi2>} : () -> tensor<32x64xi2> + %w_scale = "tfl.pseudo_const"() {value = dense<2.500000e-01> : tensor<32x1xf32>} : () -> tensor<32x1xf32> + %w_zp = "tfl.pseudo_const"() {value = dense<-5.000000e-01> : tensor<32x1xf32>} : () -> tensor<32x1xf32> + %act:3 = "tfl.blockwise_quantize"(%arg0) {block_shape = [1, 32], scale_type = f8E8M0FNU, symmetric = true, range_dilation = 1.500000e+00 : f32} : (tensor<4x64xf32>) -> (tensor<4x64xi4>, tensor<4x2xf8E8M0FNU>, none) + %act_dq = "tfl.blockwise_dequantize"(%act#0, %act#1, %act#2) {block_shape = [1, 32], symmetric = true} : (tensor<4x64xi4>, tensor<4x2xf8E8M0FNU>, none) -> tensor<4x64xf32> + %w_dq = "tfl.blockwise_dequantize"(%w, %w_scale, %w_zp) {block_shape = [1, 64], symmetric = false} : (tensor<32x64xi2>, tensor<32x1xf32>, tensor<32x1xf32>) -> tensor<32x64xf32> + %0 = "tfl.fully_connected"(%act_dq, %w_dq, %bias) {asymmetric_quantize_inputs = false, fused_activation_function = "NONE", keep_num_dims = false, weights_format = "DEFAULT"} : (tensor<4x64xf32>, tensor<32x64xf32>, none) -> tensor<4x32xf32> + func.return %0 : tensor<4x32xf32> + + // The blockwise ops collapse away entirely; the weights become a per-axis + // i2 qconst and the float activation feeds the fully_connected directly. + // CHECK-NOT: tfl.blockwise_quantize + // CHECK-NOT: tfl.blockwise_dequantize + // CHECK: %[[QW:.*]] = "tfl.pseudo_qconst"() <{qtype = tensor<32x64x!quant.uniform) -> tensor<4x32xf32> { + %bias = "tfl.no_value"() {value} : () -> none + %w = "tfl.pseudo_const"() {value = dense<1> : tensor<32x64xi2>} : () -> tensor<32x64xi2> + %w_scale = "tfl.pseudo_const"() {value = dense<2.500000e-01> : tensor<32x1xf32>} : () -> tensor<32x1xf32> + %act:3 = "tfl.blockwise_quantize"(%arg0) {block_shape = [1, 32], scale_type = f8E8M0FNU, symmetric = true, range_dilation = 1.500000e+00 : f32} : (tensor<4x64xf32>) -> (tensor<4x64xi4>, tensor<4x2xf8E8M0FNU>, none) + %act_dq = "tfl.blockwise_dequantize"(%act#0, %act#1, %act#2) {block_shape = [1, 32], symmetric = true} : (tensor<4x64xi4>, tensor<4x2xf8E8M0FNU>, none) -> tensor<4x64xf32> + %none = "tfl.no_value"() {value} : () -> none + %w_dq = "tfl.blockwise_dequantize"(%w, %w_scale, %none) {block_shape = [1, 64], symmetric = true} : (tensor<32x64xi2>, tensor<32x1xf32>, none) -> tensor<32x64xf32> + %0 = "tfl.fully_connected"(%act_dq, %w_dq, %bias) {asymmetric_quantize_inputs = false, fused_activation_function = "NONE", keep_num_dims = false, weights_format = "DEFAULT"} : (tensor<4x64xf32>, tensor<32x64xf32>, none) -> tensor<4x32xf32> + func.return %0 : tensor<4x32xf32> + + // CHECK: tfl.blockwise_dequantize + // CHECK-NOT: tfl.quant_spec +} + +// ----------------------------------------------------------------------------- +// Negative: an asymmetric weight zero point that is not the centered -0.5. +// ----------------------------------------------------------------------------- + +// CHECK-LABEL: @NoFuseNonCenteredZeroPoint +func.func @NoFuseNonCenteredZeroPoint(%arg0: tensor<4x64xf32>) -> tensor<4x32xf32> { + %bias = "tfl.no_value"() {value} : () -> none + %w = "tfl.pseudo_const"() {value = dense<1> : tensor<32x64xi2>} : () -> tensor<32x64xi2> + %w_scale = "tfl.pseudo_const"() {value = dense<2.500000e-01> : tensor<32x1xf32>} : () -> tensor<32x1xf32> + %w_zp = "tfl.pseudo_const"() {value = dense<-2.500000e-01> : tensor<32x1xf32>} : () -> tensor<32x1xf32> + %act:3 = "tfl.blockwise_quantize"(%arg0) {block_shape = [1, 32], scale_type = f8E8M0FNU, symmetric = true, range_dilation = 1.500000e+00 : f32} : (tensor<4x64xf32>) -> (tensor<4x64xi4>, tensor<4x2xf8E8M0FNU>, none) + %act_dq = "tfl.blockwise_dequantize"(%act#0, %act#1, %act#2) {block_shape = [1, 32], symmetric = true} : (tensor<4x64xi4>, tensor<4x2xf8E8M0FNU>, none) -> tensor<4x64xf32> + %w_dq = "tfl.blockwise_dequantize"(%w, %w_scale, %w_zp) {block_shape = [1, 64], symmetric = false} : (tensor<32x64xi2>, tensor<32x1xf32>, tensor<32x1xf32>) -> tensor<32x64xf32> + %0 = "tfl.fully_connected"(%act_dq, %w_dq, %bias) {asymmetric_quantize_inputs = false, fused_activation_function = "NONE", keep_num_dims = false, weights_format = "DEFAULT"} : (tensor<4x64xf32>, tensor<32x64xf32>, none) -> tensor<4x32xf32> + func.return %0 : tensor<4x32xf32> + + // CHECK: tfl.blockwise_dequantize + // CHECK-NOT: tfl.quant_spec +} + +// ----------------------------------------------------------------------------- +// Negative: an f32 activation scale. +// +// The kernel rounds the scale up to a power of two, so an unconstrained f32 +// scale would quantize the activations to different values. +// ----------------------------------------------------------------------------- + +// CHECK-LABEL: @NoFuseFloat32ActScale +func.func @NoFuseFloat32ActScale(%arg0: tensor<4x64xf32>) -> tensor<4x32xf32> { + %bias = "tfl.no_value"() {value} : () -> none + %w = "tfl.pseudo_const"() {value = dense<1> : tensor<32x64xi2>} : () -> tensor<32x64xi2> + %w_scale = "tfl.pseudo_const"() {value = dense<2.500000e-01> : tensor<32x1xf32>} : () -> tensor<32x1xf32> + %w_zp = "tfl.pseudo_const"() {value = dense<-5.000000e-01> : tensor<32x1xf32>} : () -> tensor<32x1xf32> + %act:3 = "tfl.blockwise_quantize"(%arg0) {block_shape = [1, 32], scale_type = f32, symmetric = true, range_dilation = 1.500000e+00 : f32} : (tensor<4x64xf32>) -> (tensor<4x64xi4>, tensor<4x2xf32>, none) + %act_dq = "tfl.blockwise_dequantize"(%act#0, %act#1, %act#2) {block_shape = [1, 32], symmetric = true} : (tensor<4x64xi4>, tensor<4x2xf32>, none) -> tensor<4x64xf32> + %w_dq = "tfl.blockwise_dequantize"(%w, %w_scale, %w_zp) {block_shape = [1, 64], symmetric = false} : (tensor<32x64xi2>, tensor<32x1xf32>, tensor<32x1xf32>) -> tensor<32x64xf32> + %0 = "tfl.fully_connected"(%act_dq, %w_dq, %bias) {asymmetric_quantize_inputs = false, fused_activation_function = "NONE", keep_num_dims = false, weights_format = "DEFAULT"} : (tensor<4x64xf32>, tensor<32x64xf32>, none) -> tensor<4x32xf32> + func.return %0 : tensor<4x32xf32> + + // CHECK: tfl.blockwise_quantize + // CHECK-NOT: tfl.quant_spec +} + +// ----------------------------------------------------------------------------- +// Negative: 8 bit activations. +// +// The spec is a4w2; the kernel clamps to the 4 bit range regardless. +// ----------------------------------------------------------------------------- + +// CHECK-LABEL: @NoFuseInt8Activations +func.func @NoFuseInt8Activations(%arg0: tensor<4x64xf32>) -> tensor<4x32xf32> { + %bias = "tfl.no_value"() {value} : () -> none + %w = "tfl.pseudo_const"() {value = dense<1> : tensor<32x64xi2>} : () -> tensor<32x64xi2> + %w_scale = "tfl.pseudo_const"() {value = dense<2.500000e-01> : tensor<32x1xf32>} : () -> tensor<32x1xf32> + %w_zp = "tfl.pseudo_const"() {value = dense<-5.000000e-01> : tensor<32x1xf32>} : () -> tensor<32x1xf32> + %act:3 = "tfl.blockwise_quantize"(%arg0) {block_shape = [1, 32], scale_type = f8E8M0FNU, symmetric = true, range_dilation = 1.500000e+00 : f32} : (tensor<4x64xf32>) -> (tensor<4x64xi8>, tensor<4x2xf8E8M0FNU>, none) + %act_dq = "tfl.blockwise_dequantize"(%act#0, %act#1, %act#2) {block_shape = [1, 32], symmetric = true} : (tensor<4x64xi8>, tensor<4x2xf8E8M0FNU>, none) -> tensor<4x64xf32> + %w_dq = "tfl.blockwise_dequantize"(%w, %w_scale, %w_zp) {block_shape = [1, 64], symmetric = false} : (tensor<32x64xi2>, tensor<32x1xf32>, tensor<32x1xf32>) -> tensor<32x64xf32> + %0 = "tfl.fully_connected"(%act_dq, %w_dq, %bias) {asymmetric_quantize_inputs = false, fused_activation_function = "NONE", keep_num_dims = false, weights_format = "DEFAULT"} : (tensor<4x64xf32>, tensor<32x64xf32>, none) -> tensor<4x32xf32> + func.return %0 : tensor<4x32xf32> + + // CHECK: tfl.blockwise_quantize + // CHECK-NOT: tfl.quant_spec +} + +// ----------------------------------------------------------------------------- +// Negative: an activation block size other than 32. +// ----------------------------------------------------------------------------- + +// CHECK-LABEL: @NoFuseWrongActBlockSize +func.func @NoFuseWrongActBlockSize(%arg0: tensor<4x64xf32>) -> tensor<4x32xf32> { + %bias = "tfl.no_value"() {value} : () -> none + %w = "tfl.pseudo_const"() {value = dense<1> : tensor<32x64xi2>} : () -> tensor<32x64xi2> + %w_scale = "tfl.pseudo_const"() {value = dense<2.500000e-01> : tensor<32x1xf32>} : () -> tensor<32x1xf32> + %w_zp = "tfl.pseudo_const"() {value = dense<-5.000000e-01> : tensor<32x1xf32>} : () -> tensor<32x1xf32> + %act:3 = "tfl.blockwise_quantize"(%arg0) {block_shape = [1, 16], scale_type = f8E8M0FNU, symmetric = true, range_dilation = 1.500000e+00 : f32} : (tensor<4x64xf32>) -> (tensor<4x64xi4>, tensor<4x4xf8E8M0FNU>, none) + %act_dq = "tfl.blockwise_dequantize"(%act#0, %act#1, %act#2) {block_shape = [1, 16], symmetric = true} : (tensor<4x64xi4>, tensor<4x4xf8E8M0FNU>, none) -> tensor<4x64xf32> + %w_dq = "tfl.blockwise_dequantize"(%w, %w_scale, %w_zp) {block_shape = [1, 64], symmetric = false} : (tensor<32x64xi2>, tensor<32x1xf32>, tensor<32x1xf32>) -> tensor<32x64xf32> + %0 = "tfl.fully_connected"(%act_dq, %w_dq, %bias) {asymmetric_quantize_inputs = false, fused_activation_function = "NONE", keep_num_dims = false, weights_format = "DEFAULT"} : (tensor<4x64xf32>, tensor<32x64xf32>, none) -> tensor<4x32xf32> + func.return %0 : tensor<4x32xf32> + + // CHECK: tfl.blockwise_quantize + // CHECK-NOT: tfl.quant_spec +} + +// ----------------------------------------------------------------------------- +// Negative: sub-channel weights, i.e. more than one block along the +// contracting axis. The collapsed op can only express one scale per channel. +// ----------------------------------------------------------------------------- + +// CHECK-LABEL: @NoFuseSubChannelWeights +func.func @NoFuseSubChannelWeights(%arg0: tensor<4x64xf32>) -> tensor<4x32xf32> { + %bias = "tfl.no_value"() {value} : () -> none + %w = "tfl.pseudo_const"() {value = dense<1> : tensor<32x64xi2>} : () -> tensor<32x64xi2> + %w_scale = "tfl.pseudo_const"() {value = dense<2.500000e-01> : tensor<32x2xf32>} : () -> tensor<32x2xf32> + %w_zp = "tfl.pseudo_const"() {value = dense<-5.000000e-01> : tensor<32x2xf32>} : () -> tensor<32x2xf32> + %act:3 = "tfl.blockwise_quantize"(%arg0) {block_shape = [1, 32], scale_type = f8E8M0FNU, symmetric = true, range_dilation = 1.500000e+00 : f32} : (tensor<4x64xf32>) -> (tensor<4x64xi4>, tensor<4x2xf8E8M0FNU>, none) + %act_dq = "tfl.blockwise_dequantize"(%act#0, %act#1, %act#2) {block_shape = [1, 32], symmetric = true} : (tensor<4x64xi4>, tensor<4x2xf8E8M0FNU>, none) -> tensor<4x64xf32> + %w_dq = "tfl.blockwise_dequantize"(%w, %w_scale, %w_zp) {block_shape = [1, 32], symmetric = false} : (tensor<32x64xi2>, tensor<32x2xf32>, tensor<32x2xf32>) -> tensor<32x64xf32> + %0 = "tfl.fully_connected"(%act_dq, %w_dq, %bias) {asymmetric_quantize_inputs = false, fused_activation_function = "NONE", keep_num_dims = false, weights_format = "DEFAULT"} : (tensor<4x64xf32>, tensor<32x64xf32>, none) -> tensor<4x32xf32> + func.return %0 : tensor<4x32xf32> + + // CHECK: tfl.blockwise_dequantize + // CHECK-NOT: tfl.quant_spec +} + +// ----------------------------------------------------------------------------- +// Negative: asymmetric activations. The collapsed op carries no activation +// zero point. +// ----------------------------------------------------------------------------- + +// CHECK-LABEL: @NoFuseAsymmetricActivations +func.func @NoFuseAsymmetricActivations(%arg0: tensor<4x64xf32>) -> tensor<4x32xf32> { + %bias = "tfl.no_value"() {value} : () -> none + %w = "tfl.pseudo_const"() {value = dense<1> : tensor<32x64xi2>} : () -> tensor<32x64xi2> + %w_scale = "tfl.pseudo_const"() {value = dense<2.500000e-01> : tensor<32x1xf32>} : () -> tensor<32x1xf32> + %w_zp = "tfl.pseudo_const"() {value = dense<-5.000000e-01> : tensor<32x1xf32>} : () -> tensor<32x1xf32> + %act:3 = "tfl.blockwise_quantize"(%arg0) {block_shape = [1, 32], scale_type = f8E8M0FNU, symmetric = false, range_dilation = 1.500000e+00 : f32} : (tensor<4x64xf32>) -> (tensor<4x64xi4>, tensor<4x2xf8E8M0FNU>, tensor<1x1xi4>) + %act_dq = "tfl.blockwise_dequantize"(%act#0, %act#1, %act#2) {block_shape = [1, 32], symmetric = false} : (tensor<4x64xi4>, tensor<4x2xf8E8M0FNU>, tensor<1x1xi4>) -> tensor<4x64xf32> + %w_dq = "tfl.blockwise_dequantize"(%w, %w_scale, %w_zp) {block_shape = [1, 64], symmetric = false} : (tensor<32x64xi2>, tensor<32x1xf32>, tensor<32x1xf32>) -> tensor<32x64xf32> + %0 = "tfl.fully_connected"(%act_dq, %w_dq, %bias) {asymmetric_quantize_inputs = false, fused_activation_function = "NONE", keep_num_dims = false, weights_format = "DEFAULT"} : (tensor<4x64xf32>, tensor<32x64xf32>, none) -> tensor<4x32xf32> + func.return %0 : tensor<4x32xf32> + + // CHECK: tfl.blockwise_quantize + // CHECK-NOT: tfl.quant_spec +} + +// ----------------------------------------------------------------------------- +// Negative: shuffled-weights fully_connected, which has two results. The +// rewrite builds a single-result op, so replacing it would be invalid. +// ----------------------------------------------------------------------------- + +// CHECK-LABEL: @NoFuseShuffledWeights +func.func @NoFuseShuffledWeights(%arg0: tensor<4x64xf32>) -> tensor<4x32xf32> { + %bias = "tfl.no_value"() {value} : () -> none + %w = "tfl.pseudo_const"() {value = dense<1> : tensor<32x64xi2>} : () -> tensor<32x64xi2> + %w_scale = "tfl.pseudo_const"() {value = dense<2.500000e-01> : tensor<32x1xf32>} : () -> tensor<32x1xf32> + %w_zp = "tfl.pseudo_const"() {value = dense<-5.000000e-01> : tensor<32x1xf32>} : () -> tensor<32x1xf32> + %act:3 = "tfl.blockwise_quantize"(%arg0) {block_shape = [1, 32], scale_type = f8E8M0FNU, symmetric = true, range_dilation = 1.500000e+00 : f32} : (tensor<4x64xf32>) -> (tensor<4x64xi4>, tensor<4x2xf8E8M0FNU>, none) + %act_dq = "tfl.blockwise_dequantize"(%act#0, %act#1, %act#2) {block_shape = [1, 32], symmetric = true} : (tensor<4x64xi4>, tensor<4x2xf8E8M0FNU>, none) -> tensor<4x64xf32> + %w_dq = "tfl.blockwise_dequantize"(%w, %w_scale, %w_zp) {block_shape = [1, 64], symmetric = false} : (tensor<32x64xi2>, tensor<32x1xf32>, tensor<32x1xf32>) -> tensor<32x64xf32> + %0:2 = "tfl.fully_connected"(%act_dq, %w_dq, %bias) {asymmetric_quantize_inputs = false, fused_activation_function = "NONE", keep_num_dims = false, weights_format = "SHUFFLED4x16INT8"} : (tensor<4x64xf32>, tensor<32x64xf32>, none) -> (tensor<4x32xf32>, tensor<4x32xf32>) + func.return %0#0 : tensor<4x32xf32> + + // CHECK: tfl.blockwise_dequantize + // CHECK-NOT: tfl.quant_spec +} + +// ----------------------------------------------------------------------------- +// Negative: the dequantize consumes a scale tensor that is not the one the +// matched quantize produced. Every parameter of the fused op is read off the +// quantize, so folding this away would silently change which scale is applied. +// ----------------------------------------------------------------------------- + +// CHECK-LABEL: @NoFuseMismatchedActivationScale +func.func @NoFuseMismatchedActivationScale(%arg0: tensor<4x64xf32>) -> tensor<4x32xf32> { + %bias = "tfl.no_value"() {value} : () -> none + %w = "tfl.pseudo_const"() {value = dense<1> : tensor<32x64xi2>} : () -> tensor<32x64xi2> + %w_scale = "tfl.pseudo_const"() {value = dense<2.500000e-01> : tensor<32x1xf32>} : () -> tensor<32x1xf32> + %w_zp = "tfl.pseudo_const"() {value = dense<-5.000000e-01> : tensor<32x1xf32>} : () -> tensor<32x1xf32> + %other_scale = "tfl.pseudo_const"() {value = dense<1.000000e+00> : tensor<4x2xf8E8M0FNU>} : () -> tensor<4x2xf8E8M0FNU> + %act:3 = "tfl.blockwise_quantize"(%arg0) {block_shape = [1, 32], scale_type = f8E8M0FNU, symmetric = true, range_dilation = 1.500000e+00 : f32} : (tensor<4x64xf32>) -> (tensor<4x64xi4>, tensor<4x2xf8E8M0FNU>, none) + %act_dq = "tfl.blockwise_dequantize"(%act#0, %other_scale, %act#2) {block_shape = [1, 32], symmetric = true} : (tensor<4x64xi4>, tensor<4x2xf8E8M0FNU>, none) -> tensor<4x64xf32> + %w_dq = "tfl.blockwise_dequantize"(%w, %w_scale, %w_zp) {block_shape = [1, 64], symmetric = false} : (tensor<32x64xi2>, tensor<32x1xf32>, tensor<32x1xf32>) -> tensor<32x64xf32> + %0 = "tfl.fully_connected"(%act_dq, %w_dq, %bias) {asymmetric_quantize_inputs = false, fused_activation_function = "NONE", keep_num_dims = false, weights_format = "DEFAULT"} : (tensor<4x64xf32>, tensor<32x64xf32>, none) -> tensor<4x32xf32> + func.return %0 : tensor<4x32xf32> + + // CHECK: tfl.blockwise_dequantize + // CHECK-NOT: tfl.quant_spec +} diff --git a/tensorflow/compiler/mlir/lite/transforms/fuse_a4w2_drq_fully_connected_pass.cc b/tensorflow/compiler/mlir/lite/transforms/fuse_a4w2_drq_fully_connected_pass.cc new file mode 100644 index 00000000000000..1fde90d1b6a53a --- /dev/null +++ b/tensorflow/compiler/mlir/lite/transforms/fuse_a4w2_drq_fully_connected_pass.cc @@ -0,0 +1,288 @@ +/* Copyright 2026 The TensorFlow Authors. All Rights Reserved. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +==============================================================================*/ + +#include +#include +#include +#include + +#include "llvm/ADT/APFloat.h" +#include "llvm/ADT/STLExtras.h" +#include "mlir/Dialect/Func/IR/FuncOps.h" // from @llvm-project +#include "mlir/Dialect/Quant/IR/Quant.h" // from @llvm-project +#include "mlir/Dialect/Quant/IR/QuantTypes.h" // from @llvm-project +#include "mlir/IR/Attributes.h" // from @llvm-project +#include "mlir/IR/Builders.h" // from @llvm-project +#include "mlir/IR/BuiltinAttributes.h" // from @llvm-project +#include "mlir/IR/BuiltinTypes.h" // from @llvm-project +#include "mlir/IR/MLIRContext.h" // from @llvm-project +#include "mlir/IR/Matchers.h" // from @llvm-project +#include "mlir/IR/PatternMatch.h" // from @llvm-project +#include "mlir/IR/Value.h" // from @llvm-project +#include "mlir/Pass/Pass.h" // from @llvm-project +#include "mlir/Support/LLVM.h" // from @llvm-project +#include "mlir/Support/LogicalResult.h" // from @llvm-project +#include "mlir/Support/TypeID.h" // from @llvm-project +#include "mlir/Transforms/GreedyPatternRewriteDriver.h" // from @llvm-project +#include "tensorflow/compiler/mlir/lite/ir/tfl_ops.h" +#include "tensorflow/compiler/mlir/lite/transforms/passes.h" + +namespace mlir { +namespace TFL { +namespace { + +#define GEN_PASS_DEF_FUSEA4W2DRQFULLYCONNECTEDPASS +#include "tensorflow/compiler/mlir/lite/transforms/passes.h.inc" + +// The contract named by `kA4W2SpecName`. Every one of these is baked into the +// runtime kernel rather than carried in the flatbuffer, so the pattern below +// has to verify each one explicitly: a graph that differs in any of them would +// still be serialized as `cint2_fp32_int4_e8m0_drq` and then silently evaluated +// with the wrong numerics. +constexpr char kA4W2SpecName[] = "cint2_fp32_int4_e8m0_drq"; +// Activations are quantized in blocks of 32 along the contracting axis. +constexpr int64_t kActBlockSize = 32; +// Weights use a grid centered between the integers; see +// `pure_observer.CENTERED_ZERO_POINT` on the JAX side and the `+ 0.5f` in +// `EvalA4W2DRQ` on the runtime side. +constexpr double kCenteredZeroPoint = -0.5; + +// Returns true if `value` is a constant whose elements are all exactly +// `expected`. +bool IsSplatFloatConstant(mlir::Value value, double expected) { + if (!value || mlir::isa(value.getType())) return false; + + ElementsAttr attr; + if (!matchPattern(value, m_Constant(&attr))) { + auto const_op = value.getDefiningOp(); + if (!const_op) return false; + attr = const_op.getValue(); + } + + auto fp_attr = mlir::dyn_cast(attr); + if (!fp_attr) return false; + return llvm::all_of(fp_attr.getValues(), [expected](APFloat v) { + return v.convertToDouble() == expected; + }); +} + +// Pattern to match an a4w2 dynamic range fully connected pattern: +// +// %act_q:3 = tfl.blockwise_quantize(%x, block_shape=[..., 32], ...) +// %act_dq = tfl.blockwise_dequantize(%act_q#0, %act_q#1, %act_q#2, ...) +// %w_dq = tfl.blockwise_dequantize(%q_w, %scales, %zp, block_shape=[1, K]) +// %res = tfl.fully_connected(%act_dq, %w_dq, %bias) +// +// and fuse it into a single DRQ fully connected op: +// +// %q_w_per_axis = tfl.pseudo_qconst(...) : tensor> %res = tfl.fully_connected(%x, +// %q_w_per_axis, %bias) +// { tfl.quant_spec = { spec = "cint2_fp32_int4_e8m0_drq", act_dilation +// = ... } +// } +struct FuseA4W2DRQFullyConnectedPattern + : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(TFL::FullyConnectedOp fc, + PatternRewriter& rewriter) const override { + // The rewrite below builds a single-result op, so a shuffled-weights + // fully_connected (which has a second result for the shuffled input) + // cannot be replaced by it. + if (fc.getNumResults() != 1) return failure(); + if (fc.getWeightsFormat() != "DEFAULT") return failure(); + + // 1. Match activations side. + mlir::Value act_val = fc.getInput(); + auto act_deq = act_val.getDefiningOp(); + if (!act_deq) return failure(); + + mlir::Value act_q_val = act_deq.getInput(); + auto act_q = act_q_val.getDefiningOp(); + if (!act_q) return failure(); + + // Consuming `act_q`'s quantized result is not on its own enough to make + // `act_deq` its inverse: the dequantize takes its own scales and zero + // points, while every parameter of the fused op is read off `act_q`. A + // dequantize wired to different quantization parameters describes a + // different computation and must not be folded away here. + // + // The block shape does not need its own check: the verifier ties the + // scales grid to it, so once the scales are required to be `act_q`'s own + // result a differing block shape cannot verify. + if (act_deq.getScales() != act_q.getScale()) return failure(); + if (act_deq.getZeroPoints() != act_q.getZeroPoint()) return failure(); + + if (!act_q.getSymmetric()) return failure(); + + auto act_type = + mlir::dyn_cast(act_q.getOutput().getType()); + if (!act_type || !act_type.getElementType().isInteger(4)) return failure(); + + // The kernel derives the scale as `exp2(ceil(log2(raw_scale)))`, i.e. it + // assumes an e8m0 scale. An f32 scale would quantize to different values. + if (!mlir::isa(act_q.getScaleType())) return failure(); + + // Verify block shape on activations: innermost contracting dimension block + // size must be 32, outer dims must be 1. + ArrayAttr act_block_shape = act_q.getBlockShapeAttr(); + if (!act_block_shape || act_block_shape.empty()) return failure(); + const int64_t act_rank = act_block_shape.size(); + if (mlir::cast(act_block_shape[act_rank - 1]).getInt() != + kActBlockSize) { + return failure(); + } + for (int64_t i = 0; i < act_rank - 1; ++i) { + if (mlir::cast(act_block_shape[i]).getInt() != 1) { + return failure(); + } + } + + mlir::Value real_input = act_q.getInput(); + float act_dilation = act_q.getRangeDilation().convertToFloat(); + + // 2. Match weights side. + mlir::Value filter_val = fc.getFilter(); + auto weight_deq = filter_val.getDefiningOp(); + if (!weight_deq) return failure(); + + // The kernel reconstructs the weights as `(q + 0.5) * scale`, i.e. it + // assumes a grid centered between the integers. Applying that to weights + // quantized on a plain symmetric grid would shift every value by half a + // step, so the centered zero point has to be verified rather than assumed. + if (weight_deq.getSymmetric()) return failure(); + if (!IsSplatFloatConstant(weight_deq.getZeroPoints(), kCenteredZeroPoint)) { + return failure(); + } + + mlir::Value q_weight_val = weight_deq.getInput(); + ElementsAttr q_weight_attr; + if (!matchPattern(q_weight_val, m_Constant(&q_weight_attr))) { + if (auto const_op = q_weight_val.getDefiningOp()) { + q_weight_attr = const_op.getValue(); + } else { + return failure(); + } + } + auto q_weight_type = + mlir::dyn_cast(q_weight_attr.getType()); + if (!q_weight_type || q_weight_type.getRank() != 2) return failure(); + if (!q_weight_type.getElementType().isInteger(2)) return failure(); + + int64_t num_units = q_weight_type.getDimSize(0); + int64_t input_size = q_weight_type.getDimSize(1); + + // Verify weight block shape: [1, input_size] (per-channel across output + // units). + ArrayAttr weight_block_shape = weight_deq.getBlockShapeAttr(); + if (!weight_block_shape || weight_block_shape.size() != 2) return failure(); + if (mlir::cast(weight_block_shape[0]).getInt() != 1 || + mlir::cast(weight_block_shape[1]).getInt() != input_size) { + return failure(); + } + + // Extract per-channel scales. + ElementsAttr scales_attr; + if (!matchPattern(weight_deq.getScales(), m_Constant(&scales_attr))) { + if (auto const_op = + weight_deq.getScales().getDefiningOp()) { + scales_attr = const_op.getValue(); + } else { + return failure(); + } + } + + SmallVector per_channel_scales; + per_channel_scales.reserve(num_units); + if (auto dense_scales = mlir::dyn_cast(scales_attr)) { + for (const auto& fp : dense_scales.getValues()) { + per_channel_scales.push_back(fp.convertToDouble()); + } + } else { + return failure(); + } + if (per_channel_scales.size() != static_cast(num_units)) { + return failure(); + } + + // 3. Create UniformQuantizedPerAxisType for the weights. + MLIRContext* ctx = fc.getContext(); + Type expressed_type = Float32Type::get(ctx); + Type storage_type = IntegerType::get(ctx, 2); + int64_t qmin = -2; + int64_t qmax = 1; + SmallVector zero_points(num_units, 0); + int32_t quantized_dimension = 0; + + auto quant_type = quant::UniformQuantizedPerAxisType::get( + /*flags=*/quant::QuantizationFlags::Signed, storage_type, + expressed_type, per_channel_scales, zero_points, quantized_dimension, + qmin, qmax); + + RankedTensorType new_filter_type = + RankedTensorType::get(q_weight_type.getShape(), quant_type); + + mlir::Value new_filter = rewriter.create( + fc.getLoc(), TypeAttr::get(new_filter_type), q_weight_attr); + + // 4. Create tfl.quant_spec attribute dictionary. + SmallVector spec_entries; + spec_entries.push_back( + rewriter.getNamedAttr("spec", rewriter.getStringAttr(kA4W2SpecName))); + spec_entries.push_back(rewriter.getNamedAttr( + "act_dilation", rewriter.getF32FloatAttr(act_dilation))); + DictionaryAttr quant_spec_dict = rewriter.getDictionaryAttr(spec_entries); + + // 5. Replace the FullyConnectedOp. + auto new_fc = rewriter.create( + fc.getLoc(), fc.getType(0), real_input, new_filter, fc.getBias(), + fc.getFusedActivationFunctionAttr(), fc.getWeightsFormatAttr(), + fc.getKeepNumDimsAttr(), fc.getAsymmetricQuantizeInputsAttr()); + new_fc->setAttr(TFL::kQuantSpecAttrName, quant_spec_dict); + + rewriter.replaceOp(fc, new_fc.getOutput()); + return success(); + } +}; + +struct FuseA4W2DRQFullyConnectedPass + : public impl::FuseA4W2DRQFullyConnectedPassBase< + FuseA4W2DRQFullyConnectedPass> { + public: + MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(FuseA4W2DRQFullyConnectedPass) + + void runOnOperation() override { + func::FuncOp func = getOperation(); + MLIRContext* context = func.getContext(); + + RewritePatternSet patterns(context); + patterns.add(context); + + if (failed(applyPatternsGreedily(func, std::move(patterns)))) { + signalPassFailure(); + } + } +}; + +} // namespace + +std::unique_ptr> +CreateFuseA4W2DRQFullyConnectedPass() { + return std::make_unique(); +} + +} // namespace TFL +} // namespace mlir diff --git a/tensorflow/compiler/mlir/lite/transforms/passes.h b/tensorflow/compiler/mlir/lite/transforms/passes.h index 4c44edc6e4c844..fc2a5a5b98941d 100644 --- a/tensorflow/compiler/mlir/lite/transforms/passes.h +++ b/tensorflow/compiler/mlir/lite/transforms/passes.h @@ -126,6 +126,9 @@ std::unique_ptr> CreateDefaultQuantizePass(); std::unique_ptr> CreateFoldStablehloConstantTransformsPass(); +std::unique_ptr> +CreateFuseA4W2DRQFullyConnectedPass(); + std::unique_ptr> CreateLowerQuantAnnotationsPass(); // Creates an instance of the TFLite PropagateQParams pass which propagates diff --git a/tensorflow/compiler/mlir/lite/transforms/passes.td b/tensorflow/compiler/mlir/lite/transforms/passes.td index a82604770e9ad5..08966e29514b2b 100644 --- a/tensorflow/compiler/mlir/lite/transforms/passes.td +++ b/tensorflow/compiler/mlir/lite/transforms/passes.td @@ -369,6 +369,22 @@ def FoldStablehloConstantTransformsPass : Pass<"tfl-fold-stablehlo-constant-tran ]; } +def FuseA4W2DRQFullyConnectedPass : Pass<"tfl-fuse-a4w2-drq-fully-connected", "mlir::func::FuncOp"> { + let summary = "Collapses a blockwise a4w2 dynamic range quantized pattern into a single fully_connected."; + let description = [{ + Matches a `blockwise_quantize` / `blockwise_dequantize` pair feeding the + activations of a `tfl.fully_connected` whose filter is a + `blockwise_dequantize` of an i2 constant, and rewrites it into one + `tfl.fully_connected` with per-axis i2 weights and a `tfl.quant_spec` + naming the `cint2_fp32_int4_e8m0_drq` contract. + }]; + let constructor = "CreateFuseA4W2DRQFullyConnectedPass()"; + let dependentDialects = [ + "TFL::TensorFlowLiteDialect", + "mlir::quant::QuantDialect" + ]; +} + def PropagateQParamsPass : Pass<"tfl-propagate-qparams", "mlir::ModuleOp"> { let summary = "Propagates Quantization Parameters (scale and zero point) information through the graph."; let description = [{ diff --git a/tensorflow/core/common_runtime/gpu/BUILD b/tensorflow/core/common_runtime/gpu/BUILD index fbfb30d068fe55..82f823ea0fa274 100644 --- a/tensorflow/core/common_runtime/gpu/BUILD +++ b/tensorflow/core/common_runtime/gpu/BUILD @@ -380,9 +380,10 @@ tf_cuda_cc_test( "//tensorflow/core/common_runtime:direct_session_internal", "//tensorflow/core/kernels:ops_util", "@com_google_absl//absl/synchronization", - "@xla//xla/stream_executor/gpu:gpu_cudamallocasync_allocator", "@xla//xla/tsl/framework:device_id", - ], + ] + if_cuda([ + "@xla//xla/stream_executor/gpu:gpu_cudamallocasync_allocator", + ]), ) tf_cuda_cc_test( @@ -443,8 +444,9 @@ tf_cuda_cc_test( "//tensorflow/core/common_runtime:direct_session_internal", "//tensorflow/core/kernels:ops_util", "@com_google_absl//absl/synchronization", + ] + if_cuda([ "@xla//xla/stream_executor/gpu:gpu_cudamallocasync_allocator", - ], + ]), ) tf_cuda_cc_test( diff --git a/tensorflow/core/grappler/utils/BUILD b/tensorflow/core/grappler/utils/BUILD index 165c7dce983bb9..57fe5959fae07d 100644 --- a/tensorflow/core/grappler/utils/BUILD +++ b/tensorflow/core/grappler/utils/BUILD @@ -212,6 +212,7 @@ cc_library( "//tensorflow/core:lib_internal", "//tensorflow/core:protos_all_cc", "//tensorflow/core/common_runtime:core_cpu_base_no_ops", + "//tensorflow/core/common_runtime:function_body", "//tensorflow/core/grappler:grappler_item", "//tensorflow/core/grappler:op_types", "//tensorflow/core/grappler:utils", diff --git a/tensorflow/lite/core/kernels/register.cc b/tensorflow/lite/core/kernels/register.cc index 460146c179d11d..60395705b7f69a 100644 --- a/tensorflow/lite/core/kernels/register.cc +++ b/tensorflow/lite/core/kernels/register.cc @@ -83,7 +83,7 @@ BuiltinOpResolver::BuiltinOpResolver() { Register_EMBEDDING_LOOKUP_SPARSE()); AddBuiltin(BuiltinOperator_FULLY_CONNECTED, Register_FULLY_CONNECTED(), /* min_version = */ 1, - /* max_version = */ 14); + /* max_version = */ 15); AddBuiltin(BuiltinOperator_LSH_PROJECTION, Register_LSH_PROJECTION()); AddBuiltin(BuiltinOperator_HASHTABLE_LOOKUP, Register_HASHTABLE_LOOKUP()); AddBuiltin(BuiltinOperator_SOFTMAX, Register_SOFTMAX(), diff --git a/tensorflow/lite/kernels/BUILD b/tensorflow/lite/kernels/BUILD index 855fdae1fd189b..775ec49bed7f36 100644 --- a/tensorflow/lite/kernels/BUILD +++ b/tensorflow/lite/kernels/BUILD @@ -2269,6 +2269,7 @@ cc_test( "//tensorflow/lite/schema:schema_fbs", "@com_google_absl//absl/log:absl_check", "@com_google_googletest//:gtest", + "@flatbuffers", ], ) diff --git a/tensorflow/lite/kernels/fully_connected.cc b/tensorflow/lite/kernels/fully_connected.cc index f60213231fde06..dda28a3bc449e8 100644 --- a/tensorflow/lite/kernels/fully_connected.cc +++ b/tensorflow/lite/kernels/fully_connected.cc @@ -16,6 +16,7 @@ limitations under the License. #include "tensorflow/lite/kernels/internal/optimized/integer_ops/fully_connected.h" #include +#include #include #include #include @@ -23,10 +24,15 @@ limitations under the License. #include #include +#include "Eigen/Core" // from @eigen_archive +#include "flatbuffers/flexbuffers.h" // from @flatbuffers +#include "ruy/profiler/instrumentation.h" // from @ruy #include "tensorflow/lite/core/c/builtin_op_data.h" #include "tensorflow/lite/core/c/c_api_types.h" #include "tensorflow/lite/core/c/common.h" #include "tensorflow/lite/kernels/cpu_backend_context.h" +#include "tensorflow/lite/kernels/cpu_backend_threadpool.h" +#include "tensorflow/lite/kernels/internal/common.h" #include "tensorflow/lite/kernels/internal/optimized/fully_connected_4bit.h" #include "tensorflow/lite/kernels/internal/optimized/optimized_ops.h" #include "tensorflow/lite/kernels/internal/optimized/sparse_ops/fully_connected.h" @@ -36,10 +42,12 @@ limitations under the License. #include "tensorflow/lite/kernels/internal/reference/integer_ops/fully_connected.h" #include "tensorflow/lite/kernels/internal/reference/reference_ops.h" #include "tensorflow/lite/kernels/internal/reference/sparse_ops/fully_connected.h" +#include "tensorflow/lite/kernels/internal/runtime_shape.h" #include "tensorflow/lite/kernels/internal/tensor_ctypes.h" #include "tensorflow/lite/kernels/internal/tensor_utils.h" #include "tensorflow/lite/kernels/internal/types.h" #include "tensorflow/lite/kernels/kernel_util.h" +#include "tensorflow/lite/logger.h" #include "tensorflow/lite/minimal_logging.h" #include "tensorflow/lite/util.h" #ifdef TFLITE_HAVE_CPUINFO @@ -723,6 +731,80 @@ TfLiteStatus PrepareImpl(TfLiteContext* context, TfLiteNode* node, filter->dims->data[1]); } +// Block size of the activation quantization in the `cint2_fp32_int4_e8m0_drq` +// contract. The contract fixes it rather than carrying it in the flatbuffer, +// so a filter whose contracting dimension is not a multiple of it does not +// describe a valid a4w2 op. +constexpr int kA4W2BlockSize = 32; +// Inclusive bounds of the 4 bit activation storage. +constexpr float kA4W2ActQMin = -8.0f; +constexpr float kA4W2ActQMax = 7.0f; + +// The quantization contracts this kernel implements, as named by the opaque +// `FullyConnectedOptions.quant_spec` payload. +enum class QuantSpecKind { + // No `quant_spec`; the op uses standard FullyConnected quantization. + kNone, + // Blockwise 4 bit dynamic range activations against 2 bit centered weights. + kA4W2Drq, +}; + +// Parses `params->quant_spec`. +// +// schema.fbs requires that a runtime which does not understand the contract +// named in the payload *reject* the op rather than fall back to standard +// quantization semantics, because the standard semantics would produce +// plausible but wrong numbers. So an unrecognized or malformed spec is an +// error here, not a signal to ignore the field. +TfLiteStatus ParseQuantSpec(TfLiteContext* context, + const TfLiteFullyConnectedParams* params, + QuantSpecKind* kind, float* act_dilation) { + *kind = QuantSpecKind::kNone; + *act_dilation = 0.0f; + if (params->quant_spec == nullptr || params->quant_spec_size <= 0) { + return kTfLiteOk; + } + + const uint8_t* buffer = reinterpret_cast(params->quant_spec); + const size_t buffer_size = static_cast(params->quant_spec_size); + // The payload comes straight out of the model file, so it is untrusted. + TF_LITE_ENSURE_MSG(context, + flexbuffers::VerifyBuffer(buffer, buffer_size, + /*reuse_tracker=*/nullptr), + "FullyConnected quant_spec is not a valid flexbuffer."); + + const flexbuffers::Map map = + flexbuffers::GetRoot(buffer, buffer_size).AsMap(); + const flexbuffers::String spec = map["spec"].AsString(); + if (spec.c_str() != nullptr && + std::strcmp(spec.c_str(), "cint2_fp32_int4_e8m0_drq") == 0) { + // `act_dilation` widens the denominator of the activation scale, so a + // missing key, a NaN, or a value that drives `kA4W2ActQMax + act_dilation` + // to zero or below would yield an infinite or NaN scale rather than an + // error. + const flexbuffers::Reference dilation = map["act_dilation"]; + TF_LITE_ENSURE_MSG( + context, dilation.IsNumeric(), + "cint2_fp32_int4_e8m0_drq requires a numeric 'act_dilation' entry."); + const float dilation_value = dilation.AsFloat(); + TF_LITE_ENSURE_MSG(context, + std::isfinite(dilation_value) && dilation_value >= 0.0f, + "cint2_fp32_int4_e8m0_drq requires a finite, " + "non-negative 'act_dilation'."); + *kind = QuantSpecKind::kA4W2Drq; + *act_dilation = dilation_value; + return kTfLiteOk; + } + + TF_LITE_KERNEL_LOG( + context, + "FullyConnected names quantization spec '%s', which this runtime does " + "not implement. Refusing to fall back to standard quantization " + "semantics.", + spec.c_str() != nullptr ? spec.c_str() : ""); + return kTfLiteError; +} + template TfLiteStatus Prepare(TfLiteContext* context, TfLiteNode* node) { OpData* data = reinterpret_cast(node->user_data); @@ -743,6 +825,35 @@ TfLiteStatus Prepare(TfLiteContext* context, TfLiteNode* node) { (filter->type == kTfLiteInt4) || (filter->type == kTfLiteInt2)); const bool is_hybrid = is_quantized && (input->type == kTfLiteFloat32); + // Validate the quantization contract here rather than in the eval paths: + // only `EvalHybridDense` knows how to honor one, so a `quant_spec` on any + // other path would otherwise be silently ignored. + QuantSpecKind quant_spec_kind = QuantSpecKind::kNone; + float act_dilation = 0.0f; + TF_LITE_ENSURE_OK(context, ParseQuantSpec(context, params, &quant_spec_kind, + &act_dilation)); + if (quant_spec_kind == QuantSpecKind::kA4W2Drq) { + TF_LITE_ENSURE_MSG(context, is_hybrid, + "cint2_fp32_int4_e8m0_drq requires float input with a " + "quantized filter."); + // `EvalHybrid` routes a sparse filter to `EvalHybridSparse`, which never + // reaches `EvalHybridDense` and would therefore run standard sparse hybrid + // math while ignoring the spec. + TF_LITE_ENSURE_MSG( + context, filter->sparsity == nullptr, + "cint2_fp32_int4_e8m0_drq does not support sparse filters."); + // `is_hybrid` also admits uint8, int8 and int4 filters. `EvalA4W2DRQ` + // rejects those too, but only once the graph is already being invoked. + TF_LITE_ENSURE_MSG(context, filter->type == kTfLiteInt2, + "cint2_fp32_int4_e8m0_drq requires a 2 bit filter."); + TF_LITE_ENSURE_MSG(context, filter->dims->size == 2, + "cint2_fp32_int4_e8m0_drq requires a rank 2 filter."); + TF_LITE_ENSURE_MSG(context, filter->dims->data[1] % kA4W2BlockSize == 0, + "cint2_fp32_int4_e8m0_drq requires the filter's input " + "size to be a multiple " + "of the 32 element activation block size."); + } + // Pie and hybrid path supports all kinds of fused activations, otherwise only // clipping activations are supported. if (!is_hybrid) { @@ -760,6 +871,131 @@ TfLiteStatus Prepare(TfLiteContext* context, TfLiteNode* node) { return PrepareImpl(context, node, kernel_type); } +TfLiteStatus EvalA4W2DRQ(TfLiteContext* context, + TfLiteFullyConnectedParams* params, + const TfLiteTensor* input, const TfLiteTensor* filter, + const TfLiteTensor* bias, TfLiteTensor* output, + float act_dilation) { + // This kernel is selected by an opaque `quant_spec` string rather than by the + // tensor types, so nothing upstream has checked that the operands match what + // the spec describes. Everything the loops below rely on is verified here; + // in particular the 2 bit unpack would read four times past the end of the + // filter buffer if the filter were int8. + TF_LITE_ENSURE_TYPES_EQ(context, input->type, kTfLiteFloat32); + TF_LITE_ENSURE_TYPES_EQ(context, output->type, kTfLiteFloat32); + TF_LITE_ENSURE_TYPES_EQ(context, filter->type, kTfLiteInt2); + if (bias != nullptr) { + TF_LITE_ENSURE_TYPES_EQ(context, bias->type, kTfLiteFloat32); + } + TF_LITE_ENSURE_EQ(context, filter->dims->size, 2); + + const int input_size = filter->dims->data[1]; + const int num_units = filter->dims->data[0]; + TF_LITE_ENSURE(context, input_size > 0); + TF_LITE_ENSURE(context, num_units > 0); + TF_LITE_ENSURE_MSG(context, input_size % kA4W2BlockSize == 0, + "cint2_fp32_int4_e8m0_drq requires the filter's input " + "size to be a multiple of " + "the 32 element activation block size."); + + const int total_input_size = input->bytes / sizeof(float); + TF_LITE_ENSURE_MSG(context, total_input_size % input_size == 0, + "cint2_fp32_int4_e8m0_drq requires the input size to be a " + "whole number of rows."); + const int batch_size = total_input_size / input_size; + if (bias != nullptr) { + TF_LITE_ENSURE_EQ(context, NumElements(bias), num_units); + } + TF_LITE_ENSURE_EQ(context, NumElements(output), batch_size * num_units); + + if (bias) { + tensor_utils::VectorBatchVectorAssign(GetTensorData(bias), num_units, + batch_size, + GetTensorData(output)); + } else { + std::fill_n(GetTensorData(output), batch_size * num_units, 0.0f); + } + + const size_t num_filter_elements = + static_cast(num_units) * input_size; + auto unpacked_filter = std::make_unique(num_filter_elements); + tflite::tensor_utils::UnpackPackedIntToInt8( + GetTensorData(filter), num_filter_elements, + /*bit_width=*/2, unpacked_filter.get()); + + // The contract is per-channel on the output dimension. Checked directly + // rather than through `VerifyPerChannelQuantization`, which logs an error of + // its own when the tensor is not affine quantized and which rejects a + // one-element scale array, i.e. a legitimate single output channel filter. + TF_LITE_ENSURE_EQ(context, filter->quantization.type, + kTfLiteAffineQuantization); + const auto* affine_quantization = + reinterpret_cast(filter->quantization.params); + TF_LITE_ENSURE(context, affine_quantization != nullptr); + TF_LITE_ENSURE(context, affine_quantization->scale != nullptr); + TF_LITE_ENSURE_MSG( + context, affine_quantization->scale->size == num_units, + "cint2_fp32_int4_e8m0_drq requires one weight scale per output channel."); + const float* per_channel_scale_ptr = affine_quantization->scale->data; + + const float* input_ptr = GetTensorData(input); + float* output_ptr = GetTensorData(output); + + const int num_blocks = input_size / kA4W2BlockSize; + + std::vector q_act(input_size); + std::vector act_scales(num_blocks); + + for (int b = 0; b < batch_size; ++b) { + const float* in_row = input_ptr + b * input_size; + float* out_row = output_ptr + b * num_units; + + for (int blk = 0; blk < num_blocks; ++blk) { + const float* block_ptr = in_row + blk * kA4W2BlockSize; + float max_abs = 0.0f; + for (int k = 0; k < kA4W2BlockSize; ++k) { + max_abs = std::max(max_abs, std::abs(block_ptr[k])); + } + float raw_scale = max_abs / (kA4W2ActQMax + act_dilation); + if (raw_scale <= 0.0f) raw_scale = 1.0f; + float scale = std::exp2(std::ceil(std::log2(raw_scale))); + act_scales[blk] = scale; + + for (int k = 0; k < kA4W2BlockSize; ++k) { + float q = std::nearbyint(block_ptr[k] / scale); + q = std::clamp(q, kA4W2ActQMin, kA4W2ActQMax); + q_act[blk * kA4W2BlockSize + k] = static_cast(q); + } + } + + for (int out_c = 0; out_c < num_units; ++out_c) { + const int8_t* w_row = unpacked_filter.get() + out_c * input_size; + const float s_w = per_channel_scale_ptr[out_c]; + + double acc = 0.0; + for (int blk = 0; blk < num_blocks; ++blk) { + const int8_t* a_blk = q_act.data() + blk * kA4W2BlockSize; + const int8_t* w_blk = w_row + blk * kA4W2BlockSize; + const float s_a = act_scales[blk]; + + float blk_sum = 0.0f; + for (int k = 0; k < kA4W2BlockSize; ++k) { + blk_sum += static_cast(a_blk[k]) * + (static_cast(w_blk[k]) + 0.5f); + } + acc += static_cast(s_a) * static_cast(s_w) * + static_cast(blk_sum); + } + out_row[out_c] += static_cast(acc); + } + } + + tensor_utils::ApplyActivationToVector(output_ptr, batch_size * num_units, + params->activation, output_ptr); + + return kTfLiteOk; +} + TfLiteStatus EvalHybridDense( TfLiteContext* context, TfLiteNode* node, TfLiteFullyConnectedParams* params, OpData* data, const TfLiteTensor* input, @@ -767,6 +1003,15 @@ TfLiteStatus EvalHybridDense( TfLiteTensor* input_quantized, TfLiteTensor* scaling_factors, TfLiteTensor* accum_scratch, TfLiteTensor* row_sums, TfLiteTensor* input_offsets, TfLiteTensor* output) { + QuantSpecKind quant_spec_kind = QuantSpecKind::kNone; + float act_dilation = 0.0f; + TF_LITE_ENSURE_OK(context, ParseQuantSpec(context, params, &quant_spec_kind, + &act_dilation)); + if (quant_spec_kind == QuantSpecKind::kA4W2Drq) { + return EvalA4W2DRQ(context, params, input, filter, bias, output, + act_dilation); + } + int total_input_size = 1; for (int i = 0; i < input->dims->size; i++) { total_input_size *= input->dims->data[i]; diff --git a/tensorflow/lite/kernels/fully_connected_test.cc b/tensorflow/lite/kernels/fully_connected_test.cc index 668ae435749d93..42250d4125d608 100644 --- a/tensorflow/lite/kernels/fully_connected_test.cc +++ b/tensorflow/lite/kernels/fully_connected_test.cc @@ -31,6 +31,7 @@ limitations under the License. #include #include #include "absl/log/absl_check.h" +#include "flatbuffers/flexbuffers.h" // from @flatbuffers #include "tensorflow/lite/core/interpreter.h" #include "tensorflow/lite/kernels/test_util.h" #include "tensorflow/lite/schema/schema_generated.h" @@ -2946,6 +2947,257 @@ TEST(FullyConnectedInt16FilterInt16IndexingTest, RejectsShapeProductOverflow) { kTfLiteError); } +// Builds a FullyConnected carrying an opaque `quant_spec`. +// +// The `cint2_fp32_int4_e8m0_drq` contract is selected by that string alone: +// nothing in the tensor types distinguishes it from a standard hybrid +// FullyConnected, so these tests drive the op through the flatbuffer rather +// than through the higher level helpers above. +class QuantSpecFullyConnectedOpModel : public SingleOpModel { + public: + QuantSpecFullyConnectedOpModel(int units, int batches, + const TensorData& input, + const TensorData& weights, + const std::string& spec_name, + float act_dilation, + bool corrupt_quant_spec = false) + : batches_(batches), units_(units) { + input_ = AddInput(input); + weights_ = AddInput(weights); + bias_ = AddInput({TensorType_FLOAT32, {units_}}); + output_ = AddOutput({TensorType_FLOAT32}); + + flexbuffers::Builder fbb; + const size_t map_start = fbb.StartMap(); + fbb.String("spec", spec_name); + fbb.Double("act_dilation", act_dilation); + fbb.EndMap(map_start); + fbb.Finish(); + std::vector payload = fbb.GetBuffer(); + if (corrupt_quant_spec) { + // Truncating the payload leaves the trailing byte-width/type bytes that + // the flexbuffer root is located from pointing outside the buffer. + payload.resize(payload.size() / 2); + } + + const auto quant_spec = builder_.CreateVector(payload); + const auto options = + CreateFullyConnectedOptions(builder_, ActivationFunctionType_NONE, + FullyConnectedOptionsWeightsFormat_DEFAULT, + /*keep_num_dims=*/false, + /*asymmetric_quantize_inputs=*/false, + TensorType_FLOAT32, quant_spec) + .Union(); + SetBuiltinOp(BuiltinOperator_FULLY_CONNECTED, + BuiltinOptions_FullyConnectedOptions, options); + resolver_ = std::make_unique( + BuiltinOperator_FULLY_CONNECTED, + ops::builtin::Register_FULLY_CONNECTED_REF()); + // The reference kernel is the point of these tests, so no delegate. + BuildInterpreter({GetShape(input_), GetShape(weights_), GetShape(bias_)}, + /*num_threads=*/1, /*allow_fp32_relax_to_fp16=*/false, + /*apply_delegate=*/false, /*allocate_and_delegate=*/false); + } + + using SingleOpModel::AllocateTensors; + + void SetBias(const std::vector& f) { PopulateTensor(bias_, f); } + void SetInput(const std::vector& f) { PopulateTensor(input_, f); } + + // Writes raw 2 bit codes. The centered grid this kernel implements is not + // reachable through the symmetric quantize-and-populate helpers. + void SetRawWeights(const std::vector& codes) { + PopulateTensor2bit(weights_, /*offset=*/0, codes.data(), + codes.data() + codes.size()); + } + void SetRawInt8Weights(const std::vector& codes) { + PopulateTensor(weights_, codes); + } + + std::vector GetOutput() { return ExtractVector(output_); } + + protected: + int input_; + int weights_; + int bias_; + int output_; + int batches_; + int units_; +}; + +constexpr int kA4W2TestBlockSize = 32; + +TEST(A4W2DrqFullyConnectedTest, MatchesHandComputedResult) { + // One block of 32, two output channels. + const std::vector per_channel_scales = {0.5f, 0.25f}; + QuantSpecFullyConnectedOpModel m( + /*units=*/2, /*batches=*/1, + /*input=*/{TensorType_FLOAT32, {1, kA4W2TestBlockSize}}, + /*weights=*/ + {TensorType_INT2, + {2, kA4W2TestBlockSize}, + /*min=*/0.0f, + /*max=*/0.0f, + /*scale=*/0.0f, + /*zero_point=*/0, + /*per_channel_quantization=*/true, + per_channel_scales, + /*per_channel_quantization_offsets=*/{0, 0}, + /*channel_index=*/0}, + /*spec_name=*/"cint2_fp32_int4_e8m0_drq", /*act_dilation=*/1.5f); + ASSERT_EQ(m.AllocateTensors(), kTfLiteOk); + + // max_abs is 8, so raw_scale = 8 / (7 + 1.5) = 0.941..., which rounds up to + // the power of two 2^0 = 1. Every input is then an exact 4 bit code. + std::vector input(kA4W2TestBlockSize, 1.0f); + input[0] = -8.0f; + m.SetInput(input); + + // Channel 0 is all code 1, channel 1 is all code -2. The kernel reconstructs + // them as (code + 0.5) * per_channel_scale. + std::vector weights(kA4W2TestBlockSize, 1); + weights.insert(weights.end(), kA4W2TestBlockSize, -2); + m.SetRawWeights(weights); + + m.SetBias({1.0f, 2.0f}); + ASSERT_EQ(m.Invoke(), kTfLiteOk); + + // Channel 0: sum(q_a * 1.5) = (-8 * 1.5) + (31 * 1.5) = 34.5 + // out = 1.0 * 0.5 * 34.5 + 1.0 = 18.25 + // Channel 1: sum(q_a * -1.5) = (-8 * -1.5) + (31 * -1.5) = -34.5 + // out = 1.0 * 0.25 * -34.5 + 2.0 = -6.625 + EXPECT_THAT(m.GetOutput(), ElementsAre(18.25f, -6.625f)); +} + +TEST(A4W2DrqFullyConnectedTest, RejectsUnimplementedSpec) { + // schema.fbs requires a runtime that does not implement the named contract + // to reject the op rather than silently fall back to standard hybrid + // quantization, which would produce plausible but wrong numbers. + QuantSpecFullyConnectedOpModel m( + /*units=*/2, /*batches=*/1, + /*input=*/{TensorType_FLOAT32, {1, kA4W2TestBlockSize}}, + /*weights=*/ + {TensorType_INT2, + {2, kA4W2TestBlockSize}, + 0.0f, + 0.0f, + 0.0f, + 0, + /*per_channel_quantization=*/true, + {0.5f, 0.25f}, + {0, 0}, + 0}, + /*spec_name=*/"a4w2_drq_v99", /*act_dilation=*/1.5f); + EXPECT_EQ(m.AllocateTensors(), kTfLiteError); +} + +TEST(A4W2DrqFullyConnectedTest, RejectsMalformedQuantSpec) { + QuantSpecFullyConnectedOpModel m( + /*units=*/2, /*batches=*/1, + /*input=*/{TensorType_FLOAT32, {1, kA4W2TestBlockSize}}, + /*weights=*/ + {TensorType_INT2, + {2, kA4W2TestBlockSize}, + 0.0f, + 0.0f, + 0.0f, + 0, + /*per_channel_quantization=*/true, + {0.5f, 0.25f}, + {0, 0}, + 0}, + /*spec_name=*/"cint2_fp32_int4_e8m0_drq", /*act_dilation=*/1.5f, + /*corrupt_quant_spec=*/true); + EXPECT_EQ(m.AllocateTensors(), kTfLiteError); +} + +TEST(A4W2DrqFullyConnectedTest, RejectsInt8Filter) { + // The kernel unpacks the filter as 2 bit values, so an int8 filter would + // make it read four times past the end of the buffer. `is_hybrid` alone does + // not exclude it, so `Prepare` has to. + QuantSpecFullyConnectedOpModel m( + /*units=*/2, /*batches=*/1, + /*input=*/{TensorType_FLOAT32, {1, kA4W2TestBlockSize}}, + /*weights=*/ + {TensorType_INT8, + {2, kA4W2TestBlockSize}, + 0.0f, + 0.0f, + 0.0f, + 0, + /*per_channel_quantization=*/true, + {0.5f, 0.25f}, + {0, 0}, + 0}, + /*spec_name=*/"cint2_fp32_int4_e8m0_drq", /*act_dilation=*/1.5f); + EXPECT_EQ(m.AllocateTensors(), kTfLiteError); +} + +TEST(A4W2DrqFullyConnectedTest, RejectsInputSizeThatIsNotAWholeNumberOfBlocks) { + // 48 is not a multiple of the 32 element block size; without this check the + // kernel would silently drop the trailing 16 elements of every row. + constexpr int kInputSize = 48; + QuantSpecFullyConnectedOpModel m( + /*units=*/2, /*batches=*/1, + /*input=*/{TensorType_FLOAT32, {1, kInputSize}}, + /*weights=*/ + {TensorType_INT2, + {2, kInputSize}, + 0.0f, + 0.0f, + 0.0f, + 0, + /*per_channel_quantization=*/true, + {0.5f, 0.25f}, + {0, 0}, + 0}, + /*spec_name=*/"cint2_fp32_int4_e8m0_drq", /*act_dilation=*/1.5f); + EXPECT_EQ(m.AllocateTensors(), kTfLiteError); +} + +// `act_dilation` comes out of the model buffer and lands in the denominator of +// the activation scale, so values that make that denominator non-positive, or +// that are not finite, have to be rejected rather than producing inf/NaN +// scales and silently poisoning every output. +TEST(A4W2DrqFullyConnectedTest, RejectsActDilationThatCollapsesTheScale) { + QuantSpecFullyConnectedOpModel m( + /*units=*/2, /*batches=*/1, + /*input=*/{TensorType_FLOAT32, {1, kA4W2TestBlockSize}}, + /*weights=*/ + {TensorType_INT2, + {2, kA4W2TestBlockSize}, + 0.0f, + 0.0f, + 0.0f, + 0, + /*per_channel_quantization=*/true, + {0.5f, 0.25f}, + {0, 0}, + 0}, + /*spec_name=*/"cint2_fp32_int4_e8m0_drq", /*act_dilation=*/-7.0f); + EXPECT_EQ(m.AllocateTensors(), kTfLiteError); +} + +TEST(A4W2DrqFullyConnectedTest, RejectsNonFiniteActDilation) { + QuantSpecFullyConnectedOpModel m( + /*units=*/2, /*batches=*/1, + /*input=*/{TensorType_FLOAT32, {1, kA4W2TestBlockSize}}, + /*weights=*/ + {TensorType_INT2, + {2, kA4W2TestBlockSize}, + 0.0f, + 0.0f, + 0.0f, + 0, + /*per_channel_quantization=*/true, + {0.5f, 0.25f}, + {0, 0}, + 0}, + /*spec_name=*/"cint2_fp32_int4_e8m0_drq", + /*act_dilation=*/std::numeric_limits::quiet_NaN()); + EXPECT_EQ(m.AllocateTensors(), kTfLiteError); +} + INSTANTIATE_TEST_SUITE_P( SparseQuantizedFullyConnectedOpTest, SparseQuantizedFullyConnectedOpTest, ::testing::ValuesIn(SingleOpTest::GetKernelTags(*kKernelMap))); diff --git a/tensorflow/lite/kernels/register_ref.cc b/tensorflow/lite/kernels/register_ref.cc index d8741169e20c0a..546600b631cff9 100644 --- a/tensorflow/lite/kernels/register_ref.cc +++ b/tensorflow/lite/kernels/register_ref.cc @@ -280,7 +280,7 @@ BuiltinRefOpResolver::BuiltinRefOpResolver() { Register_EMBEDDING_LOOKUP_SPARSE()); AddBuiltin(BuiltinOperator_FULLY_CONNECTED, Register_FULLY_CONNECTED_REF(), /* min_version */ 1, - /* max_version */ 14); + /* max_version */ 15); AddBuiltin(BuiltinOperator_LSH_PROJECTION, Register_LSH_PROJECTION()); AddBuiltin(BuiltinOperator_HASHTABLE_LOOKUP, Register_HASHTABLE_LOOKUP()); AddBuiltin(BuiltinOperator_SOFTMAX, Register_SOFTMAX_REF(), diff --git a/third_party/xla/.kokoro/macos/build.sh b/third_party/xla/.kokoro/macos/build.sh index c84e1814b2eac9..008638a2abf7e1 100644 --- a/third_party/xla/.kokoro/macos/build.sh +++ b/third_party/xla/.kokoro/macos/build.sh @@ -25,4 +25,4 @@ set -euox pipefail -o history # TODO(ddunleavy) figure out how to best move this into build.py cd "${KOKORO_ARTIFACTS_DIR}/github/xla" -"$KOKORO_ARTIFACTS_DIR"/github/xla/build_tools/ci/build.py --build=XLA_MACOS_X86_CPU_KOKORO +python3 "$KOKORO_ARTIFACTS_DIR"/github/xla/build_tools/ci/build.py --build=XLA_MACOS_X86_CPU_KOKORO diff --git a/third_party/xla/build_tools/ci/build.py b/third_party/xla/build_tools/ci/build.py index 8c4583fbcd9139..1193fc43f35eb3 100755 --- a/third_party/xla/build_tools/ci/build.py +++ b/third_party/xla/build_tools/ci/build.py @@ -32,8 +32,18 @@ import sys from typing import Any, ClassVar, Dict, List, Optional, Tuple +# Ensure the XLA repository root is on sys.path when invoked directly as a +# script in OSS (e.g. `python3 build_tools/ci/build.py` without PYTHONPATH). +_XLA_SRC_ROOT = os.path.dirname( + os.path.dirname(os.path.dirname(os.path.abspath(__file__))) +) +if _XLA_SRC_ROOT not in sys.path: + sys.path.insert(0, _XLA_SRC_ROOT) + +# pylint: disable=g-import-not-at-top from build_tools.ci import bazel_diff from build_tools.ci import change_detector +# pylint: enable=g-import-not-at-top # TODO(ddunleavy): move this to the bazelrc _DEFAULT_BAZEL_OPTIONS = dict( diff --git a/third_party/xla/third_party/stablehlo/temporary.patch b/third_party/xla/third_party/stablehlo/temporary.patch index 48889076fb4a16..73906a9c97345e 100644 --- a/third_party/xla/third_party/stablehlo/temporary.patch +++ b/third_party/xla/third_party/stablehlo/temporary.patch @@ -2381,6 +2381,15 @@ diff --ruN a/stablehlo/stablehlo/transforms/ChloLegalizeToStablehlo.cpp b/stable results.push_back(dotGeneral); start = limit; +@@ -2747,7 +2747,7 @@ + } + rewriter.replaceOpWithNewOp( + op, op.getResult().getType(), ValueRange{op.getLhs(), op.getRhs()}, +- attributes); ++ typename DotGeneralOp::Properties{}, attributes); + return success(); + } + @@ -2930,8 +2930,18 @@ SmallVector sliceShape(inputType.getShape().begin(), inputType.getShape().end()); @@ -2519,6 +2528,32 @@ diff --ruN a/stablehlo/stablehlo/transforms/PassUtils.h b/stablehlo/stablehlo/tr Value val); // Check if any of the given types are mlir::quant::QuantizedType. +diff --ruN a/stablehlo/stablehlo/transforms/ShapeLegalizeToStablehlo.cpp b/stablehlo/stablehlo/transforms/ShapeLegalizeToStablehlo.cpp +--- stablehlo/stablehlo/transforms/ShapeLegalizeToStablehlo.cpp ++++ stablehlo/stablehlo/transforms/ShapeLegalizeToStablehlo.cpp +@@ -520,6 +520,7 @@ + } + + rewriter.replaceOpWithNewOp(op, op->getResultTypes(), operandsI32, ++ typename OpType::Properties{}, + op->getAttrs()); + return success(); + } +diff --ruN a/stablehlo/stablehlo/transforms/StablehloCanonicalizeDynamism.cpp b/stablehlo/stablehlo/transforms/StablehloCanonicalizeDynamism.cpp +--- stablehlo/stablehlo/transforms/StablehloCanonicalizeDynamism.cpp ++++ stablehlo/stablehlo/transforms/StablehloCanonicalizeDynamism.cpp +@@ -97,8 +97,9 @@ + } + newOperands.push_back(operand.get()); + } +- rewriter.replaceOpWithNewOp(op, op.getResultTypes(), +- newOperands, newAttrs); ++ rewriter.replaceOpWithNewOp( ++ op, op.getResultTypes(), newOperands, ++ typename CustomCallOp::Properties{}, newAttrs); + return success(); + } + }; diff --ruN a/stablehlo/stablehlo/transforms/StablehloComplexMathExpander.cpp b/stablehlo/stablehlo/transforms/StablehloComplexMathExpander.cpp --- stablehlo/stablehlo/transforms/StablehloComplexMathExpander.cpp +++ stablehlo/stablehlo/transforms/StablehloComplexMathExpander.cpp @@ -2562,7 +2597,15 @@ diff --ruN a/stablehlo/stablehlo/transforms/StablehloComplexMathExpander.cpp b/s diff --ruN a/stablehlo/stablehlo/transforms/StablehloLegalizeQuantToMath.cpp b/stablehlo/stablehlo/transforms/StablehloLegalizeQuantToMath.cpp --- stablehlo/stablehlo/transforms/StablehloLegalizeQuantToMath.cpp +++ stablehlo/stablehlo/transforms/StablehloLegalizeQuantToMath.cpp -@@ -865,8 +865,9 @@ +@@ -636,6 +636,7 @@ + // Execute conversion target op. + SmallVector operands{lhsFloat32Tensor, rhsFloat32Tensor}; + rewriter.replaceOpWithNewOp(op, resFloat32TensorType, operands, ++ typename OpType::Properties{}, + op->getAttrs()); + return success(); + } +@@ -865,8 +866,9 @@ Type resultType, Value& lhs, Value& rhs, ArrayRef attrs, const DotLikeDimensionNumbers& dims) { @@ -2574,7 +2617,7 @@ diff --ruN a/stablehlo/stablehlo/transforms/StablehloLegalizeQuantToMath.cpp b/s } // Template specialization for Convolution op. -@@ -915,8 +916,9 @@ +@@ -915,8 +917,9 @@ } } } diff --git a/third_party/xla/xla/backends/cpu/codegen/ir_compiler.cc b/third_party/xla/xla/backends/cpu/codegen/ir_compiler.cc index acfc960021cb83..b99fbcce5cbe15 100644 --- a/third_party/xla/xla/backends/cpu/codegen/ir_compiler.cc +++ b/third_party/xla/xla/backends/cpu/codegen/ir_compiler.cc @@ -431,10 +431,6 @@ llvm::Error IrCompiler::RunIrPasses(llvm::Module& module, llvm::ModulePassManager pm; - if (options_.dfsan_enabled) { - pm.addPass(llvm::DataFlowSanitizerPass(options_.dfsan_abi_list_files)); - } - llvm::OptimizationLevel opt_level = GetOptimizationLevel(options_); if (opt_level == llvm::OptimizationLevel::O0) { pm.addPass(pb.buildO0DefaultPipeline(opt_level)); @@ -474,8 +470,8 @@ llvm::Error IrCompiler::RunIrPasses(llvm::Module& module, codegen::intrinsic::RunInlineAndOptPasses(module); } - // Must stay last: middle-end passes behave differently on instructions that - // already carry `contract`. + // Must run after all optimization passes: middle-end passes behave + // differently on instructions that already carry `contract`. // // TODO(b/560320144): `AllowFPOpFusion = Fast` is deliberately still set in // service/cpu/cpu_aot_loader.cc:53, tools/hlo_opt/cpu_opt.cc:217, @@ -484,9 +480,28 @@ llvm::Error IrCompiler::RunIrPasses(llvm::Module& module, // has landed. llvm_ir::SetAllowContractOnFpArithmetic(module); + // Sanitizer instrumentation must be the last IR transformation. + if (options_.dfsan_enabled) { + // The transformations immediately above are not visible to the analysis + // manager; clear its cache. + mam.clear(); + + RunSanitizerPasses(module, mam); + } + return llvm::Error::success(); } +void IrCompiler::RunSanitizerPasses(llvm::Module& module, + llvm::ModuleAnalysisManager& mam) const { + llvm::ModulePassManager pm; + + if (options_.dfsan_enabled) { + pm.addPass(llvm::DataFlowSanitizerPass(options_.dfsan_abi_list_files)); + } + pm.run(module, mam); +} + std::unique_ptr IrCompiler::EmitMachineCode( llvm::Module& module, llvm::TargetMachine* target_machine) const { // Buffer for holding machine code prior to constructing the ObjectFile. diff --git a/third_party/xla/xla/backends/cpu/codegen/ir_compiler.h b/third_party/xla/xla/backends/cpu/codegen/ir_compiler.h index c170bc0755f9c7..dcd28666904a3e 100644 --- a/third_party/xla/xla/backends/cpu/codegen/ir_compiler.h +++ b/third_party/xla/xla/backends/cpu/codegen/ir_compiler.h @@ -30,6 +30,7 @@ limitations under the License. #include "llvm/IR/FMF.h" #include "llvm/IR/LegacyPassManager.h" #include "llvm/IR/Module.h" +#include "llvm/IR/PassManager.h" #include "llvm/Object/ObjectFile.h" #include "llvm/Support/CodeGen.h" #include "llvm/Support/Error.h" @@ -140,6 +141,11 @@ class IrCompiler : public llvm::orc::IRCompileLayer::IRCompiler { // races when calling user provided compilation hooks. absl::Mutex mutex_; CompilationHooks hooks_ ABSL_GUARDED_BY(mutex_); + + // Runs the enabled sanitizer passes. Must be the last IR transformation, + // otherwise any later pass would produce uninstrumented code. + void RunSanitizerPasses(llvm::Module& module, + llvm::ModuleAnalysisManager& mam) const; }; } // namespace xla::cpu diff --git a/third_party/xla/xla/backends/cpu/codegen/ir_compiler_test.cc b/third_party/xla/xla/backends/cpu/codegen/ir_compiler_test.cc index 2d48a151a0def5..d9763c05a26112 100644 --- a/third_party/xla/xla/backends/cpu/codegen/ir_compiler_test.cc +++ b/third_party/xla/xla/backends/cpu/codegen/ir_compiler_test.cc @@ -43,6 +43,7 @@ limitations under the License. #include "llvm/Support/SourceMgr.h" #include "llvm/Support/TargetSelect.h" #include "llvm/Target/TargetMachine.h" +#include "llvm/TargetParser/Host.h" #include "llvm/TargetParser/Triple.h" #include "xla/backends/cpu/codegen/kernel_api_ir_builder.h" #include "xla/backends/cpu/codegen/object_buffer_identifier.h" @@ -284,6 +285,19 @@ TEST(IrCompilerTest, TargetMachineOptionsAreCorrectlySet) { "+foo-feature,-bar-feature"); } +TEST(IrCompilerTest, InferTargetMachineWithEmptyTriple) { + ASSERT_OK_AND_ASSIGN( + TargetMachineOptions target_machine_options, + TargetMachineOptions::FromProto(TargetMachineOptionsProto())); + ASSERT_OK_AND_ASSIGN( + std::unique_ptr target_machine, + IrCompiler::InferTargetMachine(llvm::TargetOptions(), + llvm::CodeGenOptLevel::Default, + target_machine_options)); + EXPECT_EQ(target_machine->getTargetTriple().getTriple(), + llvm::sys::getProcessTriple()); +} + TEST(IrCompilerTest, EmitIntrinsicCall) { constexpr absl::string_view kModuleName = "test_module"; constexpr absl::string_view kMemcpyCall = R"( @@ -475,6 +489,50 @@ TEST(ObjectBufferIdentifierTest, EncodeAndExtract) { EXPECT_EQ(ExtractModuleIdentifier(""), ""); } +// Compiles the given LLVM IR using IrCompiler and returns the modified module. +static absl::StatusOr> CompileIr( + llvm::LLVMContext& context, absl::string_view ir, absl::string_view name, + IrCompiler::Options options = IrCompiler::Options()) { + ABSL_ASSIGN_OR_RETURN(auto ir_module, ParseModule(context, ir, name)); + IrCompiler::CompilationHooks hooks; + std::unique_ptr ir_compiler = + IrCompiler::Create(llvm::TargetOptions(), options, hooks); + ABSL_ASSIGN_OR_RETURN(auto target_machine, ir_compiler->build_target_machine()); + ir_module->setDataLayout(target_machine->createDataLayout()); + ir_module->setTargetTriple(target_machine->getTargetTriple()); + if (llvm::Error err = (*ir_compiler)(*ir_module).takeError()) { + return Internal("IrCompiler failed: %s", llvm::toString(std::move(err))); + } + return ir_module; +} + +TEST(IrCompilerTest, DataFlowSanitizerInstrumentsPolynomialApproximations) { + constexpr absl::string_view kExpIr = R"( + declare float @llvm.exp.f32(float) + + define float @test_exp(float %x) { + %res = call float @llvm.exp.f32(float %x) + ret float %res + } + )"; + + llvm::LLVMContext context; + IrCompiler::Options options{ + /*opt_level=*/llvm::CodeGenOptLevel::None, + /*optimize_for_size=*/false, + TargetMachineOptions(kTargetTripleForHost, kTargetCpuForHost, ""), + }; + options.dfsan_enabled = true; + + ASSERT_OK_AND_ASSIGN(auto ir_module, + CompileIr(context, kExpIr, "test_exp_module", options)); + + auto ir = llvm_ir::DumpToString(ir_module.get()); + EXPECT_THAT(ir, HasSubstr("@test_exp.dfsan")); + EXPECT_THAT(ir, HasSubstr("fmul contract")); + EXPECT_THAT(ir, HasSubstr("store i8 %1, ptr @__dfsan_retval_tls")); +} + } // namespace } // namespace xla::cpu diff --git a/third_party/xla/xla/backends/cpu/target_machine_options.cc b/third_party/xla/xla/backends/cpu/target_machine_options.cc index 13317428ef4a29..66241008a9f5ac 100644 --- a/third_party/xla/xla/backends/cpu/target_machine_options.cc +++ b/third_party/xla/xla/backends/cpu/target_machine_options.cc @@ -32,6 +32,7 @@ limitations under the License. #include "absl/strings/string_view.h" #include "llvm/ADT/StringRef.h" // IWYU pragma: keep #include "llvm/TargetParser/Host.h" +#include "llvm/TargetParser/Triple.h" #include "xla/backends/cpu/codegen/cpu_features.h" #include "xla/service/cpu/executable.pb.h" #include "xla/util.h" @@ -103,15 +104,23 @@ GetEnabledAndDisabledFeatures(const std::vector& features) { return std::make_pair(enabled_features, disabled_features); } +// Preserves empty triples instead of normalizing them to "unknown" so that +// llvm::EngineBuilder::selectTarget falls back to the host process triple. +std::string NormalizeTriple(absl::string_view triple) { + return triple.empty() ? "" + : llvm::Triple::normalize( + llvm::StringRef(triple.data(), triple.size())); +} + } // namespace TargetMachineOptions::TargetMachineOptions() { - triple_ = llvm::sys::getDefaultTargetTriple(); + triple_ = NormalizeTriple(llvm::sys::getDefaultTargetTriple()); cpu_ = llvm::sys::getHostCPUName(); } TargetMachineOptions::TargetMachineOptions(const DebugOptions& debug_options) { - triple_ = llvm::sys::getDefaultTargetTriple(); + triple_ = NormalizeTriple(llvm::sys::getDefaultTargetTriple()); auto xla_cpu_max_isa = CpuFeatureFromString(debug_options.xla_cpu_max_isa()); auto detected_machine_attributes = DetectMachineAttributes(xla_cpu_max_isa); @@ -131,7 +140,7 @@ TargetMachineOptions::TargetMachineOptions(const DebugOptions& debug_options) { TargetMachineOptions::TargetMachineOptions(absl::string_view triple, absl::string_view cpu, absl::string_view features) - : triple_(triple), cpu_(cpu) { + : triple_(NormalizeTriple(triple)), cpu_(cpu) { std::vector features_vec = absl::StrSplit(features, ','); std::tie(enabled_features_, disabled_features_) = GetEnabledAndDisabledFeatures(features_vec); @@ -172,7 +181,7 @@ TargetMachineOptions::TargetMachineOptions( std::string triple, std::string cpu, std::vector enabled_features, std::vector disabled_features) - : triple_(std::move(triple)), + : triple_(NormalizeTriple(triple)), cpu_(std::move(cpu)), enabled_features_(std::move(enabled_features)), disabled_features_(std::move(disabled_features)) {} diff --git a/third_party/xla/xla/backends/cpu/target_machine_options_test.cc b/third_party/xla/xla/backends/cpu/target_machine_options_test.cc index 6ec1c3581221f8..2bab9d6c936d71 100644 --- a/third_party/xla/xla/backends/cpu/target_machine_options_test.cc +++ b/third_party/xla/xla/backends/cpu/target_machine_options_test.cc @@ -21,6 +21,7 @@ limitations under the License. #include "llvm/ADT/StringMap.h" #include "llvm/ADT/StringRef.h" #include "llvm/TargetParser/Host.h" +#include "llvm/TargetParser/Triple.h" #include "xla/debug_options_flags.h" #include "xla/service/cpu/executable.pb.h" #include "xla/tsl/lib/core/status_test_util.h" @@ -160,14 +161,36 @@ TEST(TargetMachineOptionsTest, TestTargeteMachineOptionsFeaturesAreSorted) { TEST(TargetMachineOptionsTest, TargetMachineOptionsDefaultConstructor) { TargetMachineOptions options; - EXPECT_EQ(options.triple(), llvm::sys::getDefaultTargetTriple()); + EXPECT_EQ(options.triple(), + llvm::Triple::normalize(llvm::sys::getDefaultTargetTriple())); EXPECT_EQ(options.cpu(), llvm::sys::getHostCPUName()); EXPECT_EQ(options.GetTargetMachineFeatures(), ""); } +TEST(TargetMachineOptionsTest, NormalizesTriple) { + TargetMachineOptions options("aarch64-linux-gnu", "generic", ""); + EXPECT_EQ(options.triple(), "aarch64-unknown-linux-gnu"); +} + +// An empty triple must stay empty rather than normalizing to "unknown", so that +// LLVM's EngineBuilder::selectTarget falls back to the host process triple when +// no target triple was specified. +TEST(TargetMachineOptionsTest, EmptyTripleIsNotNormalized) { + TargetMachineOptions options("", "generic", ""); + EXPECT_EQ(options.triple(), ""); +} + +TEST(TargetMachineOptionsTest, FromEmptyProtoPreservesEmptyTriple) { + TF_ASSERT_OK_AND_ASSIGN( + TargetMachineOptions options, + TargetMachineOptions::FromProto(TargetMachineOptionsProto())); + EXPECT_EQ(options.triple(), ""); +} + TEST(TargetMachineOptionsTest, NativeMatchesLLVMBehaviour) { TargetMachineOptions options = TargetMachineOptions::Native(); - EXPECT_EQ(options.triple(), llvm::sys::getDefaultTargetTriple()); + EXPECT_EQ(options.triple(), + llvm::Triple::normalize(llvm::sys::getDefaultTargetTriple())); EXPECT_EQ(options.cpu(), llvm::sys::getHostCPUName()); llvm::StringMap expected_features = llvm::sys::getHostCPUFeatures(); diff --git a/third_party/xla/xla/backends/gpu/autotuner/BUILD b/third_party/xla/xla/backends/gpu/autotuner/BUILD index fec89081d96b7b..9309c02d17b1b3 100644 --- a/third_party/xla/xla/backends/gpu/autotuner/BUILD +++ b/third_party/xla/xla/backends/gpu/autotuner/BUILD @@ -86,10 +86,13 @@ cc_library( "//xla/service/gpu/model:gpu_indexing_performance_model", "//xla/stream_executor:device_description", "//xla/stream_executor/gpu:tma_metadata", + "//xla/tsl/concurrency:executor", + "//xla/tsl/platform:env", "@com_google_absl//absl/algorithm:container", "@com_google_absl//absl/base:nullability", "@com_google_absl//absl/container:inlined_vector", "@com_google_absl//absl/log", + "@com_google_absl//absl/log:check", "@com_google_absl//absl/status", "@com_google_absl//absl/status:status_macros", "@com_google_absl//absl/status:statusor", @@ -130,14 +133,17 @@ xla_test( "//xla/service/gpu:backend_configs_cc", "//xla/service/gpu:ir_emission_utils", "//xla/service/gpu:nvptx_compiler_impl", + "//xla/service/gpu/model:gpu_indexing_performance_model", "//xla/stream_executor:platform", "//xla/stream_executor:stream_executor_h", + "//xla/tsl/platform:env", "//xla/tsl/platform:statusor", "//xla/tsl/util/proto:proto_matchers", "@com_google_absl//absl/log:check", "@com_google_absl//absl/status:status_matchers", "@com_google_absl//absl/status:statusor", "@com_google_googletest//:gtest_main", + "@llvm-project//mlir:IR", ], ) @@ -553,6 +559,7 @@ cc_library( "//xla/hlo/pass:hlo_pass_pipeline", "//xla/service:compiler", "//xla/service:hlo_cost_analysis", + "//xla/service/gpu/model:gpu_indexing_performance_model", "//xla/stream_executor:device_description", "//xla/stream_executor:stream_executor_h", "//xla/stream_executor/cuda:cuda_platform_id", @@ -688,6 +695,7 @@ cc_library( "//xla/hlo/pass:hlo_pass_pipeline", "//xla/service:compiler", "//xla/service:hlo_cost_analysis", + "//xla/service/gpu/model:gpu_indexing_performance_model", "//xla/stream_executor:device_description", "//xla/stream_executor:stream_executor_h", "//xla/stream_executor/platform:platform_object_registry", @@ -708,6 +716,7 @@ cc_library( "//xla/hlo/analysis:alias_info", "//xla/service:compiler", "//xla/service:hlo_cost_analysis", + "//xla/service/gpu/model:gpu_indexing_performance_model", "//xla/stream_executor:device_address_allocator", "//xla/stream_executor:stream_executor_h", "@com_google_absl//absl/types:span", diff --git a/third_party/xla/xla/backends/gpu/autotuner/autotuner_main.cc b/third_party/xla/xla/backends/gpu/autotuner/autotuner_main.cc index 787ebad44635c4..0e1b7223f5ffc7 100644 --- a/third_party/xla/xla/backends/gpu/autotuner/autotuner_main.cc +++ b/third_party/xla/xla/backends/gpu/autotuner/autotuner_main.cc @@ -258,7 +258,8 @@ absl::StatusOr CreateAutotunerEnvironment( ConfigAssignerPass::GetEnabledBackends( stream_executor_0, allocator.get(), target_config.get(), alias_info.get(), debug_options, mlir_context.get(), - compiler->ShapeSizeBytesFunction(), compiler.get(), platform->id())); + compiler->ShapeSizeBytesFunction(), compiler.get(), platform->id(), + thread_pool.get())); AutotuneCacheContext ctx = AutotuneCacheContext::Create( target_config->device_description, autotuner_backends); diff --git a/third_party/xla/xla/backends/gpu/autotuner/block_level_emitter.cc b/third_party/xla/xla/backends/gpu/autotuner/block_level_emitter.cc index 88ec1af6050f4b..11ba9fd2bcb239 100644 --- a/third_party/xla/xla/backends/gpu/autotuner/block_level_emitter.cc +++ b/third_party/xla/xla/backends/gpu/autotuner/block_level_emitter.cc @@ -24,7 +24,9 @@ limitations under the License. #include #include "absl/algorithm/container.h" +#include "absl/base/nullability.h" #include "absl/container/inlined_vector.h" +#include "absl/log/check.h" #include "absl/log/log.h" #include "absl/status/status.h" #include "absl/status/status_macros.h" @@ -45,6 +47,8 @@ limitations under the License. #include "xla/service/instruction_fusion.h" #include "xla/stream_executor/device_description.h" #include "xla/stream_executor/gpu/tma_metadata.h" +#include "xla/tsl/concurrency/executor.h" +#include "xla/tsl/platform/threadpool.h" #include "xla/xla.pb.h" #include "triton/Version.h" @@ -54,6 +58,19 @@ using ::xla::xtile::BlockLevelFusionConfig; namespace { +// Returns the executor to evaluate tiling candidates on, or nullptr to evaluate +// them inline on the calling thread. +tsl::Executor* absl_nullable TilingSearchExecutor( + tsl::thread::ThreadPool* absl_nullable thread_pool) { + if (thread_pool == nullptr) { + return nullptr; + } + // The callers below block on the result, so running them on a thread of the + // pool they dispatch to would deadlock. + CHECK_EQ(thread_pool->CurrentThreadId(), -1); + return thread_pool->AsExecutor(); +} + std::unique_ptr Pack( const BlockLevelFusionConfig& block_level_config) { auto config = std::make_unique(); @@ -126,7 +143,8 @@ BlockLevelEmitterBackend::GetSupportedConfigs(const HloInstruction& instr) { ABSL_ASSIGN_OR_RETURN( TopKTiledRunTimeDataOrError tiled_runtime_data, indexing_performance_model_ - .TryFindTopKBestTilingsForFusionAsync(*fusion_adaptor, num_configs) + .TryFindTopKBestTilingsForFusionAsync( + *fusion_adaptor, num_configs, TilingSearchExecutor(thread_pool_)) .Await()); if (std::holds_alternative(tiled_runtime_data)) { @@ -156,7 +174,8 @@ BlockLevelEmitterBackend::GetCostModelConfig(const HloInstruction& instr) { ABSL_ASSIGN_OR_RETURN(TiledRunTimeDataOrError tiled_runtime_data_or_error, indexing_performance_model_ - .TryFindBestTilingForFusionAsync(*fusion_adaptor) + .TryFindBestTilingForFusionAsync( + *fusion_adaptor, TilingSearchExecutor(thread_pool_)) .Await()); if (const auto* fusion_decision = diff --git a/third_party/xla/xla/backends/gpu/autotuner/block_level_emitter.h b/third_party/xla/xla/backends/gpu/autotuner/block_level_emitter.h index 215d7093016bbb..eb6cb7416d3852 100644 --- a/third_party/xla/xla/backends/gpu/autotuner/block_level_emitter.h +++ b/third_party/xla/xla/backends/gpu/autotuner/block_level_emitter.h @@ -38,6 +38,10 @@ limitations under the License. #include "xla/service/hlo_cost_analysis.h" #include "xla/xla.pb.h" +namespace tsl::thread { +class ThreadPool; +} // namespace tsl::thread + namespace xla { namespace gpu { @@ -52,7 +56,9 @@ class BlockLevelEmitterBackend : public GpuCodegenBackend { const DebugOptions* absl_nonnull debug_options, Compiler* absl_nonnull compiler, HloCostAnalysis::ShapeSizeFunction shape_size_fn, - const Compiler::GpuTargetConfig* target_config) + const Compiler::GpuTargetConfig* target_config, + tsl::thread::ThreadPool* absl_nullable thread_pool = nullptr, + MlirContextPool* absl_nullable mlir_context_pool = nullptr) : GpuCodegenBackend(autotuner::Backend::BLOCK_LEVEL_EMITTER, debug_options, compiler, target_config), shape_size_fn_(std::move(shape_size_fn)), @@ -62,9 +68,11 @@ class BlockLevelEmitterBackend : public GpuCodegenBackend { shape_size_fn_, &mlir_context_, debug_options->xla_gpu_experimental_enable_tiling_propagation(), debug_options - ->xla_gpu_experimental_enable_same_shape_multi_output_fusion()), + ->xla_gpu_experimental_enable_same_shape_multi_output_fusion(), + mlir_context_pool), xla_gpu_experimental_all_fusions_with_triton_( - debug_options->xla_gpu_experimental_all_fusions_with_triton()) { + debug_options->xla_gpu_experimental_all_fusions_with_triton()), + thread_pool_(thread_pool) { RegisterSymbolicExprStorage(&mlir_context_); } @@ -100,6 +108,7 @@ class BlockLevelEmitterBackend : public GpuCodegenBackend { GpuPerformanceModelWithIndexingAnalysis indexing_performance_model_; // If true, autotune all possible fusions with Triton. bool xla_gpu_experimental_all_fusions_with_triton_ = false; + tsl::thread::ThreadPool* absl_nullable thread_pool_ = nullptr; }; } // namespace gpu diff --git a/third_party/xla/xla/backends/gpu/autotuner/block_level_emitter_test.cc b/third_party/xla/xla/backends/gpu/autotuner/block_level_emitter_test.cc index 4dc530e3e877e5..e2140b84ef0129 100644 --- a/third_party/xla/xla/backends/gpu/autotuner/block_level_emitter_test.cc +++ b/third_party/xla/xla/backends/gpu/autotuner/block_level_emitter_test.cc @@ -24,6 +24,7 @@ limitations under the License. #include "absl/log/check.h" #include "absl/status/status_matchers.h" #include "absl/status/statusor.h" +#include "mlir/IR/MLIRContext.h" #include "xla/backends/autotuner/codegen_backend.h" #include "xla/codegen/xtile/xtile_config.pb.h" #include "xla/debug_options_flags.h" @@ -33,11 +34,14 @@ limitations under the License. #include "xla/service/executable.h" #include "xla/service/gpu/backend_configs.pb.h" #include "xla/service/gpu/ir_emission_utils.h" +#include "xla/service/gpu/model/gpu_indexing_performance_model.h" #include "xla/service/gpu/nvptx_compiler.h" #include "xla/service/platform_util.h" #include "xla/stream_executor/platform.h" #include "xla/stream_executor/stream_executor.h" +#include "xla/tsl/platform/env.h" #include "xla/tsl/platform/statusor.h" +#include "xla/tsl/platform/threadpool.h" #include "xla/tsl/util/proto/proto_matchers.h" #include "xla/xla.pb.h" @@ -293,6 +297,57 @@ ENTRY %main { EXPECT_THAT(executable, absl_testing::IsOk()); } +TEST_F(TritonBlockLevelFusionEmitterBackendTest, + GeneratesSameConfigsWithParallelTilingSearch) { + ASSERT_OK_AND_ASSIGN(std::unique_ptr module, + ParseAndReturnVerifiedModule(R"( +HloModule m +%wrapped_transpose_computation { + %param_0 = f32[16,1,64]{2,1,0} parameter(0) + ROOT %transpose.3.1 = f32[64,1,16]{2,1,0} transpose(%param_0), dimensions={2,1,0} +} + +ENTRY %main { + %p0 = f32[16,1,64]{2,1,0} parameter(0) + ROOT %wrapped_transpose = f32[64,1,16]{2,1,0} fusion(%p0), kind=kInput, + calls=%wrapped_transpose_computation +} +)")); + const HloInstruction& instr = + *module->entry_computation()->root_instruction(); + + tsl::thread::ThreadPool thread_pool(tsl::Env::Default(), "test_pool", 4); + // Mirrors the contexts GpuCompiler pools: multithreading is disabled, so the + // cost model must give each candidate its own context. + MlirContextPool mlir_context_pool( + [] { + return std::make_unique( + mlir::MLIRContext::Threading::DISABLED); + }, + /*preallocate=*/4); + BlockLevelEmitterBackend parallel_backend( + &debug_options_, &compiler_, compiler_.ShapeSizeBytesFunction(), + &target_config_, &thread_pool, &mlir_context_pool); + + ASSERT_OK_AND_ASSIGN(std::unique_ptr parallel_config, + parallel_backend.GetDefaultConfig(instr)); + ASSERT_OK_AND_ASSIGN(std::unique_ptr sequential_config, + backend_.GetDefaultConfig(instr)); + EXPECT_THAT(*parallel_config, EqualsProto(*sequential_config)); + + ASSERT_OK_AND_ASSIGN( + std::vector> parallel_configs, + parallel_backend.GetSupportedConfigs(instr)); + ASSERT_OK_AND_ASSIGN( + std::vector> sequential_configs, + backend_.GetSupportedConfigs(instr)); + ASSERT_FALSE(parallel_configs.empty()); + ASSERT_EQ(parallel_configs.size(), sequential_configs.size()); + for (int i = 0; i < parallel_configs.size(); ++i) { + EXPECT_THAT(*parallel_configs[i], EqualsProto(*sequential_configs[i])); + } +} + TEST_F(TritonBlockLevelFusionEmitterBackendTest, Version) { EXPECT_NE(backend_.version(), ""); } diff --git a/third_party/xla/xla/backends/gpu/autotuner/factory.h b/third_party/xla/xla/backends/gpu/autotuner/factory.h index 29ec8fce04ee78..5f1488b0787b23 100644 --- a/third_party/xla/xla/backends/gpu/autotuner/factory.h +++ b/third_party/xla/xla/backends/gpu/autotuner/factory.h @@ -26,10 +26,15 @@ limitations under the License. #include "xla/backends/autotuner/codegen_backend.h" #include "xla/hlo/analysis/alias_info.h" #include "xla/service/compiler.h" +#include "xla/service/gpu/model/gpu_indexing_performance_model.h" #include "xla/service/hlo_cost_analysis.h" #include "xla/stream_executor/device_address_allocator.h" #include "xla/stream_executor/stream_executor.h" +namespace tsl::thread { +class ThreadPool; +} // namespace tsl::thread + namespace xla { namespace gpu { @@ -45,7 +50,9 @@ struct GetCodegenBackends { const Compiler::GpuTargetConfig*, const AliasInfo* alias_info, mlir::MLIRContext* mlir_context, HloCostAnalysis::ShapeSizeFunction shape_size_fn, - absl::Span backend_allowlist)>; + absl::Span backend_allowlist, + tsl::thread::ThreadPool* thread_pool, + MlirContextPool* mlir_context_pool)>; }; } // namespace gpu diff --git a/third_party/xla/xla/backends/gpu/autotuner/factory_cuda.cc b/third_party/xla/xla/backends/gpu/autotuner/factory_cuda.cc index ccc89ea433a23f..c26b0da258bace 100644 --- a/third_party/xla/xla/backends/gpu/autotuner/factory_cuda.cc +++ b/third_party/xla/xla/backends/gpu/autotuner/factory_cuda.cc @@ -38,6 +38,7 @@ limitations under the License. #include "xla/hlo/analysis/alias_info.h" #include "xla/hlo/pass/hlo_pass_pipeline.h" #include "xla/service/compiler.h" +#include "xla/service/gpu/model/gpu_indexing_performance_model.h" #include "xla/service/hlo_cost_analysis.h" #include "xla/stream_executor/cuda/cuda_platform_id.h" #include "xla/stream_executor/device_description.h" @@ -77,7 +78,8 @@ std::vector> GetCodegenBackendsForCuda( const DebugOptions* debug_options, Compiler* compiler, const Compiler::GpuTargetConfig* target_config, const AliasInfo* alias_info, MLIRContext* mlir_context, HloCostAnalysis::ShapeSizeFunction shape_size_fn, - absl::Span backend_allowlist) { + absl::Span backend_allowlist, + tsl::thread::ThreadPool* thread_pool, MlirContextPool* mlir_context_pool) { // Selecting the "first' config in the autotuner is backend order dependent. // To make all tests pass we need to keep the CuDnn backend first and the // Triton backend second. @@ -97,7 +99,8 @@ std::vector> GetCodegenBackendsForCuda( backends.push_back(std::make_unique( debug_options, compiler, target_config)); backends.push_back(std::make_unique( - debug_options, compiler, shape_size_fn, target_config)); + debug_options, compiler, shape_size_fn, target_config, thread_pool, + mlir_context_pool)); if (!backend_allowlist.empty()) { backends.erase( diff --git a/third_party/xla/xla/backends/gpu/autotuner/factory_rocm.cc b/third_party/xla/xla/backends/gpu/autotuner/factory_rocm.cc index e8fbcb7e5e28fa..96a150ecd75a9d 100644 --- a/third_party/xla/xla/backends/gpu/autotuner/factory_rocm.cc +++ b/third_party/xla/xla/backends/gpu/autotuner/factory_rocm.cc @@ -40,6 +40,7 @@ limitations under the License. #include "xla/hlo/analysis/alias_info.h" #include "xla/hlo/pass/hlo_pass_pipeline.h" #include "xla/service/compiler.h" +#include "xla/service/gpu/model/gpu_indexing_performance_model.h" #include "xla/service/hlo_cost_analysis.h" #include "xla/stream_executor/device_description.h" #include "xla/stream_executor/platform/platform_object_registry.h" @@ -84,7 +85,8 @@ std::vector> GetCodegenBackendsForROCm( const DebugOptions* debug_options, Compiler* compiler, const Compiler::GpuTargetConfig* target_config, const AliasInfo* alias_info, MLIRContext* mlir_context, HloCostAnalysis::ShapeSizeFunction shape_size_fn, - absl::Span backend_allowlist) { + absl::Span backend_allowlist, + tsl::thread::ThreadPool* thread_pool, MlirContextPool* mlir_context_pool) { std::vector> backends; backends.push_back(std::make_unique( debug_options, compiler, target_config, alias_info, mlir_context)); @@ -103,7 +105,8 @@ std::vector> GetCodegenBackendsForROCm( backends.push_back(std::make_unique( debug_options, compiler, target_config)); backends.push_back(std::make_unique( - debug_options, compiler, shape_size_fn, target_config)); + debug_options, compiler, shape_size_fn, target_config, thread_pool, + mlir_context_pool)); if (!backend_allowlist.empty()) { backends.erase( diff --git a/third_party/xla/xla/backends/gpu/autotuner/factory_test.cc b/third_party/xla/xla/backends/gpu/autotuner/factory_test.cc index b9cd546490dcd5..b2de8f1f15c566 100644 --- a/third_party/xla/xla/backends/gpu/autotuner/factory_test.cc +++ b/third_party/xla/xla/backends/gpu/autotuner/factory_test.cc @@ -116,7 +116,8 @@ TEST_P(FactoryTest, GetCodegenBackends) { get_codegen_backends( stream_executor_, &allocator_, &debug_options_, compiler_.get(), &target_config_, &alias_info, &mlir_context, - /*shape_size_fn=*/[](const Shape&) { return 0; }, GetParam().names); + /*shape_size_fn=*/[](const Shape&) { return 0; }, GetParam().names, + /*thread_pool=*/nullptr, /*mlir_context_pool=*/nullptr); EXPECT_EQ(backends.size(), GetParam().expected_num_backends); } else { GTEST_SKIP() << "Skipping test for platform " << platform_->id(); diff --git a/third_party/xla/xla/backends/gpu/transforms/BUILD b/third_party/xla/xla/backends/gpu/transforms/BUILD index 09a84ce25588e7..77efce2ffb9433 100644 --- a/third_party/xla/xla/backends/gpu/transforms/BUILD +++ b/third_party/xla/xla/backends/gpu/transforms/BUILD @@ -251,6 +251,7 @@ cc_library( "//xla/service/gpu/model:fusion_analysis_cache", "//xla/service/gpu/model:gpu_indexing_performance_model", "//xla/stream_executor:device_description", + "//xla/tsl/platform:env", "@com_google_absl//absl/container:flat_hash_set", "@com_google_absl//absl/container:inlined_vector", "@com_google_absl//absl/log", @@ -281,8 +282,10 @@ xla_cc_test( "//xla/service/gpu:backend_configs_cc", "//xla/service/gpu:gpu_device_info_for_tests", "//xla/service/gpu:ir_emission_utils", + "//xla/service/gpu/model:gpu_indexing_performance_model", "//xla/stream_executor:device_description", "//xla/stream_executor/cuda:cuda_compute_capability", + "//xla/tsl/platform:env", "//xla/tsl/platform:statusor", "@com_google_absl//absl/log:check", "@com_google_absl//absl/status", diff --git a/third_party/xla/xla/backends/gpu/transforms/cudnn_fusion_compiler.cc b/third_party/xla/xla/backends/gpu/transforms/cudnn_fusion_compiler.cc index 979cab24fce067..519fa4468f07b3 100644 --- a/third_party/xla/xla/backends/gpu/transforms/cudnn_fusion_compiler.cc +++ b/third_party/xla/xla/backends/gpu/transforms/cudnn_fusion_compiler.cc @@ -547,18 +547,21 @@ class ConvDimensionAdapter { }); }; - // Pattern 1: hlo -> broadcast - if (all_users_are_broadcast(hlo)) { - return hlo->users()[0]; + // Trace through single-user chains that preserve the 1D dimensions + // (e.g. hlo -> convert -> ... -> elementwise -> ... -> broadcast). + const HloInstruction* current = hlo; + while (current->user_count() == 1) { + const HloInstruction* user = current->users()[0]; + if (user->IsElementwise() && + ShapeUtil::SameDimensions(user->shape(), hlo->shape())) { + current = user; + continue; + } + break; } - // Pattern 2: hlo -> convert -> broadcast - if (hlo->user_count() == 1 && - hlo->users()[0]->opcode() == HloOpcode::kConvert) { - const HloInstruction* convert = hlo->users()[0]; - if (all_users_are_broadcast(convert)) { - return convert->users()[0]; - } + if (all_users_are_broadcast(current)) { + return current->users()[0]; } return nullptr; diff --git a/third_party/xla/xla/backends/gpu/transforms/cudnn_fusion_compiler_deviceless_test.cc b/third_party/xla/xla/backends/gpu/transforms/cudnn_fusion_compiler_deviceless_test.cc index eb045d846b8863..7ef798c399a965 100644 --- a/third_party/xla/xla/backends/gpu/transforms/cudnn_fusion_compiler_deviceless_test.cc +++ b/third_party/xla/xla/backends/gpu/transforms/cudnn_fusion_compiler_deviceless_test.cc @@ -379,6 +379,41 @@ TEST_F(CudnnFusionCompilerDevicelessTest, GroupedFp8ConvDeliversVerdict) { DevicelessFusionSupport::kUnknown); } +TEST_F(CudnnFusionCompilerDevicelessTest, + ConvWith1DBatchBroadcastEpilogueSupported) { + constexpr absl::string_view kConvWithBatchBroadcastHlo = R"( + ENTRY e { + input = f32[2,10,10,16] parameter(0) + filter = f32[16,3,3,16] parameter(1) + mask = bf16[2] parameter(2) + mask_f32 = f32[2] convert(mask) + c_neg1 = f32[] constant(-1) + c_neg1_bcast = f32[2] broadcast(c_neg1), dimensions={} + sub = f32[2] add(mask_f32, c_neg1_bcast) + zero = f32[] constant(0) + zero_bcast = f32[2] broadcast(zero), dimensions={} + max = f32[2] maximum(zero_bcast, sub) + mask_bcast = f32[2,10,10,16] broadcast(max), dimensions={0} + conv = f32[2,10,10,16] convolution(input, filter), + window={size=3x3 pad=1_1x1_1}, dim_labels=b01f_o01i->b01f + ROOT out = f32[2,10,10,16] multiply(conv, mask_bcast) + })"; + + ASSERT_OK_AND_ASSIGN(GpuTargetConfig target_config, + DevicelessTargetConfig(GpuModel::H100_SXM)); + ASSERT_OK_AND_ASSIGN( + std::unique_ptr module, + BuildConvFusionModule( + kConvWithBatchBroadcastHlo, target_config.device_description, + se::dnn::VersionInfo(target_config.device_description.dnn_version()), + CONVOLUTION_KIND_FPROP)); + const HloFusionInstruction* fusion = FindCudnnFusion(*module); + ASSERT_NE(fusion, nullptr); + EXPECT_EQ(CuDnnFusionCompiler::SupportsFusionDeviceless( + target_config.device_description, *fusion), + DevicelessFusionSupport::kSupported); +} + // The deviceless verdict must agree with live plan enumeration on the // executor's own device: SupportsFusionDeviceless(desc, fusion) is kSupported // iff GetAvailablePlanCount(executor, desc, fusion) > 0. diff --git a/third_party/xla/xla/backends/gpu/transforms/fusion_block_level_rewriter.cc b/third_party/xla/xla/backends/gpu/transforms/fusion_block_level_rewriter.cc index 027abcc82d77a6..0d79ddb6e1612a 100644 --- a/third_party/xla/xla/backends/gpu/transforms/fusion_block_level_rewriter.cc +++ b/third_party/xla/xla/backends/gpu/transforms/fusion_block_level_rewriter.cc @@ -51,6 +51,7 @@ limitations under the License. #include "xla/service/pattern_matcher.h" #include "xla/shape.h" #include "xla/stream_executor/device_description.h" +#include "xla/tsl/platform/threadpool.h" #include "xla/xla.pb.h" #include "xla/xla_data.pb.h" @@ -215,7 +216,9 @@ absl::StatusOr ProcessFusionInstruction( HloFusionInstruction* fusion_instruction, const se::DeviceDescription& device_info, HloCostAnalysis::ShapeSizeFunction shape_size, - mlir::MLIRContext* mlir_context, bool use_experimental_tiling) { + mlir::MLIRContext* mlir_context, bool use_experimental_tiling, + tsl::thread::ThreadPool* thread_pool = nullptr, + MlirContextPool* mlir_context_pool = nullptr) { bool dump_fusion_visualization = fusion_instruction->GetModule() ->config() .debug_options() @@ -263,14 +266,17 @@ absl::StatusOr ProcessFusionInstruction( fusion_instruction->GetModule() ->config() .debug_options() - .xla_gpu_experimental_enable_same_shape_multi_output_fusion()); + .xla_gpu_experimental_enable_same_shape_multi_output_fusion(), + mlir_context_pool); auto fusion_adaptor = HloFusionAdaptor::ForInstruction( Cast(fusion_instruction)); ABSL_ASSIGN_OR_RETURN(TiledRunTimeDataOrError tiled_runtime_data_or_error, indexing_performance_model - .TryFindBestTilingForFusionAsync(*fusion_adaptor) + .TryFindBestTilingForFusionAsync( + *fusion_adaptor, + thread_pool ? thread_pool->AsExecutor() : nullptr) .Await()); if (const auto* fusion_decision = @@ -341,7 +347,8 @@ absl::StatusOr FusionBlockLevelRewriter::RunImpl( fusion_instruction, device_info_, shape_size_, mlir_context_, module->config() .debug_options() - .xla_gpu_experimental_enable_tiling_propagation())); + .xla_gpu_experimental_enable_tiling_propagation(), + thread_pool_, mlir_context_pool_)); has_changed |= changed; } diff --git a/third_party/xla/xla/backends/gpu/transforms/fusion_block_level_rewriter.h b/third_party/xla/xla/backends/gpu/transforms/fusion_block_level_rewriter.h index 39546279200157..3a2eace99a2a1a 100644 --- a/third_party/xla/xla/backends/gpu/transforms/fusion_block_level_rewriter.h +++ b/third_party/xla/xla/backends/gpu/transforms/fusion_block_level_rewriter.h @@ -22,9 +22,14 @@ limitations under the License. #include "mlir/IR/MLIRContext.h" #include "xla/hlo/ir/hlo_module.h" #include "xla/hlo/pass/hlo_pass_interface.h" +#include "xla/service/gpu/model/gpu_indexing_performance_model.h" #include "xla/service/hlo_cost_analysis.h" #include "xla/stream_executor/device_description.h" +namespace tsl::thread { +class ThreadPool; +} // namespace tsl::thread + namespace xla { namespace gpu { @@ -33,10 +38,14 @@ class FusionBlockLevelRewriter : public HloModulePass { explicit FusionBlockLevelRewriter( const se::DeviceDescription& device_info, HloCostAnalysis::ShapeSizeFunction shape_size, - mlir::MLIRContext* mlir_context) + mlir::MLIRContext* mlir_context, + tsl::thread::ThreadPool* thread_pool = nullptr, + MlirContextPool* mlir_context_pool = nullptr) : device_info_(device_info), shape_size_(shape_size), - mlir_context_(mlir_context) {} + mlir_context_(mlir_context), + thread_pool_(thread_pool), + mlir_context_pool_(mlir_context_pool) {} absl::string_view name() const override { return "fusion-block-level-rewriter"; @@ -51,6 +60,8 @@ class FusionBlockLevelRewriter : public HloModulePass { const se::DeviceDescription& device_info_; HloCostAnalysis::ShapeSizeFunction shape_size_; mlir::MLIRContext* mlir_context_; + tsl::thread::ThreadPool* thread_pool_ = nullptr; + MlirContextPool* mlir_context_pool_ = nullptr; }; } // namespace gpu diff --git a/third_party/xla/xla/backends/gpu/transforms/fusion_block_level_rewriter_test.cc b/third_party/xla/xla/backends/gpu/transforms/fusion_block_level_rewriter_test.cc index 6b5c098f98a3f8..c6724f060d7f42 100644 --- a/third_party/xla/xla/backends/gpu/transforms/fusion_block_level_rewriter_test.cc +++ b/third_party/xla/xla/backends/gpu/transforms/fusion_block_level_rewriter_test.cc @@ -40,11 +40,14 @@ License. #include "xla/service/gpu/backend_configs.pb.h" #include "xla/service/gpu/gpu_device_info_for_tests.h" #include "xla/service/gpu/ir_emission_utils.h" +#include "xla/service/gpu/model/gpu_indexing_performance_model.h" #include "xla/service/hlo_cost_analysis.h" #include "xla/service/hlo_module_config.h" #include "xla/stream_executor/cuda/cuda_compute_capability.h" #include "xla/stream_executor/device_description.h" +#include "xla/tsl/platform/env.h" #include "xla/tsl/platform/statusor.h" +#include "xla/tsl/platform/threadpool.h" #include "xla/xla.pb.h" namespace xla::gpu { @@ -461,6 +464,37 @@ ENTRY entry { absl_testing::IsOkAndHolds(false)); } +TEST_P(FusionBlockLevelRewriterTest, RewritesFusionWithParallelTilingSearch) { + const absl::string_view hlo_text = R"( +fusion_computation { + param_0 = f32[128,128] parameter(0) + ROOT exp = f32[128,128] exponential(param_0) +} + +ENTRY entry { + param_0 = f32[128,128] parameter(0) + ROOT fusion = f32[128,128] fusion(param_0), kind=kCustom, + calls=fusion_computation, + backend_config={"fusion_backend_config":{"kind":"__triton"}} +})"; + ASSERT_OK_AND_ASSIGN(std::unique_ptr module, + ParseAndReturnVerifiedModule(hlo_text)); + tsl::thread::ThreadPool thread_pool(tsl::Env::Default(), "test_pool", 4); + // Mirrors the contexts GpuCompiler pools: multithreading is disabled, so the + // cost model must give each candidate its own context. + MlirContextPool mlir_context_pool( + [] { + return std::make_unique( + mlir::MLIRContext::Threading::DISABLED); + }, + /*preallocate=*/4); + EXPECT_THAT( + FusionBlockLevelRewriter(device_info_, HloCostAnalysis::DefaultShapeSize, + &mlir_context_, &thread_pool, &mlir_context_pool) + .Run(module.get()), + absl_testing::IsOkAndHolds(true)); +} + TEST_F(FusionBlockLevelRewriterTestBase, RewritesSameShapeMultiOutputFusionWithoutGeneralBlockLevelRewriter) { const absl::string_view hlo_text = R"hlo( diff --git a/third_party/xla/xla/backends/profiler/gpu/BUILD b/third_party/xla/xla/backends/profiler/gpu/BUILD index 8208e5ad72e141..f8355d54413e86 100644 --- a/third_party/xla/xla/backends/profiler/gpu/BUILD +++ b/third_party/xla/xla/backends/profiler/gpu/BUILD @@ -144,8 +144,10 @@ cc_library( "@com_google_absl//absl/base:no_destructor", "@com_google_absl//absl/container:flat_hash_map", "@com_google_absl//absl/log", + "@com_google_absl//absl/strings:string_view", "@com_google_absl//absl/types:span", "@local_config_cuda//cuda:cuda_headers", + "@local_config_cuda//cuda:cupti_headers", ], ) @@ -764,6 +766,7 @@ cc_library( ], visibility = ["//visibility:public"], deps = [ + ":cuda_version_variants", ":cupti_interface", ":cupti_marker_data_parser", ":cupti_utils", @@ -894,11 +897,13 @@ xla_cc_test( "no_mac", ], deps = [ + ":cuda_version_variants", ":cupti_buffer_events", ":cupti_collector", # buildcleaner: keep ":cupti_utils", "@com_google_googletest//:gtest", "@com_google_googletest//:gtest_main", + "@local_config_cuda//cuda:cupti_headers", ], ) diff --git a/third_party/xla/xla/backends/profiler/gpu/cuda_version_12080_newer.cc b/third_party/xla/xla/backends/profiler/gpu/cuda_version_12080_newer.cc index c0abb99b85a7c1..80f7160170cab6 100644 --- a/third_party/xla/xla/backends/profiler/gpu/cuda_version_12080_newer.cc +++ b/third_party/xla/xla/backends/profiler/gpu/cuda_version_12080_newer.cc @@ -14,6 +14,8 @@ limitations under the License. ==============================================================================*/ #include "absl/base/no_destructor.h" +#include "absl/strings/string_view.h" +#include "third_party/gpus/cuda/extras/CUPTI/include/cupti_activity.h" #include "third_party/gpus/cuda/extras/CUPTI/include/cupti_driver_cbid.h" #include "xla/backends/profiler/gpu/cuda_version_variants.h" @@ -35,6 +37,24 @@ const CbidCategoryMap& GetExtraCallbackIdCategories12080() { return *kCbidCategoryMap; } +absl::string_view GetExtraActivityOverheadKindString12080( + CUpti_ActivityOverheadKind kind) { + switch (kind) { + case CUPTI_ACTIVITY_OVERHEAD_RUNTIME_TRIGGERED_MODULE_LOADING: + return "RUNTIME_TRIGGERED_MODULE_LOADING"; + case CUPTI_ACTIVITY_OVERHEAD_LAZY_FUNCTION_LOADING: + return "LAZY_FUNCTION_LOADING"; + case CUPTI_ACTIVITY_OVERHEAD_COMMAND_BUFFER_FULL: + return "COMMAND_BUFFER_FULL"; + case CUPTI_ACTIVITY_OVERHEAD_ACTIVITY_BUFFER_REQUEST: + return "ACTIVITY_BUFFER_REQUEST"; + case CUPTI_ACTIVITY_OVERHEAD_UVM_ACTIVITY_INIT: + return "UVM_ACTIVITY_INIT"; + default: + return ""; + } +} + } // namespace cuda_versions } // namespace profiler diff --git a/third_party/xla/xla/backends/profiler/gpu/cuda_version_12080_older.cc b/third_party/xla/xla/backends/profiler/gpu/cuda_version_12080_older.cc index d385d64805a834..bc0776ee77802e 100644 --- a/third_party/xla/xla/backends/profiler/gpu/cuda_version_12080_older.cc +++ b/third_party/xla/xla/backends/profiler/gpu/cuda_version_12080_older.cc @@ -13,6 +13,8 @@ See the License for the specific language governing permissions and limitations under the License. ==============================================================================*/ +#include "absl/strings/string_view.h" +#include "third_party/gpus/cuda/extras/CUPTI/include/cupti_activity.h" #include "xla/backends/profiler/gpu/cuda_version_variants.h" namespace xla { @@ -23,6 +25,11 @@ const CbidCategoryMap& GetExtraCallbackIdCategories12080() { return EmptyCallbackIdCategories(); } +absl::string_view GetExtraActivityOverheadKindString12080( + CUpti_ActivityOverheadKind kind) { + return ""; +} + } // namespace cuda_versions } // namespace profiler diff --git a/third_party/xla/xla/backends/profiler/gpu/cuda_version_variants.h b/third_party/xla/xla/backends/profiler/gpu/cuda_version_variants.h index 81bd3dd2426c14..053acdc62305d0 100644 --- a/third_party/xla/xla/backends/profiler/gpu/cuda_version_variants.h +++ b/third_party/xla/xla/backends/profiler/gpu/cuda_version_variants.h @@ -17,8 +17,10 @@ limitations under the License. #define XLA_BACKENDS_PROFILER_GPU_CUDA_VERSION_VARIANTS_H_ #include "absl/container/flat_hash_map.h" +#include "absl/strings/string_view.h" #include "absl/types/span.h" #include "third_party/gpus/cuda/extras/CUPTI/include/cupti.h" +#include "third_party/gpus/cuda/extras/CUPTI/include/cupti_activity.h" #include "third_party/gpus/cuda/extras/CUPTI/include/cupti_callbacks.h" namespace xla { @@ -58,6 +60,10 @@ const CbidCategoryMap& GetExtraCallbackIdCategories12080(); // Resource CBIDs impacted only before/after 12.0. absl::Span GetCudaGraphTracingResourceCbids(); +// Overhead kinds introduced after CUDA 12.0 (available in CUDA 12.8+). +absl::string_view GetExtraActivityOverheadKindString12080( + CUpti_ActivityOverheadKind kind); + } // namespace cuda_versions } // namespace profiler } // namespace xla diff --git a/third_party/xla/xla/backends/profiler/gpu/cupti_buffer_events.cc b/third_party/xla/xla/backends/profiler/gpu/cupti_buffer_events.cc index 239809bd64bb53..a7caf563ab076b 100644 --- a/third_party/xla/xla/backends/profiler/gpu/cupti_buffer_events.cc +++ b/third_party/xla/xla/backends/profiler/gpu/cupti_buffer_events.cc @@ -25,6 +25,7 @@ limitations under the License. #include "absl/strings/string_view.h" #include "third_party/gpus/cuda/extras/CUPTI/include/cupti_activity.h" #include "third_party/gpus/cuda/include/cuda.h" +#include "xla/backends/profiler/gpu/cuda_version_variants.h" #include "xla/backends/profiler/gpu/cupti_interface.h" #include "xla/backends/profiler/gpu/cupti_marker_data_parser.h" #include "xla/backends/profiler/gpu/cupti_utils.h" @@ -115,23 +116,6 @@ using CuptiActivityMarkerTy = CUpti_ActivityMarker; constexpr int kCuptiActivityMarkerVersion = 1; #endif // CUDA_VERSION >= 11070 -// Maps an OverheadKind enum to a const string. -const char *getActivityOverheadKindString(CUpti_ActivityOverheadKind kind) { - switch (kind) { - case CUPTI_ACTIVITY_OVERHEAD_DRIVER_COMPILER: - return "COMPILER"; - case CUPTI_ACTIVITY_OVERHEAD_CUPTI_BUFFER_FLUSH: - return "BUFFER_FLUSH"; - case CUPTI_ACTIVITY_OVERHEAD_CUPTI_INSTRUMENTATION: - return "INSTRUMENTATION"; - case CUPTI_ACTIVITY_OVERHEAD_CUPTI_RESOURCE: - return "RESOURCE"; - default: - break; - } - return ""; -} - const char *getActivityUnifiedMemoryKindString( CUpti_ActivityUnifiedMemoryCounterKind kind) { switch (kind) { @@ -421,7 +405,7 @@ void AddCuptiOverheadActivityEvent(CuptiEventCollectorDelegate &collector, const CUpti_ActivityOverhead *overhead) { CuptiTracerEvent event{}; event.type = CuptiTracerEventType::Overhead; - event.name = getActivityOverheadKindString(overhead->overheadKind); + event.name = GetActivityOverheadKindString(overhead->overheadKind); event.source = CuptiTracerEventSource::Activity; event.start_time_ns = overhead->start; event.end_time_ns = overhead->end; @@ -840,5 +824,25 @@ absl::string_view GetMemoryKindName(int8_t memory_kind) { } } +std::string GetActivityOverheadKindString(CUpti_ActivityOverheadKind kind) { + switch (kind) { + case CUPTI_ACTIVITY_OVERHEAD_DRIVER_COMPILER: + return "COMPILER"; + case CUPTI_ACTIVITY_OVERHEAD_CUPTI_BUFFER_FLUSH: + return "BUFFER_FLUSH"; + case CUPTI_ACTIVITY_OVERHEAD_CUPTI_INSTRUMENTATION: + return "INSTRUMENTATION"; + case CUPTI_ACTIVITY_OVERHEAD_CUPTI_RESOURCE: + return "RESOURCE"; + default: + if (absl::string_view extra_str = + cuda_versions::GetExtraActivityOverheadKindString12080(kind); + !extra_str.empty()) { + return std::string(extra_str); + } + return absl::StrCat("Overhead::UNKNOWN:", static_cast(kind)); + } +} + } // namespace profiler } // namespace xla diff --git a/third_party/xla/xla/backends/profiler/gpu/cupti_buffer_events.h b/third_party/xla/xla/backends/profiler/gpu/cupti_buffer_events.h index a969c002b188d1..e454d30cb23b85 100644 --- a/third_party/xla/xla/backends/profiler/gpu/cupti_buffer_events.h +++ b/third_party/xla/xla/backends/profiler/gpu/cupti_buffer_events.h @@ -31,6 +31,7 @@ limitations under the License. #include "absl/strings/str_cat.h" #include "absl/strings/string_view.h" #include "absl/synchronization/mutex.h" +#include "third_party/gpus/cuda/extras/CUPTI/include/cupti_activity.h" #include "third_party/gpus/cuda/extras/CUPTI/include/cupti_callbacks.h" #include "xla/backends/profiler/gpu/string_deduper.h" #include "xla/tsl/profiler/utils/buffer_pool.h" @@ -169,6 +170,9 @@ inline std::string ToXStat(const KernelDetails& kernel_info, // Gets the name of the CUpti_ActivityMemoryKind value. absl::string_view GetMemoryKindName(int8_t memory_kind); +// Maps an OverheadKind enum to a string. +std::string GetActivityOverheadKindString(CUpti_ActivityOverheadKind kind); + enum class CuptiTracerEventType { Unsupported = 0, Kernel = 1, diff --git a/third_party/xla/xla/backends/profiler/gpu/cupti_buffer_events_test.cc b/third_party/xla/xla/backends/profiler/gpu/cupti_buffer_events_test.cc index ebb875f48ee9f0..d03fde5aeaef00 100644 --- a/third_party/xla/xla/backends/profiler/gpu/cupti_buffer_events_test.cc +++ b/third_party/xla/xla/backends/profiler/gpu/cupti_buffer_events_test.cc @@ -16,6 +16,8 @@ limitations under the License. #include "xla/backends/profiler/gpu/cupti_buffer_events.h" #include +#include "third_party/gpus/cuda/extras/CUPTI/include/cupti_activity.h" +#include "xla/backends/profiler/gpu/cuda_version_variants.h" namespace xla { namespace profiler { @@ -62,6 +64,45 @@ TEST(CuptiBufferEventsTest, EventInitialization) { EXPECT_EQ(event.graph_node_id, 11); } +TEST(CuptiBufferEventsTest, GetActivityOverheadKindString) { + EXPECT_EQ(GetActivityOverheadKindString(CUPTI_ACTIVITY_OVERHEAD_UNKNOWN), + "Overhead::UNKNOWN:0"); + EXPECT_EQ( + GetActivityOverheadKindString(CUPTI_ACTIVITY_OVERHEAD_DRIVER_COMPILER), + "COMPILER"); + EXPECT_EQ( + GetActivityOverheadKindString(CUPTI_ACTIVITY_OVERHEAD_CUPTI_BUFFER_FLUSH), + "BUFFER_FLUSH"); + EXPECT_EQ(GetActivityOverheadKindString( + CUPTI_ACTIVITY_OVERHEAD_CUPTI_INSTRUMENTATION), + "INSTRUMENTATION"); + EXPECT_EQ( + GetActivityOverheadKindString(CUPTI_ACTIVITY_OVERHEAD_CUPTI_RESOURCE), + "RESOURCE"); + if (!cuda_versions::GetExtraActivityOverheadKindString12080( + static_cast(4 << 16)) + .empty()) { + EXPECT_EQ(GetActivityOverheadKindString( + static_cast(4 << 16)), + "RUNTIME_TRIGGERED_MODULE_LOADING"); + EXPECT_EQ(GetActivityOverheadKindString( + static_cast(5 << 16)), + "LAZY_FUNCTION_LOADING"); + EXPECT_EQ(GetActivityOverheadKindString( + static_cast(6 << 16)), + "COMMAND_BUFFER_FULL"); + EXPECT_EQ(GetActivityOverheadKindString( + static_cast(7 << 16)), + "ACTIVITY_BUFFER_REQUEST"); + EXPECT_EQ(GetActivityOverheadKindString( + static_cast(8 << 16)), + "UVM_ACTIVITY_INIT"); + } + EXPECT_EQ(GetActivityOverheadKindString( + static_cast(999)), + "Overhead::UNKNOWN:999"); +} + } // namespace } // namespace test } // namespace profiler diff --git a/third_party/xla/xla/benchmarks/pallas_microbenchmarks/cost_model.py b/third_party/xla/xla/benchmarks/pallas_microbenchmarks/cost_model.py index 01af02c1b86cc8..5ca5b0f2884f4a 100644 --- a/third_party/xla/xla/benchmarks/pallas_microbenchmarks/cost_model.py +++ b/third_party/xla/xla/benchmarks/pallas_microbenchmarks/cost_model.py @@ -25,6 +25,24 @@ from xla.benchmarks.core import platform_info +def vmem_for_operand( + d1: int, + d2: int, + block_d1: int | np.ndarray, + block_d2: int | np.ndarray, + dtype: jax.typing.DTypeLike, +) -> int | np.ndarray: + """Calculates VMEM buffer usage for an operand including double buffering.""" + return ( + # Double-buffer if the operand doesn't fit in the window. + (2 - ((d1 == block_d1) & (d2 == block_d2))) + * block_d1 + * block_d2 + * jax.dtypes.itemsize_bits(dtype) + // 8 + ) + + def _vmem_usage_bytes( m: int, k: int, @@ -61,29 +79,18 @@ def _vmem_usage_bytes( The estimated VMEM usage in bytes, or an array of VMEM usages for all block size combinations. """ - - def _vmem_for_operand(d1, d2, block_d1, block_d2, dtype): - return ( - # Double-buffer if the operand doesn't fit in the window. - (2 - ((d1 == block_d1) & (d2 == block_d2))) - * block_d1 - * block_d2 - * jax.dtypes.itemsize_bits(dtype) - // 8 - ) - - rhs_vmem_usage = _vmem_for_operand(k, n, block_k, block_n, rhs_dtype) + rhs_vmem_usage = vmem_for_operand(k, n, block_k, block_n, rhs_dtype) if sparse_rhs: - rhs_sp_indices_vmem_usage = _vmem_for_operand( + rhs_sp_indices_vmem_usage = vmem_for_operand( k, n, block_k, block_n, jnp.int2 ) rhs_vmem_usage = ( (rhs_vmem_usage + rhs_sp_indices_vmem_usage) * sp_n // sp_m ) return ( - _vmem_for_operand(m, k, block_m, block_k, lhs_dtype) + vmem_for_operand(m, k, block_m, block_k, lhs_dtype) + rhs_vmem_usage - + _vmem_for_operand(m, n, block_m, block_n, out_dtype) + + vmem_for_operand(m, n, block_m, block_n, out_dtype) # Only one accumulator tile is needed. + block_m * block_n * jax.dtypes.itemsize_bits(acc_dtype) // 8 ) diff --git a/third_party/xla/xla/mosaic/dialect/tpu/tpu_ops.cc b/third_party/xla/xla/mosaic/dialect/tpu/tpu_ops.cc index 09432b11e5d080..8688155bcadb36 100644 --- a/third_party/xla/xla/mosaic/dialect/tpu/tpu_ops.cc +++ b/third_party/xla/xla/mosaic/dialect/tpu/tpu_ops.cc @@ -980,67 +980,112 @@ OpFoldResult TransposeOp::fold(FoldAdaptor adaptor) { LogicalResult MemRefBitcastOp::verify() { auto src_ty = getInput().getType(); auto tgt_ty = getType(); - if (tgt_ty.getMemorySpace() != src_ty.getMemorySpace()) { - return emitOpError("Memory spaces do not match."); - } - if (src_ty.getRank() != tgt_ty.getRank()) { - return emitOpError("Ranks do not match."); - } - if (src_ty.getRank() <= 1) { - return emitOpError("Not implemented: 1d memref bitcast."); - } - auto src_bitwidth = getElementTypeBitwidth(src_ty); - auto tgt_bitwidth = getElementTypeBitwidth(tgt_ty); - for (int i = 0; i < src_ty.getRank(); ++i) { - auto src_dim_size = src_ty.getDimSize(i); - auto tgt_dim_size = tgt_ty.getDimSize(i); - if (i == src_ty.getRank() - 2) { - auto src_bits = src_dim_size * src_bitwidth; - auto tgt_bits = tgt_dim_size * tgt_bitwidth; - if (src_bits != tgt_bits) { - return emitOpError( - "Expected the same number of bits on the 2nd minormost " - "dim: (") - << src_dim_size << " * " << src_bitwidth << ") vs (" - << tgt_dim_size << " * " << tgt_bitwidth << ")"; - ; - } - } else { - if (src_dim_size != tgt_dim_size) { - return emitOpError("Expected the same dim size on dim ") - << i << ": " << src_dim_size << " vs " << tgt_dim_size; - } - } - } - // Source and target attributes may be different before propagation is done by - // the canonicalizer, so we allow this when attributes are "unset" in the - // target type. - auto tgt_layout = dyn_cast(tgt_ty.getLayout()); - if (!tgt_layout) { - return success(); - } - auto src_layout = dyn_cast(src_ty.getLayout()); - if (!src_layout) { - return emitOpError("Expected a tiled layout for the input memref."); + FAILUREOR_ASSIGN_OR_RETURN( + auto result_type, inferResultType(src_ty, tgt_ty.getElementType(), + [this]() { return emitOpError(); })); + if (result_type != tgt_ty) { + return emitOpError("Expected result type to be ") << result_type; } - return verifyTiling(); + return success(); } -mlir::InFlightDiagnostic MemRefBitcastOp::verifyTiling() { - auto src_ty = getMemRefType(getInput()); - auto tgt_ty = getType(); - auto src_bitwidth = getElementTypeBitwidth(src_ty); - auto tgt_bitwidth = getElementTypeBitwidth(tgt_ty); - auto src_layout = cast(src_ty.getLayout()); - auto tgt_layout = cast(tgt_ty.getLayout()); - // TODO(jevinjiang): verify memref tiling is valid. Here we just assume the - // source and target tilings are valid. - auto src_tile = src_layout.getTiles().front().dimensions(); - auto tgt_tile = tgt_layout.getTiles().front().dimensions(); - if (src_tile[0] * src_bitwidth != tgt_tile[0] * tgt_bitwidth) { - return emitOpError("Invalid memref bitcast."); +mlir::FailureOr MemRefBitcastOp::inferResultType( + const MemRefType input_type, const Type result_elem_type, + function_ref emit_error) { + const int8_t input_bitwidth = getElementTypeBitwidth(input_type); + const int8_t result_bitwidth = getTypeBitwidth(result_elem_type); + if (input_bitwidth == result_bitwidth) { + return MemRefType( + MemRefType::Builder(input_type).setElementType(result_elem_type)); } - return {}; + const int64_t rank = input_type.getRank(); + if (rank < 2) { + return emit_error() << "Cannot bitcast along 2nd minor in 1D memref"; + } + if (input_type.isDynamicDim(rank - 2)) { + return emit_error() << "Not implemented: Dynamic 2nd minor dimension"; + } + if (input_type.getDimSize(rank - 2) * input_bitwidth % result_bitwidth != 0) { + return emit_error() << "Input 2nd minor dimension bits not a multiple of " + "result bitwidth"; + } + SmallVector result_shape(input_type.getShape()); + result_shape[rank - 2] = + input_type.getDimSize(rank - 2) * input_bitwidth / result_bitwidth; + + if (auto affine_layout = dyn_cast(input_type.getLayout()); + affine_layout && affine_layout.isIdentity()) { + // An affine map layout is interpreted as "unset"/"unknown" (non-standard + // semantics) + return MemRefType(MemRefType::Builder(input_type) + .setShape(result_shape) + .setElementType(result_elem_type)); + } + const auto input_layout = + dyn_cast(input_type.getLayout()); + if (input_layout == nullptr) { + return emit_error() << "Expected a tiled layout for the input memref."; + } + ArrayRef input_tiles = input_layout.getTiles(); + while (!input_tiles.empty() && llvm::all_of(input_tiles.back().dimensions(), + llvm::equal_to(1))) { + input_tiles = input_tiles.drop_back(1); + } + const int64_t num_input_tiles = input_tiles.size(); + for (int64_t i = 0; i < num_input_tiles - 1; ++i) { + if (input_tiles[i].dimensions().size() < + input_tiles[i + 1].dimensions().size()) { + // NOTE: TiledLayoutAttr verification allows tiles that tile across + // previous levels, like T(256)(128)(2, 1), but it enforces that + // previous levels are evenly divided by later levels. + return emit_error() << "Not implemented: Tile at level " << i + 1 + << " tiles across previous tiles."; + } + } + + const bool has_interleaving_tile = + !input_tiles.empty() && input_tiles.back().dimensions().size() >= 2 && + *(input_tiles.back().dimensions().end() - 1) == 1; + const int64_t input_factor = + has_interleaving_tile ? *(input_tiles.back().dimensions().end() - 2) : 1; + if (input_factor * input_bitwidth % result_bitwidth != 0) { + return emit_error() << "Not implemented: Input 2nd minor interleaving tile " + "bits not a multiple of result bitwidth"; + } + const int64_t result_factor = input_factor * input_bitwidth / result_bitwidth; + SmallVector result_tiles; + result_tiles.reserve(input_tiles.size()); + for (const xla::Tile& input_tile : input_tiles) { + const int64_t tile_rank = input_tile.dimensions().size(); + SmallVector tile; + if (tile_rank == 0) { + continue; // Skip + } + if (tile_rank == 1) { + // Recall we checked tiles are of decreasing rank. This is part of a + // suffix of 1D tiles. The input factor must be 1 since the last tile + // is not of the form (..., A, 1). + // The suffix ...(A)(B)(C) becomes ...(1, A)(1, B)(1, C) + CHECK_EQ(input_factor, 1); + CHECK_NE(result_factor, 1); // Identical bitwidths handled above + tile.push_back(1); + } + tile.append(input_tile.dimensions().begin(), input_tile.dimensions().end()); + CHECK_EQ(*(tile.end() - 2) * result_factor % input_factor, 0); + *(tile.end() - 2) = *(tile.end() - 2) * result_factor / input_factor; + // Since tiles are of decreasing rank and higher level tiles divide lower + // level tiles, we can safely remove all-1s tiles. + if (!llvm::all_of(tile, llvm::equal_to(1))) { + result_tiles.emplace_back(tile); + } + } + if (!has_interleaving_tile && result_factor != 1) { + result_tiles.emplace_back(xla::Tile({result_factor, 1})); + } + auto result_layout = tpu::TiledLayoutAttr::get( + input_type.getContext(), result_tiles, input_layout.getTileStrides()); + return MemRefType::get(result_shape, result_elem_type, result_layout, + input_type.getMemorySpace()); } LogicalResult MemRefBitcastOp::canonicalize(MemRefBitcastOp op, diff --git a/third_party/xla/xla/mosaic/dialect/tpu/tpu_ops.td b/third_party/xla/xla/mosaic/dialect/tpu/tpu_ops.td index 1130919108db8c..e31e411588e5c3 100644 --- a/third_party/xla/xla/mosaic/dialect/tpu/tpu_ops.td +++ b/third_party/xla/xla/mosaic/dialect/tpu/tpu_ops.td @@ -1419,7 +1419,9 @@ def TPU_MemRefBitcastOp : TPU_Op<"memref_bitcast", [Pure]> { $input attr-dict `:` type($input) `->` type($result) }]; let extraClassDeclaration = [{ - mlir::InFlightDiagnostic verifyTiling(); + static ::mlir::FailureOr<::mlir::MemRefType> inferResultType( + ::mlir::MemRefType input_type, ::mlir::Type result_elem_type, + ::mlir::function_ref<::mlir::InFlightDiagnostic()> emit_error); }]; let hasVerifier = 1; let hasCanonicalizeMethod = 1; diff --git a/third_party/xla/xla/pjrt/common_pjrt_client.cc b/third_party/xla/xla/pjrt/common_pjrt_client.cc index 6610bf56efdabe..777483f56d717e 100644 --- a/third_party/xla/xla/pjrt/common_pjrt_client.cc +++ b/third_party/xla/xla/pjrt/common_pjrt_client.cc @@ -182,29 +182,31 @@ CommonPjRtClient::AllocateForDelinearizationAsync( "AllocateForDelinearizationAsync is not supported")); } -void CommonPjRtClient::DelinearizeAsync( +static void DelinearizeWhenReady( + CommonPjRtClient* client, tsl::AsyncValueRef staging_buffer, PjRtMemorySpace* memory_space, const Shape& shape, MutableLiteralBase* literal, tsl::Promise promise) { tsl::Context context(tsl::ContextKind::kThread); - staging_buffer.AndThen([this, staging_buffer, shape, literal, + staging_buffer.AndThen([client, staging_buffer, memory_space, shape, literal, context = std::move(context), promise = std::move(promise)]() mutable { if (auto* error = staging_buffer.GetErrorIfPresent()) { promise.Set(*error); return; } - auto run_delinearize = [this, staging_buffer, shape, literal, - context = std::move(context), + auto run_delinearize = [client, staging_buffer, memory_space, shape, + literal, context = std::move(context), promise = std::move(promise)]() mutable { tsl::WithContext wc(context); absl::Span input_data = staging_buffer->const_data(); - absl::Status status = DelinearizeHostBuffer(input_data, shape, literal); + absl::Status status = + client->Delinearize(input_data, shape, literal, memory_space); staging_buffer.reset(); promise.Set(status); }; - if (async_work_runner() != nullptr) { - async_work_runner()->Execute(std::move(run_delinearize)); + if (client->async_work_runner() != nullptr) { + client->async_work_runner()->Execute(std::move(run_delinearize)); } else { run_delinearize(); } @@ -758,9 +760,10 @@ absl::StatusOr CommonPjRtClient::LinearizeIntoImpl( return event.value(); } -absl::Status CommonPjRtClient::DelinearizeHostBuffer( - absl::Span input_data, const Shape& shape, - MutableLiteralBase* literal) { +absl::Status CommonPjRtClient::Delinearize(absl::Span input_data, + const Shape& shape, + MutableLiteralBase* literal, + PjRtMemorySpace* memory_space) { xla::Layout literal_layout; bool need_transpose = false; if (shape.IsArray()) { @@ -3919,9 +3922,9 @@ Future<> CommonPjRtBufferImpl::ToLiteralImpl( return common_client->AllocateForDelinearizationAsync( size, memory_space); }); - common_client->DelinearizeAsync(std::move(staging_buffer), - memory_space, shape, literal, - std::move(promise)); + DelinearizeWhenReady(common_client, std::move(staging_buffer), + memory_space, shape, literal, + std::move(promise)); }; if (literal != nullptr) { diff --git a/third_party/xla/xla/pjrt/common_pjrt_client.h b/third_party/xla/xla/pjrt/common_pjrt_client.h index cfd960f46b0a02..e2a23dcb6ca790 100644 --- a/third_party/xla/xla/pjrt/common_pjrt_client.h +++ b/third_party/xla/xla/pjrt/common_pjrt_client.h @@ -90,10 +90,12 @@ class CommonPjRtClient : public PjRtClient { virtual tsl::AsyncValueRef AllocateForDelinearizationAsync( size_t size, PjRtMemorySpace* memory_space); - virtual void DelinearizeAsync( - tsl::AsyncValueRef staging_buffer, - PjRtMemorySpace* memory_space, const Shape& shape, - MutableLiteralBase* literal, tsl::Promise promise); + // Delinearizes `input_data`, which has the on-device layout of `shape`, into + // `literal`. + virtual absl::Status Delinearize(absl::Span input_data, + const Shape& shape, + MutableLiteralBase* literal, + PjRtMemorySpace* memory_space); // TODO(parkers): Properly support error buffers on GPU and CPU. virtual bool include_raw_buffer_in_ready_event() const { return false; } @@ -590,10 +592,6 @@ class CommonPjRtClient : public PjRtClient { "GetDeviceAddressAlignment is not implemented."); } - absl::Status DelinearizeHostBuffer(absl::Span input_data, - const Shape& shape, - MutableLiteralBase* literal); - // Does the provided shape require runtime shape metadata when being // linearized into the provided memory space? bool RequiresRuntimeShapeMetadata( diff --git a/third_party/xla/xla/pjrt/gpu/BUILD b/third_party/xla/xla/pjrt/gpu/BUILD index 2be58d722890c2..726f16f9c12e10 100644 --- a/third_party/xla/xla/pjrt/gpu/BUILD +++ b/third_party/xla/xla/pjrt/gpu/BUILD @@ -136,6 +136,7 @@ cc_library( "//xla/pjrt/plugin/xla_gpu:xla_gpu_allocator_config", "//xla/pjrt/plugin/xla_gpu:xla_gpu_client_options", "//xla/pjrt/proto:compile_options_proto_cc", + "//xla/pjrt/proto:topology_description_proto_cc", "//xla/pjrt/se:event_pool", "//xla/pjrt/se:local_device_state", "//xla/pjrt/se:pjrt_stream_executor_client", diff --git a/third_party/xla/xla/pjrt/gpu/se_gpu_pjrt_client.cc b/third_party/xla/xla/pjrt/gpu/se_gpu_pjrt_client.cc index a7b8bc56730be8..aa03a6ffb84ef8 100644 --- a/third_party/xla/xla/pjrt/gpu/se_gpu_pjrt_client.cc +++ b/third_party/xla/xla/pjrt/gpu/se_gpu_pjrt_client.cc @@ -97,6 +97,7 @@ limitations under the License. #include "xla/pjrt/pjrt_executable.h" #include "xla/pjrt/plugin/xla_gpu/xla_gpu_allocator_config.h" #include "xla/pjrt/plugin/xla_gpu/xla_gpu_client_options.h" +#include "xla/pjrt/proto/topology_description.pb.h" #include "xla/pjrt/raw_buffer.h" #include "xla/pjrt/se/buffer_sequencing_event.h" #include "xla/pjrt/se/local_device_state.h" @@ -1564,7 +1565,7 @@ CreateAllocatorMemoryRegistration(GpuAllocatorConfig* allocator_config) { struct PjRtDevicesAndTopology { std::vector> devices; - GpuTopologyProto topology; + std::shared_ptr topology; std::vector> local_device_states; }; @@ -1576,6 +1577,7 @@ absl::StatusOr BuildDistributedDevices( std::shared_ptr kv_store, bool enable_mock_nccl, std::optional mock_gpu_topology = std::nullopt, std::optional partition_index = std::nullopt, + bool verify_topology_fingerprint = true, absl::Duration get_local_topology_timeout = absl::Minutes(2), absl::Duration get_global_topology_timeout = absl::Minutes(5)) { std::vector> devices; @@ -1799,6 +1801,43 @@ absl::StatusOr BuildDistributedDevices( gpu_executable_run_options->set_gpu_global_device_ids( std::move(gpu_device_ids)); + ABSL_ASSIGN_OR_RETURN(GpuTopologyProto gpu_topology_proto, + BuildGpuTopology(global_topology, *gpu_target_config, + cpu::TargetMachineOptions())); + ABSL_ASSIGN_OR_RETURN(std::shared_ptr gpu_topology, + GpuTopology::FromProto(gpu_topology_proto)); + std::shared_ptr se_gpu_topology = + CreateSEGpuTopology(platform_name, std::move(gpu_topology), + GetFirstExecutor(local_device_states_vec)); + ABSL_RETURN_IF_ERROR(tsl::ReadBoolFromEnvVar("XLA_PJRT_GPU_VALIDATE_TOPOLOGY", + verify_topology_fingerprint, + &verify_topology_fingerprint)); + if (verify_topology_fingerprint && !enable_mock_nccl && num_nodes > 1) { + constexpr absl::string_view kTopologyFingerprintKey = + "topology_fingerprint"; + ABSL_ASSIGN_OR_RETURN(PjRtTopologyDescriptionProto topology_proto, + se_gpu_topology->ToProto()); + LOG(INFO) << "GPU topology for process " << process_id << ":\n" + << topology_proto; + ABSL_ASSIGN_OR_RETURN(uint64_t topology_fingerprint, + se_gpu_topology->Fingerprint()); + const std::string fingerprint_str = absl::StrCat(topology_fingerprint); + if (process_id == 0) { + ABSL_RETURN_IF_ERROR(kv_store->Set(kTopologyFingerprintKey, fingerprint_str)); + } else { + ABSL_ASSIGN_OR_RETURN( + std::string expected_fingerprint_str, + kv_store->Get(kTopologyFingerprintKey, get_global_topology_timeout)); + if (fingerprint_str != expected_fingerprint_str) { + return absl::FailedPreconditionError(absl::StrFormat( + "Topology fingerprint mismatch between process 0 (%s) and " + "process %d (%s); different hosts may have different GPU " + "topologies or driver versions", + expected_fingerprint_str, process_id, fingerprint_str)); + } + } + } + auto* gpu_collectives = gpu_executable_run_options->collectives(); if (gpu_collectives == nullptr) { gpu_collectives = gpu::GpuCollectives::Resolve(platform_name); @@ -1815,10 +1854,7 @@ absl::StatusOr BuildDistributedDevices( std::move(clique_id_callback)); } - ABSL_ASSIGN_OR_RETURN(GpuTopologyProto gpu_topology, - BuildGpuTopology(global_topology, *gpu_target_config, - cpu::TargetMachineOptions())); - return PjRtDevicesAndTopology{std::move(devices), std::move(gpu_topology), + return PjRtDevicesAndTopology{std::move(devices), std::move(se_gpu_topology), std::move(local_device_states_vec)}; } @@ -2037,14 +2073,10 @@ absl::StatusOr> GetStreamExecutorGpuClient( pjrt_platform_name, std::move(local_device_states), options.node_id, options.num_nodes, gpu_run_options.get(), kv_store, options.enable_mock_nccl, options.mock_gpu_topology, - options.partition_index)); + options.partition_index, options.verify_topology_fingerprint)); - ABSL_ASSIGN_OR_RETURN(std::shared_ptr gpu_topology, - GpuTopology::FromProto(devices_and_topology.topology)); se::StreamExecutor* first_executor = GetFirstExecutor(devices_and_topology.local_device_states); - auto se_gpu_topology = CreateSEGpuTopology( - pjrt_platform_name, std::move(gpu_topology), first_executor); auto raw_client = std::make_unique( tsl::Fingerprint64(pjrt_platform_name), std::move(devices_and_topology.local_device_states), std::move(allocator), @@ -2059,7 +2091,7 @@ absl::StatusOr> GetStreamExecutorGpuClient( return MakeStreamExecutorGpuClient( pjrt_platform_name, std::move(devices_and_topology.devices), options.node_id, std::move(raw_client), std::move(kv_store), - std::move(se_gpu_topology), options.num_nodes); + std::move(devices_and_topology.topology), options.num_nodes); } absl::StatusOr> GetSharedStreamExecutorGpuClient( @@ -2081,10 +2113,12 @@ absl::StatusOr> GetSharedStreamExecutorGpuClient( } ABSL_ASSIGN_OR_RETURN( auto devices_and_topology, - BuildDistributedDevices(platform_name, std::move(local_device_states), - options.node_id, options.num_nodes, - gpu_run_options.get(), kv_store, - /*enable_mock_nccl=*/false)); + BuildDistributedDevices( + platform_name, std::move(local_device_states), options.node_id, + options.num_nodes, gpu_run_options.get(), kv_store, + /*enable_mock_nccl=*/false, /*mock_gpu_topology=*/std::nullopt, + /*partition_index=*/std::nullopt, + options.verify_topology_fingerprint)); VLOG(2) << "Distributed devices built with size=" << devices_and_topology.devices.size(); @@ -2098,13 +2132,8 @@ absl::StatusOr> GetSharedStreamExecutorGpuClient( << "nullptr"; } } - ABSL_ASSIGN_OR_RETURN(auto gpu_topology, - absl::StatusOr>( - GpuTopology::FromProto(devices_and_topology.topology))); se::StreamExecutor* first_executor = GetFirstExecutor(devices_and_topology.local_device_states); - auto se_gpu_topology = CreateSEGpuTopology( - platform_name, std::move(gpu_topology), first_executor); auto raw_client = std::make_unique( tsl::Fingerprint64(platform_name), std::move(devices_and_topology.local_device_states), std::move(allocator), @@ -2119,7 +2148,8 @@ absl::StatusOr> GetSharedStreamExecutorGpuClient( return MakeStreamExecutorGpuClient( platform_name, std::move(devices_and_topology.devices), /*process_index=*/options.node_id, std::move(raw_client), - std::move(kv_store), /*topology=*/std::move(se_gpu_topology), + std::move(kv_store), + /*topology=*/std::move(devices_and_topology.topology), /*num_nodes=*/options.num_nodes); } diff --git a/third_party/xla/xla/pjrt/gpu/se_gpu_pjrt_client_multi_gpu_test.cc b/third_party/xla/xla/pjrt/gpu/se_gpu_pjrt_client_multi_gpu_test.cc index e31a33843a4801..f2f995d81a877e 100644 --- a/third_party/xla/xla/pjrt/gpu/se_gpu_pjrt_client_multi_gpu_test.cc +++ b/third_party/xla/xla/pjrt/gpu/se_gpu_pjrt_client_multi_gpu_test.cc @@ -250,6 +250,94 @@ TEST(StreamExecutorGpuClientTest, CopyDelayedErrorBufferToDevice) { EXPECT_THAT(recv_buffer->ToLiteral().Await(), error); } +TEST(StreamExecutorGpuClientTest, + PropagateAsyncHostToDeviceDelayedErrorCollective) { + ASSERT_OK_AND_ASSIGN(auto client, + GetStreamExecutorGpuClient(GetTestGpuClientOptions(2))); + ASSERT_GE(client->addressable_devices().size(), 2); + + PjRtDevice* d0 = client->addressable_devices()[0]; + PjRtDevice* d1 = client->addressable_devices()[1]; + ASSERT_OK_AND_ASSIGN(PjRtMemorySpace * d0_memory_space, + d0->default_memory_space()); + ASSERT_OK_AND_ASSIGN(PjRtMemorySpace * d1_memory_space, + d1->default_memory_space()); + + static constexpr absl::string_view kAllReduceProgram = R"( +HloModule AllReduce + +sum { + lhs = f32[] parameter(0) + rhs = f32[] parameter(1) + ROOT add = f32[] add(lhs, rhs) +} + +ENTRY main { + p0 = f32[4]{0} parameter(0) + p1 = f32[4]{0} parameter(1) + sum_inputs = f32[4]{0} add(p0, p1) + ROOT ar = f32[4]{0} all-reduce(sum_inputs), replica_groups={{0,1}}, to_apply=sum +} +)"; + + CompileOptions compile_options; + compile_options.executable_build_options.set_num_replicas(2); + compile_options.executable_build_options.mutable_debug_options() + ->set_xla_gpu_executable_terminate_timeout_seconds(5); + ASSERT_OK_AND_ASSIGN( + std::unique_ptr executable, + CompileExecutable(kAllReduceProgram, *client, compile_options)); + + Shape shape = ShapeUtil::MakeShape(F32, {4}); + ASSERT_OK_AND_ASSIGN(auto txm_d0_0, client->CreateBuffersForAsyncHostToDevice( + {shape}, d0_memory_space)); + ASSERT_OK_AND_ASSIGN(auto txm_d0_1, client->CreateBuffersForAsyncHostToDevice( + {shape}, d0_memory_space)); + ASSERT_OK_AND_ASSIGN(auto txm_d1_0, client->CreateBuffersForAsyncHostToDevice( + {shape}, d1_memory_space)); + + std::unique_ptr p0_d0 = txm_d0_0->RetrieveBuffer(0); + std::unique_ptr p1_d0 = txm_d0_1->RetrieveBuffer(0); + std::unique_ptr p0_d1 = txm_d1_0->RetrieveBuffer(0); + ASSERT_OK_AND_ASSIGN(std::unique_ptr p1_d1, + p1_d0->CopyToMemorySpace(d1_memory_space)); + + absl::Status input_error = + absl::UnavailableError("ReadHostBuffer connection timeout"); + std::unique_ptr error_thread(tsl::Env::Default()->StartThread( + tsl::ThreadOptions(), "set_buffer_error", [&]() { + // Wait for both devices' launch_on_device() callbacks to block in + // p0's BufferSequencingEvent::WaitForEventOnStream(). + absl::SleepFor(absl::Milliseconds(100)); + // Poison p0 on d0 first, then p1 (which also poisons p1_d1 via + // CopyToMemorySpace), and then p0 on d1. If IsPredeterminedError() is + // checked before WaitForEventOnStream(), d0 misses both errors and + // launches the collective while d1 sees p1_d1's error and skips + // RunAsync(), deadlocking the collective. + txm_d0_0->SetBufferError(0, input_error); + absl::SleepFor(absl::Milliseconds(50)); + txm_d0_1->SetBufferError(0, input_error); + absl::SleepFor(absl::Milliseconds(50)); + txm_d1_0->SetBufferError(0, input_error); + })); + + std::optional>> returned_futures = + std::vector>(); + ASSERT_OK_AND_ASSIGN(auto results, + executable->Execute({{p0_d0.get(), p1_d0.get()}, + {p0_d1.get(), p1_d1.get()}}, + ExecuteOptions(), returned_futures)); + + ASSERT_EQ(results.size(), 2); + ASSERT_EQ(returned_futures->size(), 2); + for (int i = 0; i < 2; ++i) { + EXPECT_THAT((*returned_futures)[i].Await(), + StatusIs(input_error.code(), HasSubstr(input_error.message()))); + EXPECT_THAT(results[i][0]->GetReadyFuture().Await(), + StatusIs(input_error.code(), HasSubstr(input_error.message()))); + } +} + TEST(StreamExecutorGpuClientTest, DistributedInit) { auto kv_store = std::make_shared(); tsl::thread::ThreadPool thread_pool(tsl::Env::Default(), "DistributeInit", 4); diff --git a/third_party/xla/xla/pjrt/gpu/se_gpu_pjrt_client_test.cc b/third_party/xla/xla/pjrt/gpu/se_gpu_pjrt_client_test.cc index 971c42c5673635..6ea543ba53cb21 100644 --- a/third_party/xla/xla/pjrt/gpu/se_gpu_pjrt_client_test.cc +++ b/third_party/xla/xla/pjrt/gpu/se_gpu_pjrt_client_test.cc @@ -504,6 +504,63 @@ ENTRY %Add.6 (a.1: f32[], b.2: f32[]) -> (f32[], f32[]) { } } +TEST(StreamExecutorGpuClientTest, PropagateAsyncHostToDeviceDelayedError) { + ASSERT_OK_AND_ASSIGN(auto client, + GetStreamExecutorGpuClient(GetTestGpuClientOptions())); + + Shape shape = xla::ShapeUtil::MakeScalarShape(xla::F32); + ASSERT_OK_AND_ASSIGN( + auto* memory_space, + client->addressable_devices()[0]->default_memory_space()); + ASSERT_OK_AND_ASSIGN( + auto transfer_manager, + client->CreateBuffersForAsyncHostToDevice({shape}, memory_space)); + std::unique_ptr buffer = transfer_manager->RetrieveBuffer(0); + + static constexpr char const* kAddProgram = + R"( +HloModule Add.6, entry_computation_layout={(f32[], f32[])->(f32[], f32[])} + +ENTRY %Add.6 (a.1: f32[], b.2: f32[]) -> (f32[], f32[]) { + %a.1 = f32[] parameter(0) + %b.2 = f32[] parameter(1) + %add.3 = f32[] add(f32[] %a.1, f32[] %b.2) + %add.4 = f32[] add(f32[] %add.3, f32[] %add.3) + ROOT %tuple.5 = (f32[], f32[]) tuple(f32[] %add.3, f32[] %add.4) +} +)"; + ASSERT_OK_AND_ASSIGN(auto executable, + CompileExecutable(kAddProgram, *client)); + + absl::Status input_error = + absl::UnavailableError("ReadHostBuffer connection timeout"); + std::unique_ptr error_thread(tsl::Env::Default()->StartThread( + tsl::ThreadOptions(), "set_buffer_error", [&]() { + // Allow Execute() to enter + // launch_on_device() and block in + // BufferSequencingEvent::WaitForEventOnStream() + // before poisoning the transfer. + absl::SleepFor(absl::Milliseconds(100)); + transfer_manager->SetBufferError(0, input_error); + })); + + std::optional>> returned_futures = + std::vector>(); + ASSERT_OK_AND_ASSIGN(auto result, + executable->Execute({{buffer.get(), buffer.get()}}, + /*options=*/{}, returned_futures)); + + ASSERT_EQ(result.size(), 1); + ASSERT_EQ(result[0].size(), 2); + ASSERT_EQ(returned_futures->size(), 1); + EXPECT_THAT((*returned_futures)[0].Await(), + StatusIs(input_error.code(), HasSubstr(input_error.message()))); + for (const auto& b : result[0]) { + EXPECT_THAT(b->GetReadyFuture().Await(), + StatusIs(input_error.code(), HasSubstr(input_error.message()))); + } +} + TEST(StreamExecutorGpuClientTest, DonateWithControlDependency) { TF_ASSERT_OK_AND_ASSIGN( auto client, GetStreamExecutorGpuClient(GetTestGpuClientOptions())); diff --git a/third_party/xla/xla/pjrt/pjrt_executable.cc b/third_party/xla/xla/pjrt/pjrt_executable.cc index 0a341f34c2019d..093aee78de76ec 100644 --- a/third_party/xla/xla/pjrt/pjrt_executable.cc +++ b/third_party/xla/xla/pjrt/pjrt_executable.cc @@ -703,7 +703,7 @@ absl::Status CompileOptions::ApplyOption(const std::string& key, case tsl::protobuf::FieldDescriptor::TYPE_FLOAT: { if (std::holds_alternative(value)) { double double_value = std::get(value); - if (double_value >= std::numeric_limits::min() && + if (double_value >= std::numeric_limits::lowest() && double_value <= std::numeric_limits::max()) { return ApplyFloatOption(xla_field, static_cast(double_value), debug_options); diff --git a/third_party/xla/xla/pjrt/plugin/xla_gpu/xla_gpu_client_options.h b/third_party/xla/xla/pjrt/plugin/xla_gpu/xla_gpu_client_options.h index 4fdd53bb691eb8..24491b5f44edbc 100644 --- a/third_party/xla/xla/pjrt/plugin/xla_gpu/xla_gpu_client_options.h +++ b/third_party/xla/xla/pjrt/plugin/xla_gpu/xla_gpu_client_options.h @@ -67,6 +67,8 @@ struct GpuClientOptions { std::optional use_async_dispatch; std::optional max_inflight_computations = 32; + + bool verify_topology_fingerprint = true; }; } // namespace xla diff --git a/third_party/xla/xla/pjrt/se/pjrt_stream_executor_client.cc b/third_party/xla/xla/pjrt/se/pjrt_stream_executor_client.cc index 036dbe82c319f0..8297a0e76a82d3 100644 --- a/third_party/xla/xla/pjrt/se/pjrt_stream_executor_client.cc +++ b/third_party/xla/xla/pjrt/se/pjrt_stream_executor_client.cc @@ -1502,12 +1502,12 @@ PjRtStreamExecutorRawLoadedExecutable::Execute( for (size_t i = 0; i < extra_deps.size(); ++i) { const auto& event = extra_deps[i]; if (auto ev = event.down_cast()) { + ev->WaitForEventOnStream(device_state->compute_stream()); if (ev->IsPredeterminedError()) { if (predetermined_error.ok()) { predetermined_error = ev->GetDefinedStatus(); } } - ev->WaitForEventOnStream(device_state->compute_stream()); } else if (event) { xla::BlockUntilReady(event); if (auto error = event.GetErrorIfPresent()) { diff --git a/third_party/xla/xla/service/gpu/BUILD b/third_party/xla/xla/service/gpu/BUILD index f93d35c2da47d0..d2271f05752e61 100644 --- a/third_party/xla/xla/service/gpu/BUILD +++ b/third_party/xla/xla/service/gpu/BUILD @@ -1729,6 +1729,7 @@ cc_library( "//xla/hlo/pass:hlo_pass_pipeline", "//xla/hlo/transforms/simplifiers:hlo_dce", "//xla/service:hlo_cost_analysis", + "//xla/service/gpu/model:gpu_indexing_performance_model", "//xla/stream_executor:device_description", "@llvm-project//mlir:IR", ], diff --git a/third_party/xla/xla/service/gpu/autotuning/BUILD b/third_party/xla/xla/service/gpu/autotuning/BUILD index 32594035905f8e..7ac210fd3058cf 100644 --- a/third_party/xla/xla/service/gpu/autotuning/BUILD +++ b/third_party/xla/xla/service/gpu/autotuning/BUILD @@ -308,6 +308,7 @@ cc_library( "//xla/service/gpu:backend_configs_cc", "//xla/service/gpu:cublas_cudnn", "//xla/service/gpu:ir_emission_utils", + "//xla/service/gpu/model:gpu_indexing_performance_model", "//xla/stream_executor:device_address_allocator", "//xla/stream_executor:device_description", "//xla/stream_executor:platform_id", diff --git a/third_party/xla/xla/service/gpu/autotuning/config_assigner_pass.cc b/third_party/xla/xla/service/gpu/autotuning/config_assigner_pass.cc index 67c4d9f647cade..33973b180aebe8 100644 --- a/third_party/xla/xla/service/gpu/autotuning/config_assigner_pass.cc +++ b/third_party/xla/xla/service/gpu/autotuning/config_assigner_pass.cc @@ -58,6 +58,7 @@ limitations under the License. #include "xla/service/gpu/backend_configs.pb.h" #include "xla/service/gpu/cublas_cudnn.h" #include "xla/service/gpu/ir_emission_utils.h" +#include "xla/service/gpu/model/gpu_indexing_performance_model.h" #include "xla/service/hlo_cost_analysis.h" #include "xla/stream_executor/device_address_allocator.h" #include "xla/stream_executor/device_description.h" @@ -393,7 +394,8 @@ ConfigAssignerPass::GetEnabledBackends( const Compiler::GpuTargetConfig* target_config, const AliasInfo* alias_info, const DebugOptions& debug_options, mlir::MLIRContext* mlir_context, HloCostAnalysis::ShapeSizeFunction shape_size_fn, Compiler* compiler, - se::PlatformId platform_id) { + se::PlatformId platform_id, tsl::thread::ThreadPool* thread_pool, + MlirContextPool* mlir_context_pool) { std::vector autotune_backends; for (const auto& backend : debug_options.xla_gpu_experimental_autotune_backends()) { @@ -436,7 +438,8 @@ ConfigAssignerPass::GetEnabledBackends( registry.FindObject(platform_id)); std::vector> backends = get_codegen_backends( stream_exec, device_allocator, &debug_options, compiler, target_config, - alias_info, mlir_context, shape_size_fn, autotune_backends); + alias_info, mlir_context, shape_size_fn, autotune_backends, thread_pool, + mlir_context_pool); return backends; } diff --git a/third_party/xla/xla/service/gpu/autotuning/config_assigner_pass.h b/third_party/xla/xla/service/gpu/autotuning/config_assigner_pass.h index dc049ccc00636d..28d4741177e18f 100644 --- a/third_party/xla/xla/service/gpu/autotuning/config_assigner_pass.h +++ b/third_party/xla/xla/service/gpu/autotuning/config_assigner_pass.h @@ -38,6 +38,7 @@ limitations under the License. #include "xla/hlo/pass/hlo_pass_interface.h" #include "xla/pjrt/distributed/key_value_store_interface.h" #include "xla/service/compiler.h" +#include "xla/service/gpu/model/gpu_indexing_performance_model.h" #include "xla/service/hlo_cost_analysis.h" #include "xla/stream_executor/device_address_allocator.h" #include "xla/stream_executor/device_description.h" @@ -86,7 +87,9 @@ class ConfigAssignerPass : public HloModulePass { const DebugOptions& debug_options, mlir::MLIRContext* mlir_context, HloCostAnalysis::ShapeSizeFunction shape_size_fn, - Compiler* compiler, se::PlatformId platform_id); + Compiler* compiler, se::PlatformId platform_id, + tsl::thread::ThreadPool* thread_pool = nullptr, + MlirContextPool* mlir_context_pool = nullptr); // Note: the target_config must outlive the pass. static absl::StatusOr> Create( diff --git a/third_party/xla/xla/service/gpu/fusion_dispatch_pipeline.cc b/third_party/xla/xla/service/gpu/fusion_dispatch_pipeline.cc index 558f59a223609b..a337411c2b0d02 100644 --- a/third_party/xla/xla/service/gpu/fusion_dispatch_pipeline.cc +++ b/third_party/xla/xla/service/gpu/fusion_dispatch_pipeline.cc @@ -19,6 +19,7 @@ limitations under the License. #include "xla/backends/gpu/transforms/fusion_block_level_rewriter.h" #include "xla/hlo/pass/hlo_pass_pipeline.h" #include "xla/hlo/transforms/simplifiers/hlo_dce.h" +#include "xla/service/gpu/model/gpu_indexing_performance_model.h" #include "xla/service/hlo_cost_analysis.h" #include "xla/stream_executor/device_description.h" #include "xla/xla.pb.h" @@ -29,11 +30,13 @@ namespace gpu { HloPassPipeline FusionDispatchPipeline( const se::DeviceDescription& device_description, HloCostAnalysis::ShapeSizeFunction shape_size_fn, - mlir::MLIRContext* mlir_context) { + mlir::MLIRContext* mlir_context, tsl::thread::ThreadPool* thread_pool, + MlirContextPool* mlir_context_pool) { HloPassPipeline pipeline("fusion-dispatch-pipeline"); pipeline.AddPass(); pipeline.AddPass(device_description, shape_size_fn, - mlir_context); + mlir_context, thread_pool, + mlir_context_pool); return pipeline; } diff --git a/third_party/xla/xla/service/gpu/fusion_dispatch_pipeline.h b/third_party/xla/xla/service/gpu/fusion_dispatch_pipeline.h index 2fae2a309bf1ac..6c5f4d82641d3c 100644 --- a/third_party/xla/xla/service/gpu/fusion_dispatch_pipeline.h +++ b/third_party/xla/xla/service/gpu/fusion_dispatch_pipeline.h @@ -18,10 +18,15 @@ limitations under the License. #include "mlir/IR/MLIRContext.h" #include "xla/hlo/pass/hlo_pass_pipeline.h" +#include "xla/service/gpu/model/gpu_indexing_performance_model.h" #include "xla/service/hlo_cost_analysis.h" #include "xla/stream_executor/device_description.h" #include "xla/xla.pb.h" +namespace tsl::thread { +class ThreadPool; +} // namespace tsl::thread + namespace xla { namespace gpu { @@ -30,7 +35,9 @@ namespace gpu { HloPassPipeline FusionDispatchPipeline( const se::DeviceDescription& device_description, HloCostAnalysis::ShapeSizeFunction shape_size_fn, - mlir::MLIRContext* mlir_context); + mlir::MLIRContext* mlir_context, + tsl::thread::ThreadPool* thread_pool = nullptr, + MlirContextPool* mlir_context_pool = nullptr); } // namespace gpu } // namespace xla diff --git a/third_party/xla/xla/service/gpu/gpu_compiler.cc b/third_party/xla/xla/service/gpu/gpu_compiler.cc index 6a828612bf9bb7..21b539b21b6e10 100644 --- a/third_party/xla/xla/service/gpu/gpu_compiler.cc +++ b/third_party/xla/xla/service/gpu/gpu_compiler.cc @@ -2884,9 +2884,6 @@ GpuCompiler::CompileToBackendResult( HloPassPipeline pipeline("scheduled-gpu-module"); AddHloVerifier(&pipeline); ABSL_RETURN_IF_ERROR(pipeline.Run(module).status()); - ABSL_RETURN_IF_ERROR( - RunPostSchedulingPipelines(module, schedule_metadata.scheduler_mem_limit, - gpu_topology, alias_info.get(), mlir_context)); MaybeOwningThreadPool thread_pool = CreateMaybeOwningThreadPool( /*parallelism=*/module->config() @@ -2895,6 +2892,10 @@ GpuCompiler::CompileToBackendResult( /*default_thread_pool=*/options.thread_pool, /*default_parallelism=*/tsl::port::MaxParallelism()); + ABSL_RETURN_IF_ERROR(RunPostSchedulingPipelines( + module, schedule_metadata.scheduler_mem_limit, gpu_topology, + alias_info.get(), mlir_context, thread_pool.get_mutable())); + absl::Mutex module_stats_m_; ModuleStats module_stats; CompileModuleResults compile_module_results; @@ -3341,7 +3342,7 @@ HloRematerialization::Options CreateRematOpts( absl::Status GpuCompiler::RunPostSchedulingPipelines( HloModule* module, int64_t scheduler_mem_limit, const GpuTopology& gpu_topology, const GpuAliasInfo* alias_info, - mlir::MLIRContext* mlir_context) { + mlir::MLIRContext* mlir_context, tsl::thread::ThreadPool* thread_pool) { tsl::profiler::TraceMe traceme("RunPostSchedulingPipelines"); ABSL_RETURN_IF_ERROR( RunPostSchedulingCopyInsertion(module, &gpu_topology, alias_info)); @@ -3405,8 +3406,9 @@ absl::Status GpuCompiler::RunPostSchedulingPipelines( if (cuda_cc != nullptr && cuda_cc->IsAtLeastAmpere()) { // This needs to run after every pass affecting fusions. The last passes // that create new fusions are FusionWrapper and StreamAttributeAnnotator. - main_pipeline.AddPass(FusionDispatchPipeline( - gpu_device_info, ShapeSizeBytesFunction(), mlir_context)); + main_pipeline.AddPass( + FusionDispatchPipeline(gpu_device_info, ShapeSizeBytesFunction(), + mlir_context, thread_pool, &mlir_context_pool_)); } // Sanitize constant names. This is in its own pipeline to ensure it always @@ -3531,7 +3533,8 @@ absl::Status GpuCompiler::AddConfigAssignerPass( [&]() -> absl::StatusOr>> { return ConfigAssignerPass::GetEnabledBackends( stream_exec, options.device_allocator, target_config, alias_info, - debug_options, mlir_context, shape_size_fn, this, PlatformId()); + debug_options, mlir_context, shape_size_fn, this, PlatformId(), + thread_pool, &mlir_context_pool_); }; ABSL_ASSIGN_OR_RETURN( diff --git a/third_party/xla/xla/service/gpu/gpu_compiler.h b/third_party/xla/xla/service/gpu/gpu_compiler.h index 53dc27069edb23..169dbadc58e379 100644 --- a/third_party/xla/xla/service/gpu/gpu_compiler.h +++ b/third_party/xla/xla/service/gpu/gpu_compiler.h @@ -103,11 +103,11 @@ class GpuCompiler : public LLVMCompiler { absl::StatusOr> Export( Executable* executable) override; - absl::Status RunPostSchedulingPipelines(HloModule* module, - int64_t scheduler_mem_limit, - const GpuTopology& gpu_topology, - const GpuAliasInfo* alias_info, - mlir::MLIRContext* mlir_context); + absl::Status RunPostSchedulingPipelines( + HloModule* module, int64_t scheduler_mem_limit, + const GpuTopology& gpu_topology, const GpuAliasInfo* alias_info, + mlir::MLIRContext* mlir_context, + tsl::thread::ThreadPool* thread_pool = nullptr); std::string target_triple() const { return target_triple_; } std::string data_layout() const { return data_layout_; } diff --git a/third_party/xla/xla/service/gpu_topology.cc b/third_party/xla/xla/service/gpu_topology.cc index bf83bb006dc446..68a1c58c02c2d5 100644 --- a/third_party/xla/xla/service/gpu_topology.cc +++ b/third_party/xla/xla/service/gpu_topology.cc @@ -82,7 +82,7 @@ GetHostTargetMachineOptions(absl::string_view platform_version) { } if (platform_version == "oberon_b200" || platform_version == "oberon_b300") { return cpu::TargetMachineOptions{ - "aarch64-linux-gnu", "neoverse-n1", + "aarch64-unknown-linux-gnu", "neoverse-n1", "+aes,+crc,+fp-armv8,+lse,+neon,+sha2,+sha3,+sm4,+sve-aes,+sve-sha3,+" "sve-sm4,-rand,-sve,-sve2"}; } diff --git a/third_party/xla/xla/service/gpu_topology_test.cc b/third_party/xla/xla/service/gpu_topology_test.cc index ef04f68704831c..accdbfec7f23ba 100644 --- a/third_party/xla/xla/service/gpu_topology_test.cc +++ b/third_party/xla/xla/service/gpu_topology_test.cc @@ -161,7 +161,7 @@ TEST(GpuTopologyTest, GetGpuTopologyForPlatformOberonB200) { EXPECT_TRUE(topology.has_gpu_target_config()); EXPECT_THAT(topology.host_target_machine_options(), Optional(Property(&cpu::TargetMachineOptions::triple, - "aarch64-linux-gnu"))); + "aarch64-unknown-linux-gnu"))); } TEST(GpuTopologyTest, GetGpuTopologyForPlatformOberonB300) { @@ -175,7 +175,7 @@ TEST(GpuTopologyTest, GetGpuTopologyForPlatformOberonB300) { EXPECT_TRUE(topology.has_gpu_target_config()); EXPECT_THAT(topology.host_target_machine_options(), Optional(Property(&cpu::TargetMachineOptions::triple, - "aarch64-linux-gnu"))); + "aarch64-unknown-linux-gnu"))); } TEST(GpuTopologyTest, GetGpuTopologyForPlatformInvalid) { diff --git a/third_party/xla/xla/service/layout_assignment.h b/third_party/xla/xla/service/layout_assignment.h index 934b35780a8590..a431c1665f89a9 100644 --- a/third_party/xla/xla/service/layout_assignment.h +++ b/third_party/xla/xla/service/layout_assignment.h @@ -1038,11 +1038,10 @@ class LayoutAssignment : public HloModulePass { std::string ToString(const LayoutConstraints& constraints) const; int64_t current_priority() const { return current_priority_; } - - private: // Returns whether the given instruction is in a copy-disabled while loop. bool IsWhileLoopCopyDisabled(const HloInstruction& instruction) const; + private: // Map containing the layouts of all computations assigned so // far. Computations are handled in a topological sort where computations are // handled before their caller instructions so the layouts of caller diff --git a/third_party/xla/xla/service/while_loop_simplifier.cc b/third_party/xla/xla/service/while_loop_simplifier.cc index 3e99ac8dcbc57d..2c420839aa4e7c 100644 --- a/third_party/xla/xla/service/while_loop_simplifier.cc +++ b/third_party/xla/xla/service/while_loop_simplifier.cc @@ -92,19 +92,21 @@ static absl::StatusOr TryRemoveTrivialCompare(HloInstruction* while_op) { std::optional constant_value = LiteralUtil::LiteralAsScalarInt64(constant->literal()); if (constant_value.has_value()) { - // x <= c && i >= c --> i > x - if (constant_value.value() <= init_value.value()) { - if (body_instr->comparison_direction() == - ComparisonDirection::kLt) { - ABSL_RETURN_IF_ERROR(while_op->while_body()->ReplaceInstruction( - body_instr, MakeScalarLike(body_instr, false))); - return true; - } else if (body_instr->comparison_direction() == - ComparisonDirection::kGt) { - ABSL_RETURN_IF_ERROR(while_op->while_body()->ReplaceInstruction( - body_instr, MakeScalarLike(body_instr, true))); - return true; - } + // x <= c && i >= c --> !(i < x) + // x < c && i >= c --> i > x + if (constant_value.value() <= init_value.value() && + body_instr->comparison_direction() == + ComparisonDirection::kLt) { + ABSL_RETURN_IF_ERROR(while_op->while_body()->ReplaceInstruction( + body_instr, MakeScalarLike(body_instr, false))); + return true; + } + if (constant_value.value() < init_value.value() && + body_instr->comparison_direction() == + ComparisonDirection::kGt) { + ABSL_RETURN_IF_ERROR(while_op->while_body()->ReplaceInstruction( + body_instr, MakeScalarLike(body_instr, true))); + return true; } // x >= c + k && i < c + k --> i < x if (constant_value.value() >= diff --git a/third_party/xla/xla/service/while_loop_simplifier_test.cc b/third_party/xla/xla/service/while_loop_simplifier_test.cc index dc348d5027c9c4..4a58124d6d370b 100644 --- a/third_party/xla/xla/service/while_loop_simplifier_test.cc +++ b/third_party/xla/xla/service/while_loop_simplifier_test.cc @@ -1347,7 +1347,7 @@ TEST_F(WhileLoopSimplifierTest, RemoveTrivialCompare) { )"; for (std::string dir : {"LT", "GT"}) { - for (int i = 1; i > -5; i--) { + for (int i = (dir == "LT" ? 1 : 0); i > -5; i--) { std::string hlo_string = absl::StrReplaceAll( hlo_template, {{"{{LOOP_CONSTANT}}", absl::StrCat(i)}, {"{{DIRECTION}}", dir}}); @@ -1366,6 +1366,15 @@ TEST_F(WhileLoopSimplifierTest, RemoveTrivialCompare) { .IsAll(dir == "GT")); } + if (dir == "GT") { + std::string hlo_string = absl::StrReplaceAll( + hlo_template, {{"{{LOOP_CONSTANT}}", "1"}, {"{{DIRECTION}}", dir}}); + auto m = ParseAndReturnVerifiedModule(hlo_string).value(); + EXPECT_FALSE(WhileLoopSimplifier(/*simplify_compare_instrs=*/true) + .Run(m.get()) + .value()); + } + for (int i = 11; i < 15; i++) { std::string hlo_string = absl::StrReplaceAll( hlo_template, diff --git a/third_party/xla/xla/stream_executor/cuda/topk_kernel_cuda_common.cu.h b/third_party/xla/xla/stream_executor/cuda/topk_kernel_cuda_common.cu.h index 3bb05e6822e6dc..116d2dbf3f4ef2 100644 --- a/third_party/xla/xla/stream_executor/cuda/topk_kernel_cuda_common.cu.h +++ b/third_party/xla/xla/stream_executor/cuda/topk_kernel_cuda_common.cu.h @@ -30,6 +30,7 @@ limitations under the License. #include "xla/stream_executor/gpu/gpu_kernel_registry.h" #include "xla/stream_executor/gpu/topk_kernel.h" #include "xla/tsl/lib/math/math_util.h" +#include "xla/types.h" #define WAVEFRONT_SIZE 32 @@ -68,33 +69,63 @@ __device__ __forceinline__ NT GpuShuffle(NT val, uint32_t idx, // ordering, properly handling special values such as NaNs and signed zeroes // during integer sorting. namespace details { + template -__device__ __forceinline__ auto ToOrdered(T x) { - if constexpr (sizeof(T) == 4 && !std::is_integral_v) { +struct OrderedTraits { + using Type = T; + static __device__ __forceinline__ Type ToOrdered(T x) { return x; } + static __device__ __forceinline__ T FromOrdered(Type val) { return val; } +}; + +template <> +struct OrderedTraits { + using Type = uint32_t; + + static __device__ __forceinline__ Type ToOrdered(float x) { uint32_t val = absl::bit_cast(x); return (val & 0x80000000u) ? ~val : (val | 0x80000000u); - } else if constexpr (sizeof(T) == 2 && !std::is_integral_v) { + } + + static __device__ __forceinline__ float FromOrdered(Type val) { + uint32_t u = (val & 0x80000000u) ? (val ^ 0x80000000u) : ~val; + return absl::bit_cast(u); + } +}; + +template <> +struct OrderedTraits { + using Type = uint16_t; + + static __device__ __forceinline__ Type ToOrdered(xla::bfloat16 x) { uint16_t val = absl::bit_cast(x); return (val & 0x8000u) ? static_cast(~val) : static_cast(val | 0x8000u); - } else { - return x; } -} -template -__device__ __forceinline__ T FromOrdered(OrderedT val) { - if constexpr (sizeof(T) == 4 && !std::is_integral_v) { - uint32_t u = (val & 0x80000000u) ? (val ^ 0x80000000u) : ~val; - return absl::bit_cast(u); - } else if constexpr (sizeof(T) == 2 && !std::is_integral_v) { + static __device__ __forceinline__ xla::bfloat16 FromOrdered(Type val) { uint16_t u = (val & 0x8000u) ? static_cast(val ^ 0x8000u) : static_cast(~val); - return absl::bit_cast(u); - } else { - return val; + return absl::bit_cast(u); } -} +}; + +template <> +struct OrderedTraits { + using Type = uint16_t; + + static __device__ __forceinline__ Type ToOrdered(xla::half x) { + uint16_t val = absl::bit_cast(x); + return (val & 0x8000u) ? static_cast(~val) + : static_cast(val | 0x8000u); + } + + static __device__ __forceinline__ xla::half FromOrdered(Type val) { + uint16_t u = (val & 0x8000u) ? static_cast(val ^ 0x8000u) + : static_cast(~val); + return absl::bit_cast(u); + } +}; + } // namespace details // Default implementation for KV holder. Useful for testing while adding support @@ -102,7 +133,7 @@ __device__ __forceinline__ T FromOrdered(OrderedT val) { // implementations below. template struct Descending { - using OrderedKey = decltype(details::ToOrdered(T{})); + using OrderedKey = typename details::OrderedTraits::Type; struct KVT { OrderedKey key; V idx; @@ -202,7 +233,7 @@ struct TopK { // TODO(doak): Use bitonic sort. #pragma unroll for (int i = 0; i < K; i++) { - tmp[i] = {details::ToOrdered(key[Idx(i)]), VT(Idx(i))}; + tmp[i] = {details::OrderedTraits::ToOrdered(key[Idx(i)]), VT(Idx(i))}; } #pragma unroll for (int i = 0; i < K; i++) { @@ -218,7 +249,8 @@ struct TopK { constexpr uint32_t WarpSize = WAVEFRONT_SIZE; for (int idx = K; idx < n; idx++) { - KVT kv{details::ToOrdered(key[Idx(idx)]), VT(Idx(idx))}; + KVT kv{details::OrderedTraits::ToOrdered(key[Idx(idx)]), + VT(Idx(idx))}; Push(tmp, kv); } Reduce(tmp, WarpSize); @@ -252,7 +284,7 @@ struct TopK { return; } for (int i = 0; i < num_outputs_; ++i) { - keys[i] = details::FromOrdered(tmp[i].key); + keys[i] = details::OrderedTraits::FromOrdered(tmp[i].key); idxs[i] = tmp[i].idx; } } diff --git a/third_party/xla/xla/stream_executor/rocm/topk_kernel_rocm_common.cu.h b/third_party/xla/xla/stream_executor/rocm/topk_kernel_rocm_common.cu.h index 6742958e8d85c8..240942a73befa9 100644 --- a/third_party/xla/xla/stream_executor/rocm/topk_kernel_rocm_common.cu.h +++ b/third_party/xla/xla/stream_executor/rocm/topk_kernel_rocm_common.cu.h @@ -30,6 +30,7 @@ limitations under the License. #include "xla/stream_executor/kernel_symbol_registry.h" #include "xla/stream_executor/rocm/rocm_platform_id.h" #include "xla/tsl/lib/math/math_util.h" +#include "xla/types.h" // https://rocm.docs.amd.com/en/latest/about/release-notes.html#amdgpu-wavefront-size-compiler-macro-deprecation #if defined(__GFX9__) @@ -72,33 +73,63 @@ __device__ __forceinline__ NT GpuShuffle(NT val, uint32_t idx, // well-defined total ordering, properly handling special values such as NaNs // and signed zeroes during integer sorting. namespace details { + template -__device__ __forceinline__ auto ToOrdered(T x) { - if constexpr (sizeof(T) == 4 && !std::is_integral_v) { +struct OrderedTraits { + using Type = T; + static __device__ __forceinline__ Type ToOrdered(T x) { return x; } + static __device__ __forceinline__ T FromOrdered(Type val) { return val; } +}; + +template <> +struct OrderedTraits { + using Type = uint32_t; + + static __device__ __forceinline__ Type ToOrdered(float x) { uint32_t val = absl::bit_cast(x); return (val & 0x80000000u) ? ~val : (val | 0x80000000u); - } else if constexpr (sizeof(T) == 2 && !std::is_integral_v) { + } + + static __device__ __forceinline__ float FromOrdered(Type val) { + uint32_t u = (val & 0x80000000u) ? (val ^ 0x80000000u) : ~val; + return absl::bit_cast(u); + } +}; + +template <> +struct OrderedTraits { + using Type = uint16_t; + + static __device__ __forceinline__ Type ToOrdered(xla::bfloat16 x) { uint16_t val = absl::bit_cast(x); return (val & 0x8000u) ? static_cast(~val) : static_cast(val | 0x8000u); - } else { - return x; } -} -template -__device__ __forceinline__ T FromOrdered(OrderedT val) { - if constexpr (sizeof(T) == 4 && !std::is_integral_v) { - uint32_t u = (val & 0x80000000u) ? (val ^ 0x80000000u) : ~val; - return absl::bit_cast(u); - } else if constexpr (sizeof(T) == 2 && !std::is_integral_v) { + static __device__ __forceinline__ xla::bfloat16 FromOrdered(Type val) { uint16_t u = (val & 0x8000u) ? static_cast(val ^ 0x8000u) : static_cast(~val); - return absl::bit_cast(u); - } else { - return val; + return absl::bit_cast(u); } -} +}; + +template <> +struct OrderedTraits { + using Type = uint16_t; + + static __device__ __forceinline__ Type ToOrdered(xla::half x) { + uint16_t val = absl::bit_cast(x); + return (val & 0x8000u) ? static_cast(~val) + : static_cast(val | 0x8000u); + } + + static __device__ __forceinline__ xla::half FromOrdered(Type val) { + uint16_t u = (val & 0x8000u) ? static_cast(val ^ 0x8000u) + : static_cast(~val); + return absl::bit_cast(u); + } +}; + } // namespace details // Default implementation for KV holder. Useful for testing while adding support @@ -106,7 +137,7 @@ __device__ __forceinline__ T FromOrdered(OrderedT val) { // implementations below. template struct Descending { - using OrderedKey = decltype(details::ToOrdered(T{})); + using OrderedKey = typename details::OrderedTraits::Type; struct KVT { OrderedKey key; V idx; @@ -206,7 +237,7 @@ struct TopK { // TODO(doak): Use bitonic sort. #pragma unroll for (int i = 0; i < K; i++) { - tmp[i] = {details::ToOrdered(key[Idx(i)]), VT(Idx(i))}; + tmp[i] = {details::OrderedTraits::ToOrdered(key[Idx(i)]), VT(Idx(i))}; } #pragma unroll for (int i = 0; i < K; i++) { @@ -222,7 +253,8 @@ struct TopK { constexpr uint32_t WarpSize = WAVEFRONT_SIZE; for (int idx = K; idx < n; idx++) { - KVT kv{details::ToOrdered(key[Idx(idx)]), VT(Idx(idx))}; + KVT kv{details::OrderedTraits::ToOrdered(key[Idx(idx)]), + VT(Idx(idx))}; Push(tmp, kv); } Reduce(tmp, WarpSize); @@ -250,7 +282,7 @@ struct TopK { Reduce(tmp, blockDim.x / WarpSize); if (threadIdx.x != 0) return; for (int i = 0; i < num_outputs_; ++i) { - keys[i] = details::FromOrdered(tmp[i].key); + keys[i] = details::OrderedTraits::FromOrdered(tmp[i].key); idxs[i] = tmp[i].idx; } } diff --git a/third_party/xla/xla/tools/benchmarks/benchmark_registry.pbtxt b/third_party/xla/xla/tools/benchmarks/benchmark_registry.pbtxt index 9a35bc7726329d..3e49eba395d705 100644 --- a/third_party/xla/xla/tools/benchmarks/benchmark_registry.pbtxt +++ b/third_party/xla/xla/tools/benchmarks/benchmark_registry.pbtxt @@ -20,7 +20,7 @@ benchmarks { environment_configs { id: "gpu_l4" - runner_label: "linux-x86-g2-16-l4-1gpu" + runner_label: "linux-x86-16cpu-l4-1gpu" container_image: "us-docker.pkg.dev/ml-oss-artifacts-published/ml-public-container/ml-build-cuda13.2-cudnn9.15:latest" workload_action_inputs { key: "hardware_category" @@ -69,7 +69,7 @@ benchmarks { environment_configs { id: "cpu_x86" - runner_label: "linux-x86-n2-128" + runner_label: "linux-x86-n4-80" container_image: "us-docker.pkg.dev/ml-oss-artifacts-published/ml-public-container/ml-build:latest" workload_action_inputs { key: "hardware_category" @@ -124,7 +124,7 @@ benchmarks { environment_configs { id: "gpu_l4" - runner_label: "linux-x86-g2-16-l4-1gpu" + runner_label: "linux-x86-16cpu-l4-1gpu" container_image: "us-docker.pkg.dev/ml-oss-artifacts-published/ml-public-container/ml-build-cuda13.2-cudnn9.15:latest" workload_action_inputs { key: "hardware_category" @@ -208,7 +208,7 @@ benchmarks { environment_configs { id: "cpu_x86" - runner_label: "linux-x86-n2-128" + runner_label: "linux-x86-n4-80" container_image: "us-docker.pkg.dev/ml-oss-artifacts-published/ml-public-container/ml-build:latest" workload_action_inputs { key: "hardware_category" @@ -284,7 +284,7 @@ benchmarks { environment_configs { id: "cpu_x86" - runner_label: "linux-x86-n2-128" + runner_label: "linux-x86-n4-80" container_image: "us-docker.pkg.dev/ml-oss-artifacts-published/ml-public-container/ml-build:latest" workload_action_inputs { key: "hardware_category"