From 3152ae31a547b7f49774b56c552a2cc238fc97eb Mon Sep 17 00:00:00 2001 From: Junwhan Ahn Date: Tue, 29 Sep 2026 11:41:11 -0700 Subject: [PATCH 01/24] [PJRT:GPU] Fail distributed initialization when hosts have different topology fingerprints. When different hosts in a multi-host GPU job run different driver versions or have mismatched target configurations, each host builds a local `StreamExecutorGpuTopologyDescription` with a different fingerprint after `ExchangeTopologies`, causing compiled programs to diverge and hang in NCCL collective rendezvous. Add `GpuClientOptions::verify_topology_fingerprint` (enabled by default, overridable via `XLA_PJRT_GPU_VALIDATE_TOPOLOGY`) so that process 0 writes the expected topology fingerprint to `kv_store` and all other processes read and verify that their topology fingerprint matches. PiperOrigin-RevId: 990436201 --- .../xla/xla/pjrt/gpu/se_gpu_pjrt_client.cc | 67 +++++++++++++------ .../plugin/xla_gpu/xla_gpu_client_options.h | 2 + 2 files changed, 48 insertions(+), 21 deletions(-) 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..cc2a32b9cd8a66 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 @@ -1564,7 +1564,7 @@ CreateAllocatorMemoryRegistration(GpuAllocatorConfig* allocator_config) { struct PjRtDevicesAndTopology { std::vector> devices; - GpuTopologyProto topology; + std::shared_ptr topology; std::vector> local_device_states; }; @@ -1576,6 +1576,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 +1800,39 @@ 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(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 +1849,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 +2068,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 +2086,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 +2108,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 +2127,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 +2143,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/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 From 9eb09c736a9d89d36b8200058cb398db352b3882 Mon Sep 17 00:00:00 2001 From: "A. Unique TensorFlower" Date: Tue, 29 Sep 2026 11:43:09 -0700 Subject: [PATCH 02/24] Fix PjRt StreamExecutor execution to check BufferSequencingEvent error after waiting on stream. In `PjRtStreamExecutorRawLoadedExecutable::Execute`, `BufferSequencingEvent::WaitForEventOnStream()` blocks until the definition event has completed or failed. Previously, `ev->IsPredeterminedError()` was checked before `ev->WaitForEventOnStream()`. When an asynchronous host-to-device transfer (or any in-flight dependency) was still executing or pending, `ev->IsPredeterminedError()` returned false. `ev->WaitForEventOnStream()` then blocked until the event completed with error. But because `ev->IsPredeterminedError()` was not checked after the wait, the error was ignored, and `launch_on_device` proceeded to invoke `RunAsync()` on an uninitialized/poisoned buffer (or asymmetrically skipped `RunAsync` on only a subset of devices, causing cross-node / collective GPU hangs and deadlocks). Move `ev->WaitForEventOnStream()` before checking `ev->IsPredeterminedError()`, matching the behavior of the `else if (event)` branch where `xla::BlockUntilReady(event)` is called before inspecting `event.GetErrorIfPresent()`. Add regression unit tests in `se_gpu_pjrt_client_test.cc` and `se_gpu_pjrt_client_multi_gpu_test.cc`. PiperOrigin-RevId: 990437414 --- .../gpu/se_gpu_pjrt_client_multi_gpu_test.cc | 88 +++++++++++++++++++ .../xla/pjrt/gpu/se_gpu_pjrt_client_test.cc | 57 ++++++++++++ .../pjrt/se/pjrt_stream_executor_client.cc | 2 +- 3 files changed, 146 insertions(+), 1 deletion(-) 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/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()) { From 25a757887c047f0994ea4222a6cd235787eb0577 Mon Sep 17 00:00:00 2001 From: Dmitri Latushko Date: Tue, 29 Sep 2026 12:34:37 -0700 Subject: [PATCH 03/24] Add function_body dependency to Grappler functions library to fix Windows build link error. PiperOrigin-RevId: 990466106 --- tensorflow/core/grappler/utils/BUILD | 1 + 1 file changed, 1 insertion(+) 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", From 2f442bb17776eeffc6b36bf10dd0fdb455506026 Mon Sep 17 00:00:00 2001 From: "A. Unique TensorFlower" Date: Tue, 29 Sep 2026 12:45:46 -0700 Subject: [PATCH 04/24] [XLA] Fix TryRemoveTrivialCompare folding i > c to true when x == c In `TryRemoveTrivialCompare`, for a while loop induction variable `i` starting at initial value `c` (`i >= c`) compared against a constant `x`: - `i < x` (`ComparisonDirection::kLt`) is trivially `false` when `x <= c`. - `i > x` (`ComparisonDirection::kGt`) is trivially `true` only when `x < c` (strict inequality), not `x <= c`. When `x == c`, `i > c` evaluates to `false` on the first iteration (`i == c`) and `true` on subsequent iterations, so folding it to `true` miscompiles the first iteration. PiperOrigin-RevId: 990471976 --- .../xla/xla/service/while_loop_simplifier.cc | 28 ++++++++++--------- .../xla/service/while_loop_simplifier_test.cc | 11 +++++++- 2 files changed, 25 insertions(+), 14 deletions(-) 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, From 07de375eb1acd0b8fdae61701ddb549b434a363b Mon Sep 17 00:00:00 2001 From: Junwhan Ahn Date: Tue, 29 Sep 2026 13:05:47 -0700 Subject: [PATCH 05/24] [PJRT:GPU] Log GPU topology proto before checking topology fingerprint. Log `se_gpu_topology->ToProto()` for each process before computing and comparing topology fingerprints during multi-host GPU initialization to make topology mismatches easier to debug. PiperOrigin-RevId: 990482998 --- third_party/xla/xla/pjrt/gpu/BUILD | 1 + third_party/xla/xla/pjrt/gpu/se_gpu_pjrt_client.cc | 5 +++++ 2 files changed, 6 insertions(+) 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 cc2a32b9cd8a66..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" @@ -1814,6 +1815,10 @@ absl::StatusOr BuildDistributedDevices( 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); From b7c15f96070b881317d4dfdb92ebfaaeb4fac35f Mon Sep 17 00:00:00 2001 From: Alex Pivovarov Date: Tue, 29 Sep 2026 13:12:35 -0700 Subject: [PATCH 06/24] Fix nvcc compilation by using explicit traits in CUDA TopK This replaces the generic `ToOrdered`/`FromOrdered` free functions and `if constexpr` logic with an explicit `OrderedTraits` template struct, specialized for `float`, `xla::bfloat16`, and `xla::half`. This maintains functional parity while ensuring strict type safety and compatibility with older toolchains in the OSS CI. PiperOrigin-RevId: 990486882 --- .../cuda/topk_kernel_cuda_common.cu.h | 72 +++++++++++++------ .../rocm/topk_kernel_rocm_common.cu.h | 72 +++++++++++++------ 2 files changed, 104 insertions(+), 40 deletions(-) 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; } } From 7393fa678ee5fe999286cc551ba55257f5556d3b Mon Sep 17 00:00:00 2001 From: Quoc Truong Date: Tue, 29 Sep 2026 13:35:20 -0700 Subject: [PATCH 07/24] Change n2-128 runners to n4-80 and change the old L4 1GPU runner to the new L4 1GPU runner for benchmark. PiperOrigin-RevId: 990499379 --- .../xla/xla/tools/benchmarks/benchmark_registry.pbtxt | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) 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" From 9624b74b3e257e6231ad437f087dd5c066ed99ad Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Tom=C3=A1s=20Longeri?= Date: Tue, 29 Sep 2026 14:03:28 -0700 Subject: [PATCH 08/24] [Mosaic] Fix MemRefBitcastOp verification and tiling propagation The previous implementation completely only ever looked at the first element of the first tile, and scaled that. It didn't check trailing tiles and assumed the tiling was always of the form (A, B)(32 / bitwidth, 1). The new implementation does not assume anything about the tiling (e.g. it can handle a 16-bit type with a (4, 1) tile instead of (2, 1)) and can handle scaling n-D tiles such as (2, 4, 6)(2, 1) -> (2, 8, 6)(4, 1). PiperOrigin-RevId: 990515869 --- .../xla/xla/mosaic/dialect/tpu/tpu_ops.cc | 159 +++++++++++------- .../xla/xla/mosaic/dialect/tpu/tpu_ops.td | 4 +- 2 files changed, 105 insertions(+), 58 deletions(-) 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; From 099f9d4ffad80f39dd2617039bf4885817b0411f Mon Sep 17 00:00:00 2001 From: Mohammed Anany Date: Tue, 29 Sep 2026 14:10:05 -0700 Subject: [PATCH 09/24] [XLA:GPU] Evaluate tiling candidates in parallel in FusionBlockLevelRewriter PiperOrigin-RevId: 990520367 --- .../xla/xla/backends/gpu/transforms/BUILD | 3 ++ .../transforms/fusion_block_level_rewriter.cc | 15 +++++--- .../transforms/fusion_block_level_rewriter.h | 15 ++++++-- .../fusion_block_level_rewriter_test.cc | 34 +++++++++++++++++++ third_party/xla/xla/service/gpu/BUILD | 1 + .../service/gpu/fusion_dispatch_pipeline.cc | 7 ++-- .../service/gpu/fusion_dispatch_pipeline.h | 9 ++++- .../xla/xla/service/gpu/gpu_compiler.cc | 14 ++++---- .../xla/xla/service/gpu/gpu_compiler.h | 10 +++--- 9 files changed, 88 insertions(+), 20 deletions(-) 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/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/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/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..cb3d0763d34ec0 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 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_; } From 9eea59909643c19e7bccc5053a4f14441e527198 Mon Sep 17 00:00:00 2001 From: Dirk Hornung Date: Tue, 29 Sep 2026 14:12:18 -0700 Subject: [PATCH 10/24] [XLA:GPU] Support 1D broadcast parameters after pointwise chains in cuDNN conv fusion. PiperOrigin-RevId: 990521795 --- .../gpu/transforms/cudnn_fusion_compiler.cc | 23 ++++++------ .../cudnn_fusion_compiler_deviceless_test.cc | 35 +++++++++++++++++++ 2 files changed, 48 insertions(+), 10 deletions(-) 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. From 256b984545ff785b3d842d831884bb607e694f86 Mon Sep 17 00:00:00 2001 From: Nitin Srinivasan Date: Tue, 29 Sep 2026 14:15:41 -0700 Subject: [PATCH 11/24] Resolve ModuleNotFoundError in XLA CI and benchmark workflows when invoking the build script directly. Executing the CI build entrypoint outside of Bazel fails to resolve sibling package imports after recent cross-module refactoring removed the inline shell interpreter directive. Add the repository root to the Python module search path at runtime and invoke the Python interpreter explicitly across remaining workflow scripts. PiperOrigin-RevId: 990523829 --- third_party/xla/.kokoro/macos/build.sh | 2 +- third_party/xla/build_tools/ci/build.py | 10 ++++++++++ 2 files changed, 11 insertions(+), 1 deletion(-) 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( From c1c14eefc7b0591999e771b906d96f9c2bc9485b Mon Sep 17 00:00:00 2001 From: "A. Unique TensorFlower" Date: Tue, 29 Sep 2026 14:17:19 -0700 Subject: [PATCH 12/24] Enhance CUPTI activity overhead event naming for newer CUDA versions. Newer versions of CUPTI (up to CUDA 12.8+) introduce additional `CUpti_ActivityOverheadKind` enum values for runtime-triggered module loading, lazy function loading, command buffer full, activity buffer request, and UVM activity initialization overheads. Previously, these overhead kinds were unrecognized by XProf and appeared as `` in traces. Add `GetExtraActivityOverheadKindString12080` to `cuda_version_variants` (selected at build time via `if_cuda_newer_than("12_8", ...)` without `#if CUDA_VERSION` macros) and use it in `GetActivityOverheadKindString`. Format unknown or unrecognized overhead kinds as `Overhead::UNKNOWN:` so the underlying enum value remains visible for diagnosis. PiperOrigin-RevId: 990524975 --- .../xla/xla/backends/profiler/gpu/BUILD | 5 +++ .../profiler/gpu/cuda_version_12080_newer.cc | 20 +++++++++ .../profiler/gpu/cuda_version_12080_older.cc | 7 ++++ .../profiler/gpu/cuda_version_variants.h | 6 +++ .../profiler/gpu/cupti_buffer_events.cc | 40 ++++++++++-------- .../profiler/gpu/cupti_buffer_events.h | 4 ++ .../profiler/gpu/cupti_buffer_events_test.cc | 41 +++++++++++++++++++ 7 files changed, 105 insertions(+), 18 deletions(-) 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 From 65ae169f72648c2562d2d531821327a7e1b6da73 Mon Sep 17 00:00:00 2001 From: Praveen Batra Date: Tue, 29 Sep 2026 14:25:04 -0700 Subject: [PATCH 13/24] Factor vmem_for_operand into module level helper. PiperOrigin-RevId: 990529446 --- .../pallas_microbenchmarks/cost_model.py | 37 +++++++++++-------- 1 file changed, 22 insertions(+), 15 deletions(-) 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 ) From b7b6ee838c9aa0203ace9b02c7e923e2c77449cd Mon Sep 17 00:00:00 2001 From: "A. Unique TensorFlower" Date: Tue, 29 Sep 2026 14:25:39 -0700 Subject: [PATCH 14/24] Bug fix: Isolate cuda-only dependencies to if_cuda in gpu_device tests. PiperOrigin-RevId: 990529913 --- tensorflow/core/common_runtime/gpu/BUILD | 8 +++++--- 1 file changed, 5 insertions(+), 3 deletions(-) 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( From 3b5cfa5a0debe035c28a72a850bae47ce626f758 Mon Sep 17 00:00:00 2001 From: Matthias Guenther Date: Tue, 29 Sep 2026 14:34:26 -0700 Subject: [PATCH 15/24] Fix more deprecated MLIR API call sites The latest LLVM integration deprecates an MLIR API we use. A prior change migrated most call sites but missed some that were abstracted away by template functions; this change updates those call sites. PiperOrigin-RevId: 990535948 --- .../xla/third_party/stablehlo/temporary.patch | 47 ++++++++++++++++++- 1 file changed, 45 insertions(+), 2 deletions(-) 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 @@ } } } From f7d3e1d241077239d4b531b25d2f2f5d274f357a Mon Sep 17 00:00:00 2001 From: Emilio Cota Date: Tue, 29 Sep 2026 14:36:48 -0700 Subject: [PATCH 16/24] [xla:cpu] TargetMachineOptions: normalize target triple llvm::sys::getDefaultTargetTriple() can return a triple that omits the vendor component, e.g. "aarch64-linux-gnu". Constructing an llvm::Triple from it parses the components positionally ($arch,$vendor,$OS), which in this case results in an empty $OS. Reading $OS from the resulting triple results in a crash with `LLVM ERROR: unsupported operating system`. (This can be reproduced with the upcoming full msan support, which calls `TargetTriple.getOS()`.) Normalize the triple string consistently so that we always store canonicalized triples. That is, we don't save "aarch64-linux-gnu"; we save "aarch64-unknown-linux-gnu". Note that we do preserve empty triples ("") (vs. normalizing them to "unknown") so that llvm::EngineBuilder::selectTarget continues to fall back to the host process triple (it checks for ""). Also update the hardcoded host triple for oberon_b200 and oberon_b300 in gpu_topology and its test from "aarch64-linux-gnu" to the normalized "aarch64-unknown-linux-gnu". (this was the original, correct name---see cl/893544830---but it was changed to accommodate XLA:CPU's TargetMachineOptions. This change restores the original, correct string. PiperOrigin-RevId: 990537506 --- .../backends/cpu/codegen/ir_compiler_test.cc | 14 ++++++++++ .../backends/cpu/target_machine_options.cc | 17 +++++++++--- .../cpu/target_machine_options_test.cc | 27 +++++++++++++++++-- third_party/xla/xla/service/gpu_topology.cc | 2 +- .../xla/xla/service/gpu_topology_test.cc | 4 +-- 5 files changed, 55 insertions(+), 9 deletions(-) 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..f575f3869a746d 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"( 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/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) { From 2a50edbabe630972a99f204134ea572fd9e5c7c7 Mon Sep 17 00:00:00 2001 From: Parker Schuh Date: Tue, 29 Sep 2026 14:41:28 -0700 Subject: [PATCH 17/24] Rollforward of: Make an overrideable Delinearize function instead of allowing DelinearizeAsync to be overloaded explicitly. Reverts changelist 990024422 PiperOrigin-RevId: 990540455 --- .../xla/xla/pjrt/common_pjrt_client.cc | 29 ++++++++++--------- third_party/xla/xla/pjrt/common_pjrt_client.h | 14 ++++----- 2 files changed, 22 insertions(+), 21 deletions(-) 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( From 543710f084ca4513a5ac794fdc1951d47a4c8351 Mon Sep 17 00:00:00 2001 From: Maxim Ermilov Date: Tue, 29 Sep 2026 14:50:03 -0700 Subject: [PATCH 18/24] mark tsl::Future as ABSL_MUST_USE_RESULT Ignoring returned future is always an error PiperOrigin-RevId: 990545500 --- third_party/xla/xla/python/ifrt_proxy/client/rpc_helper.h | 7 ++++++- third_party/xla/xla/tsl/concurrency/future.h | 2 +- 2 files changed, 7 insertions(+), 2 deletions(-) diff --git a/third_party/xla/xla/python/ifrt_proxy/client/rpc_helper.h b/third_party/xla/xla/python/ifrt_proxy/client/rpc_helper.h index bff0bdeb4018b8..37d4579e99fe2f 100644 --- a/third_party/xla/xla/python/ifrt_proxy/client/rpc_helper.h +++ b/third_party/xla/xla/python/ifrt_proxy/client/rpc_helper.h @@ -77,8 +77,13 @@ class RpcHelper { return host_buffer_store_; } + // Returned ResponseFuture is safe to ignore. template - using ResponseFuture = tsl::Future>; + class ResponseFuture : public tsl::Future> { + public: + ResponseFuture(tsl::Future> future) + : tsl::Future>::Future(std::move(future)) {} + }; class Batcher; enum BatchOperation { kDeleteArray, kDestructArray, kSentinelDoNotUse }; diff --git a/third_party/xla/xla/tsl/concurrency/future.h b/third_party/xla/xla/tsl/concurrency/future.h index 487057effb65f5..f2f642b462ee35 100644 --- a/third_party/xla/xla/tsl/concurrency/future.h +++ b/third_party/xla/xla/tsl/concurrency/future.h @@ -653,7 +653,7 @@ class PromiseOnceMaker; // Future is a copyable type, although all copies share the same underlying // async value. template -class Future : public internal::FutureBase> { +class [[nodiscard]] Future : public internal::FutureBase> { using Base = internal::FutureBase>; static constexpr bool is_move_only = Base::IsMoveOnly(); // NOLINT From 2d81e2cdd114f18f6e6ec493d3430033ad55dab9 Mon Sep 17 00:00:00 2001 From: Emilio Cota Date: Tue, 29 Sep 2026 15:31:07 -0700 Subject: [PATCH 19/24] [xla:cpu] move sanitizer passes to the very end of the pipeline So that all generated code is instrumented. This paves the way for the upcoming msan support. PiperOrigin-RevId: 990569171 --- .../xla/backends/cpu/codegen/ir_compiler.cc | 27 +++++++++--- .../xla/backends/cpu/codegen/ir_compiler.h | 6 +++ .../backends/cpu/codegen/ir_compiler_test.cc | 44 +++++++++++++++++++ 3 files changed, 71 insertions(+), 6 deletions(-) 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 f575f3869a746d..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 @@ -489,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 From ef5b0adcce9791c78caa276dd431b6392e90e78a Mon Sep 17 00:00:00 2001 From: "A. Unique TensorFlower" Date: Tue, 29 Sep 2026 15:33:47 -0700 Subject: [PATCH 20/24] [XLA:TPU] Pin convolution output layout inside copy-disabled while loops. Inside a while loop carrying xla_disable_while_loop_copies, XLA is not free to insert a relayout copy. This change exposes IsWhileLoopCopyDisabled on ComputationLayoutConstraints and threads it to the TPU convolution output layout tie-break to prevent inserting relayout copies inside copy-disabled while loops. PiperOrigin-RevId: 990570717 --- third_party/xla/xla/service/layout_assignment.h | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) 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 From 8d2c456dd1d53c2fc7458cc6433e0f9e20a2ff3e Mon Sep 17 00:00:00 2001 From: Mohammed Anany Date: Tue, 29 Sep 2026 15:40:59 -0700 Subject: [PATCH 21/24] [XLA:GPU] Evaluate tiling candidates in parallel in BlockLevelEmitterBackend PiperOrigin-RevId: 990574854 --- .../xla/xla/backends/gpu/autotuner/BUILD | 9 +++ .../backends/gpu/autotuner/autotuner_main.cc | 3 +- .../gpu/autotuner/block_level_emitter.cc | 23 +++++++- .../gpu/autotuner/block_level_emitter.h | 15 ++++- .../gpu/autotuner/block_level_emitter_test.cc | 55 +++++++++++++++++++ .../xla/xla/backends/gpu/autotuner/factory.h | 9 ++- .../backends/gpu/autotuner/factory_cuda.cc | 7 ++- .../backends/gpu/autotuner/factory_rocm.cc | 7 ++- .../backends/gpu/autotuner/factory_test.cc | 3 +- .../xla/xla/service/gpu/autotuning/BUILD | 1 + .../gpu/autotuning/config_assigner_pass.cc | 7 ++- .../gpu/autotuning/config_assigner_pass.h | 5 +- .../xla/xla/service/gpu/gpu_compiler.cc | 3 +- 13 files changed, 131 insertions(+), 16 deletions(-) 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/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/gpu_compiler.cc b/third_party/xla/xla/service/gpu/gpu_compiler.cc index cb3d0763d34ec0..21b539b21b6e10 100644 --- a/third_party/xla/xla/service/gpu/gpu_compiler.cc +++ b/third_party/xla/xla/service/gpu/gpu_compiler.cc @@ -3533,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( From 6f21121abbe0333b9c5138abd4297d63ca5a6b64 Mon Sep 17 00:00:00 2001 From: Majid Dadashi Date: Tue, 29 Sep 2026 15:47:38 -0700 Subject: [PATCH 22/24] Add `cint2_fp32_int4_e8m0_drq` fusion pass and reference `FullyConnected` kernel. - Add `FuseA4W2DRQFullyConnectedPass` to collapse blockwise Q/DQ patterns (symmetric 32-element i4/e8m0 dynamic activations + centered per-channel i2 weights) into a single `tfl.fully_connected` carrying `tfl.quant_spec = {spec = "cint2_fp32_int4_e8m0_drq", act_dilation = ...}` and a per-axis `i2` `tfl.pseudo_qconst`. - Implement `ParseQuantSpec` and `EvalA4W2DRQ` in the `FullyConnected` reference kernel (bumping max version to 15) to evaluate `a4w2_drq_v1` and reject unrecognized `quant_spec` payloads. PiperOrigin-RevId: 990578615 --- tensorflow/compiler/mlir/lite/BUILD | 1 + .../lite/python/stablehlo_tfl_pipeline.cc | 5 + .../tests/fuse-a4w2-drq-fully-connected.mlir | 241 +++++++++++++++ .../fuse_a4w2_drq_fully_connected_pass.cc | 288 ++++++++++++++++++ .../compiler/mlir/lite/transforms/passes.h | 3 + .../compiler/mlir/lite/transforms/passes.td | 16 + tensorflow/lite/core/kernels/register.cc | 2 +- tensorflow/lite/kernels/BUILD | 1 + tensorflow/lite/kernels/fully_connected.cc | 245 +++++++++++++++ .../lite/kernels/fully_connected_test.cc | 252 +++++++++++++++ tensorflow/lite/kernels/register_ref.cc | 2 +- 11 files changed, 1054 insertions(+), 2 deletions(-) create mode 100644 tensorflow/compiler/mlir/lite/tests/fuse-a4w2-drq-fully-connected.mlir create mode 100644 tensorflow/compiler/mlir/lite/transforms/fuse_a4w2_drq_fully_connected_pass.cc 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/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(), From 17803f757062ee41a0f21658aac325d359e3167a Mon Sep 17 00:00:00 2001 From: Parker Schuh Date: Tue, 29 Sep 2026 15:47:43 -0700 Subject: [PATCH 23/24] Prefer `std::numeric_limits::lowest()` over `::min()` Based on https://en.cppreference.com/cpp/types/numeric_limits: `::min()` returns "the smallest positive normal value of the given floating-point type" while `::lowest()` returns "the lowest finite value". PiperOrigin-RevId: 990578654 --- third_party/xla/xla/pjrt/pjrt_executable.cc | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) 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); From 39489a634c5111835e265ff7ae32f04b22b254b1 Mon Sep 17 00:00:00 2001 From: "A. Unique TensorFlower" Date: Tue, 29 Sep 2026 15:51:11 -0700 Subject: [PATCH 24/24] mark tsl::Future as ABSL_MUST_USE_RESULT Ignoring returned future is always an error Reverts 543710f084ca4513a5ac794fdc1951d47a4c8351 PiperOrigin-RevId: 990580556 --- third_party/xla/xla/python/ifrt_proxy/client/rpc_helper.h | 7 +------ third_party/xla/xla/tsl/concurrency/future.h | 2 +- 2 files changed, 2 insertions(+), 7 deletions(-) diff --git a/third_party/xla/xla/python/ifrt_proxy/client/rpc_helper.h b/third_party/xla/xla/python/ifrt_proxy/client/rpc_helper.h index 37d4579e99fe2f..bff0bdeb4018b8 100644 --- a/third_party/xla/xla/python/ifrt_proxy/client/rpc_helper.h +++ b/third_party/xla/xla/python/ifrt_proxy/client/rpc_helper.h @@ -77,13 +77,8 @@ class RpcHelper { return host_buffer_store_; } - // Returned ResponseFuture is safe to ignore. template - class ResponseFuture : public tsl::Future> { - public: - ResponseFuture(tsl::Future> future) - : tsl::Future>::Future(std::move(future)) {} - }; + using ResponseFuture = tsl::Future>; class Batcher; enum BatchOperation { kDeleteArray, kDestructArray, kSentinelDoNotUse }; diff --git a/third_party/xla/xla/tsl/concurrency/future.h b/third_party/xla/xla/tsl/concurrency/future.h index f2f642b462ee35..487057effb65f5 100644 --- a/third_party/xla/xla/tsl/concurrency/future.h +++ b/third_party/xla/xla/tsl/concurrency/future.h @@ -653,7 +653,7 @@ class PromiseOnceMaker; // Future is a copyable type, although all copies share the same underlying // async value. template -class [[nodiscard]] Future : public internal::FutureBase> { +class Future : public internal::FutureBase> { using Base = internal::FutureBase>; static constexpr bool is_move_only = Base::IsMoveOnly(); // NOLINT