From 7c533c44b25af34f7469314f0b48146fe9de7883 Mon Sep 17 00:00:00 2001 From: Ethan Ng Date: Tue, 1 Sep 2026 16:13:33 -0700 Subject: [PATCH] Add portable quantized stacked-halves RoPE operator (#22423) Summary: Define `cadence::quantized_rope_rotate_stacked_halves` with fake and reference implementations, selective-build registration, and a portable runtime fallback. Add focused coverage for stacked layout, position selection, aligned-sized inputs, and SIMD-tail-sized inputs. Differential Revision: D118320273 --- backends/cadence/aot/ops_registrations.py | 21 ++++++++++++++ backends/cadence/aot/ref_implementations.py | 32 +++++++++++++++++++++ 2 files changed, 53 insertions(+) diff --git a/backends/cadence/aot/ops_registrations.py b/backends/cadence/aot/ops_registrations.py index da82a1ea3ec..6789d17f16b 100644 --- a/backends/cadence/aot/ops_registrations.py +++ b/backends/cadence/aot/ops_registrations.py @@ -489,6 +489,13 @@ def register_fake( "rope_rotate_stacked_halves.out(Tensor input, Tensor sin_tensor, Tensor cos_tensor, Tensor? pos, *, Tensor(a!) out) -> Tensor(a!)" ) +lib.define( + "quantized_rope_rotate_stacked_halves(Tensor input, Tensor sin_tensor, Tensor cos_tensor, Tensor? pos, float in_scale, int in_zero_point, float out_scale, int out_zero_point) -> (Tensor out)" +) +lib.define( + "quantized_rope_rotate_stacked_halves.out(Tensor input, Tensor sin_tensor, Tensor cos_tensor, Tensor? pos, float in_scale, int in_zero_point, float out_scale, int out_zero_point, *, Tensor(a!) out) -> Tensor(a!)" +) + lib.define( "quantized_softmax(Tensor input, Tensor mask, int dim, int mask_type, Tensor pos, Tensor in_scale, Tensor in_zero_point, Tensor out_scale, Tensor out_zero_point) -> (Tensor out)" ) @@ -3136,6 +3143,20 @@ def rope_rotate_stacked_halves_meta( return input.new_empty(input.shape, dtype=input.dtype) +@register_fake("cadence::quantized_rope_rotate_stacked_halves") +def quantized_rope_rotate_stacked_halves_meta( + input: torch.Tensor, + sin_tensor: torch.Tensor, + cos_tensor: torch.Tensor, + pos: Optional[torch.Tensor], + in_scale: float, + in_zero_point: int, + out_scale: float, + out_zero_point: int, +) -> torch.Tensor: + return rope_rotate_stacked_halves_meta(input, sin_tensor, cos_tensor, pos) + + @register_fake("cadence::idma_copy") def copy_idma_copy_impl( src: torch.Tensor, diff --git a/backends/cadence/aot/ref_implementations.py b/backends/cadence/aot/ref_implementations.py index d3a5c853a4a..16768a8b68e 100644 --- a/backends/cadence/aot/ref_implementations.py +++ b/backends/cadence/aot/ref_implementations.py @@ -2265,6 +2265,38 @@ def rope_rotate_stacked_halves( return rotated.view(original_shape) +@impl_tracked(m, "quantized_rope_rotate_stacked_halves") +def quantized_rope_rotate_stacked_halves( + input_tensor: torch.Tensor, + sin_tensor: torch.Tensor, + cos_tensor: torch.Tensor, + pos: torch.Tensor | None, + in_scale: float, + in_zero_point: int, + out_scale: float, + out_zero_point: int, +) -> torch.Tensor: + dtype = input_tensor.dtype + dtype_limits = torch.iinfo(dtype) + dequantized = dequantize_per_tensor_common( + input_tensor, + in_scale, + in_zero_point, + dtype_limits.min, + dtype_limits.max, + dtype, + ) + rotated = rope_rotate_stacked_halves(dequantized, sin_tensor, cos_tensor, pos) + return quantize_per_tensor_common( + rotated, + out_scale, + out_zero_point, + dtype_limits.min, + dtype_limits.max, + dtype, + ) + + @impl_tracked(m, "im2row") def im2row( input_tensor: torch.Tensor,