Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
24 commits
Select commit Hold shift + click to select a range
3152ae3
[PJRT:GPU] Fail distributed initialization when hosts have different …
junwhanahn Sep 29, 2026
9eb09c7
Fix PjRt StreamExecutor execution to check BufferSequencingEvent erro…
tensorflower-gardener Sep 29, 2026
25a7578
Add function_body dependency to Grappler functions library to fix Win…
dmiltr3 Sep 29, 2026
2f442bb
[XLA] Fix TryRemoveTrivialCompare folding i > c to true when x == c
tensorflower-gardener Sep 29, 2026
07de375
[PJRT:GPU] Log GPU topology proto before checking topology fingerprint.
junwhanahn Sep 29, 2026
b7c15f9
Fix nvcc compilation by using explicit traits in CUDA TopK
apivovarov Sep 29, 2026
7393fa6
Change n2-128 runners to n4-80 and change the old L4 1GPU runner to t…
quoctruong Sep 29, 2026
9624b74
[Mosaic] Fix MemRefBitcastOp verification and tiling propagation
tlongeri Sep 29, 2026
099f9d4
[XLA:GPU] Evaluate tiling candidates in parallel in FusionBlockLevelR…
Moerafaat Sep 29, 2026
9eea599
[XLA:GPU] Support 1D broadcast parameters after pointwise chains in c…
derdrdirk Sep 29, 2026
256b984
Resolve ModuleNotFoundError in XLA CI and benchmark workflows when in…
nitins17 Sep 29, 2026
c1c14ee
Enhance CUPTI activity overhead event naming for newer CUDA versions.
tensorflower-gardener Sep 29, 2026
65ae169
Factor vmem_for_operand into module level helper.
18praveenb Sep 29, 2026
b7b6ee8
Bug fix: Isolate cuda-only dependencies to if_cuda in gpu_device tests.
tensorflower-gardener Sep 29, 2026
3b5cfa5
Fix more deprecated MLIR API call sites
mrguenther Sep 29, 2026
f7d3e1d
[xla:cpu] TargetMachineOptions: normalize target triple
cota Sep 29, 2026
2a50edb
Rollforward of: Make an overrideable Delinearize function instead of …
pschuh Sep 29, 2026
543710f
mark tsl::Future as ABSL_MUST_USE_RESULT
ermilovmaxim Sep 29, 2026
2d81e2c
[xla:cpu] move sanitizer passes to the very end of the pipeline
cota Sep 29, 2026
ef5b0ad
[XLA:TPU] Pin convolution output layout inside copy-disabled while lo…
tensorflower-gardener Sep 29, 2026
8d2c456
[XLA:GPU] Evaluate tiling candidates in parallel in BlockLevelEmitter…
Moerafaat Sep 29, 2026
6f21121
Add `cint2_fp32_int4_e8m0_drq` fusion pass and reference `FullyConnec…
majiddadashi Sep 29, 2026
17803f7
Prefer `std::numeric_limits<T>::lowest()` over `::min()`
pschuh Sep 29, 2026
39489a6
mark tsl::Future as ABSL_MUST_USE_RESULT
tensorflower-gardener Sep 29, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions tensorflow/compiler/mlir/lite/BUILD
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -301,6 +301,11 @@ void AddPipelinePasses(mlir::OpPassManager& pass_manager,
mlir::createCanonicalizerPass());
pass_manager.addNestedPass<mlir::func::FuncOp>(mlir::createCSEPass());
pass_manager.addPass(CreatePruneDeadResourcesPass());
pass_manager.addNestedPass<mlir::func::FuncOp>(
mlir::TFL::CreateFuseA4W2DRQFullyConnectedPass());
pass_manager.addNestedPass<mlir::func::FuncOp>(
mlir::createCanonicalizerPass());
pass_manager.addPass(mlir::createSymbolDCEPass());
pass_manager.addPass(mlir::TFL::CreateCleanupOptimizationBarrierPass());
pass_manager.addPass(mlir::odml::createLegalizeStablehloToVhloPass());
pass_manager.addPass(mlir::createReconcileUnrealizedCastsPass());
Expand Down
241 changes: 241 additions & 0 deletions tensorflow/compiler/mlir/lite/tests/fuse-a4w2-drq-fully-connected.mlir

Large diffs are not rendered by default.

Original file line number Diff line number Diff line change
@@ -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 <cstddef>
#include <cstdint>
#include <memory>
#include <utility>

#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<NoneType>(value.getType())) return false;

ElementsAttr attr;
if (!matchPattern(value, m_Constant(&attr))) {
auto const_op = value.getDefiningOp<TFL::ConstOp>();
if (!const_op) return false;
attr = const_op.getValue();
}

auto fp_attr = mlir::dyn_cast<DenseFPElementsAttr>(attr);
if (!fp_attr) return false;
return llvm::all_of(fp_attr.getValues<APFloat>(), [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<NxK x
// !quant.uniform<i2:f32:0, {scales}>> %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<TFL::FullyConnectedOp> {
using OpRewritePattern<TFL::FullyConnectedOp>::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<BlockwiseDequantizeOp>();
if (!act_deq) return failure();

mlir::Value act_q_val = act_deq.getInput();
auto act_q = act_q_val.getDefiningOp<BlockwiseQuantizeOp>();
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<RankedTensorType>(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<Float8E8M0FNUType>(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<IntegerAttr>(act_block_shape[act_rank - 1]).getInt() !=
kActBlockSize) {
return failure();
}
for (int64_t i = 0; i < act_rank - 1; ++i) {
if (mlir::cast<IntegerAttr>(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<BlockwiseDequantizeOp>();
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<TFL::ConstOp>()) {
q_weight_attr = const_op.getValue();
} else {
return failure();
}
}
auto q_weight_type =
mlir::dyn_cast<RankedTensorType>(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<IntegerAttr>(weight_block_shape[0]).getInt() != 1 ||
mlir::cast<IntegerAttr>(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<TFL::ConstOp>()) {
scales_attr = const_op.getValue();
} else {
return failure();
}
}

SmallVector<double> per_channel_scales;
per_channel_scales.reserve(num_units);
if (auto dense_scales = mlir::dyn_cast<DenseFPElementsAttr>(scales_attr)) {
for (const auto& fp : dense_scales.getValues<APFloat>()) {
per_channel_scales.push_back(fp.convertToDouble());
}
} else {
return failure();
}
if (per_channel_scales.size() != static_cast<size_t>(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<int64_t> 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<TFL::QConstOp>(
fc.getLoc(), TypeAttr::get(new_filter_type), q_weight_attr);

// 4. Create tfl.quant_spec attribute dictionary.
SmallVector<NamedAttribute, 2> 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<TFL::FullyConnectedOp>(
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<FuseA4W2DRQFullyConnectedPattern>(context);

if (failed(applyPatternsGreedily(func, std::move(patterns)))) {
signalPassFailure();
}
}
};

} // namespace

std::unique_ptr<OperationPass<func::FuncOp>>
CreateFuseA4W2DRQFullyConnectedPass() {
return std::make_unique<FuseA4W2DRQFullyConnectedPass>();
}

} // namespace TFL
} // namespace mlir
3 changes: 3 additions & 0 deletions tensorflow/compiler/mlir/lite/transforms/passes.h
Original file line number Diff line number Diff line change
Expand Up @@ -126,6 +126,9 @@ std::unique_ptr<OperationPass<func::FuncOp>> CreateDefaultQuantizePass();
std::unique_ptr<OperationPass<func::FuncOp>>
CreateFoldStablehloConstantTransformsPass();

std::unique_ptr<OperationPass<func::FuncOp>>
CreateFuseA4W2DRQFullyConnectedPass();

std::unique_ptr<OperationPass<ModuleOp>> CreateLowerQuantAnnotationsPass();

// Creates an instance of the TFLite PropagateQParams pass which propagates
Expand Down
16 changes: 16 additions & 0 deletions tensorflow/compiler/mlir/lite/transforms/passes.td
Original file line number Diff line number Diff line change
Expand Up @@ -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 = [{
Expand Down
8 changes: 5 additions & 3 deletions tensorflow/core/common_runtime/gpu/BUILD
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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(
Expand Down
1 change: 1 addition & 0 deletions tensorflow/core/grappler/utils/BUILD
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
2 changes: 1 addition & 1 deletion tensorflow/lite/core/kernels/register.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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(),
Expand Down
1 change: 1 addition & 0 deletions tensorflow/lite/kernels/BUILD
Original file line number Diff line number Diff line change
Expand Up @@ -2269,6 +2269,7 @@ cc_test(
"//tensorflow/lite/schema:schema_fbs",
"@com_google_absl//absl/log:absl_check",
"@com_google_googletest//:gtest",
"@flatbuffers",
],
)

Expand Down
Loading
Loading