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,