Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
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
36 changes: 36 additions & 0 deletions backends/cadence/aot/BUCK
Original file line number Diff line number Diff line change
Expand Up @@ -97,6 +97,8 @@ fbcode_target(_kind = runtime.python_library,
],
deps = [
":fuse_ops",
":pass_utils",
":weight_packing",
":remove_ops",
":reorder_ops",
":replace_ops",
Expand Down Expand Up @@ -131,6 +133,7 @@ fbcode_target(_kind = runtime.python_library,
deps = [
"fbcode//caffe2:torch",
"fbcode//executorch/backends/cadence/aot:utils",
"fbcode//executorch/backends/cadence/aot:weight_packing",
"fbcode//executorch/exir:scalar_type",
"fbcode//executorch/kernels/quantized:custom_ops_generated_lib",
],
Expand All @@ -157,6 +160,7 @@ fbcode_target(_kind = executorch_generated_lib,
"//executorch/backends/cadence/generic/operators:op_quantized_depthwise_conv1d_ncl",
"//executorch/backends/cadence/generic/operators:op_quantized_depthwise_conv1d_nlc",
"//executorch/backends/cadence/generic/operators:op_quantized_fully_connected",
"//executorch/backends/cadence/generic/operators:op_quantized_fully_connected_packed",
"//executorch/backends/cadence/generic/operators:op_quantized_layer_norm",
"//executorch/backends/cadence/generic/operators:op_quantized_linear",
"//executorch/backends/cadence/generic/operators:op_quantized_matmul",
Expand Down Expand Up @@ -637,6 +641,37 @@ fbcode_target(_kind = python_unittest,
],
)

fbcode_target(_kind = runtime.python_library,
name = "weight_packing",
srcs = [
"weight_packing.py",
],
deps = [
"fbcode//caffe2:torch",
"//executorch/backends/cadence/aot/quantizer:pattern_utils",
"//executorch/exir:lib",
],
)

fbcode_target(_kind = python_unittest,
name = "test_weight_packing",
srcs = [
"tests/test_weight_packing.py",
],
supports_static_listing = False,
deps = [
"//caffe2:torch",
"//executorch/backends/cadence/aot:compiler",
"//executorch/backends/cadence/aot:ops_registrations",
"//executorch/backends/cadence/aot:pass_utils",
"//executorch/backends/cadence/aot:weight_packing",
"//executorch/backends/cadence/aot:graph_builder",
"//executorch/backends/cadence/aot/quantizer:fusion_pass",
"//executorch/backends/cadence/aot/quantizer:quantizer",
"//pytorch/ao:torchao",
]
)

fbcode_target(_kind = python_unittest,
name = "test_ref_implementations",
srcs = [
Expand All @@ -646,6 +681,7 @@ fbcode_target(_kind = python_unittest,
deps = [
":typing_stubs",
"//executorch/backends/cadence/aot:ops_registrations",
"//executorch/backends/cadence/aot:weight_packing",
"//executorch/backends/cadence/aot/quantizer:utils",
"//caffe2:torch",
]
Expand Down
5 changes: 5 additions & 0 deletions backends/cadence/aot/functions.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -494,6 +494,11 @@
- arg_meta: null
kernel_name: impl::generic::quantized_fully_connected_asym8uxasym8u_asym8u_per_tensor_out

- func: cadence::quantized_fully_connected_packed.out(Tensor src, Tensor weight, Tensor bias, int in_dim, int weight_bits, int src_zero_point, Tensor? weight_zero_point, Tensor out_multiplier, Tensor out_shift, int out_zero_point, Tensor? offset, *, Tensor(a!) out) -> Tensor(a!)
kernels:
- arg_meta: null
kernel_name: impl::generic::quantized_fully_connected_packed_out

- func: cadence::requantize.out(Tensor input, Tensor in_scale, Tensor in_zero_point, Tensor out_scale, Tensor out_zero_point, ScalarType out_dtype, *, Tensor(a!) out) -> Tensor(a!)
kernels:
- arg_meta: null
Expand Down
7 changes: 7 additions & 0 deletions backends/cadence/aot/functions_hifi.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -460,6 +460,13 @@
- arg_meta: null
kernel_name: impl::HiFi::quantized_conv2d_nhwc_depthwise_asym8uxsym8u_asym8u_per_tensor_out

- func: cadence::quantized_fully_connected_packed.out(Tensor src, Tensor weight, Tensor bias, int in_dim, int weight_bits, int src_zero_point, Tensor? weight_zero_point, Tensor out_multiplier, Tensor out_shift, int out_zero_point, Tensor? offset, *, Tensor(a!) out) -> Tensor(a!)
kernels:
- arg_meta: null
# No packed kernel on this backend yet: fall back to the generic one
# rather than failing to resolve the operator.
kernel_name: impl::generic::quantized_fully_connected_packed_out

- func: cadence::quantized_layer_norm.out(Tensor input, Tensor in_scale, Tensor in_zero_point, int[] normalized_shape, Tensor weight, Tensor bias, float eps, float output_scale, int output_zero_point, *, Tensor(a!) out) -> Tensor(a!)
kernels:
- arg_meta: null
Expand Down
39 changes: 39 additions & 0 deletions backends/cadence/aot/ops_registrations.py
Original file line number Diff line number Diff line change
Expand Up @@ -482,6 +482,13 @@ def register_fake(
"quantized_fully_connected_asym8uxasym8u_asym8u.per_tensor(Tensor src, Tensor weight, Tensor bias, int src_zero_point, "
"int weight_zero_point, int out_multiplier, int out_shift, int out_zero_point, Tensor? offset) -> (Tensor Z)"
)
# Sub-byte weights: `weight` is a [out_dim, packed_row_bytes] int8 blob, so
# in_dim and the bit width can no longer be read off its shape.
lib.define(
"quantized_fully_connected_packed(Tensor src, Tensor weight, Tensor bias, int in_dim, int weight_bits, "
"int src_zero_point, Tensor? weight_zero_point, Tensor out_multiplier, Tensor out_shift, int out_zero_point, "
"Tensor? offset) -> (Tensor Z)"
)
lib.define("where_Scalar(Tensor condition, float self, float other) -> (Tensor Z)")
lib.define(
"where_Scalar.out(Tensor condition, float self, float other, *, Tensor(a!) out) -> Tensor(a!)"
Expand Down Expand Up @@ -666,6 +673,11 @@ def register_fake(
"quantized_fully_connected_asym8uxasym8u_asym8u.per_tensor_out(Tensor src, Tensor weight, Tensor bias, int src_zero_point, "
"int weight_zero_point, int out_multiplier, int out_shift, int out_zero_point, Tensor? offset, *, Tensor(a!) out) -> Tensor(a!)"
)
lib.define(
"quantized_fully_connected_packed.out(Tensor src, Tensor weight, Tensor bias, int in_dim, int weight_bits, "
"int src_zero_point, Tensor? weight_zero_point, Tensor out_multiplier, Tensor out_shift, int out_zero_point, "
"Tensor? offset, *, Tensor(a!) out) -> Tensor(a!)"
)
lib.define(
"quantized_embedding_byte.out(Tensor weight, Tensor weight_scales, Tensor? weight_zero_points, "
"Tensor indices, bool pruned_weights=False, *, Tensor(a!) out) -> Tensor(a!)"
Expand Down Expand Up @@ -2769,6 +2781,33 @@ def quantized_fully_connected_per_tensor_meta(
return src.new_empty(out_size, dtype=src.dtype)


@register_fake("cadence::quantized_fully_connected_packed")
def quantized_fully_connected_packed_meta(
src: torch.Tensor,
weight: torch.Tensor,
bias: torch.Tensor,
in_dim: int,
weight_bits: int,
in_zero_point: int,
weight_zero_point: Optional[torch.Tensor],
out_multiplier: torch.Tensor,
out_shift: torch.Tensor,
out_zero_point: int,
offset: Optional[torch.Tensor],
) -> torch.Tensor:
torch._check(bias.dtype == torch.int32, lambda: "expected int32")
torch._check(weight.dim() == 2, lambda: "expected 2D tensor")
torch._check(src.size(0) == 1, lambda: "expected batch size of 1")
# src comes in shape [leading_dims, in_dim]
# weight comes in shape [out_dim, packed_row_bytes], so out_dim is still
# size(0) - that is what packing per row buys us - but in_dim is an
# argument rather than size(1).
torch._check(src.size(-1) == in_dim, lambda: "src last dim must equal in_dim")
out_size = list(src.size())
out_size[-1] = weight.size(0)
return src.new_empty(out_size, dtype=src.dtype)


@register_fake("cadence::quantized_fully_connected_asym8sxasym8s_asym8s.per_tensor")
def quantized_fully_connected_asym8sxasym8s_asym8s_per_tensor_meta(
src: torch.Tensor,
Expand Down
4 changes: 4 additions & 0 deletions backends/cadence/aot/pass_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,10 @@ class CompileMode(Enum):
@dataclass(frozen=True)
class EdgePassesConfig:
use_im2row_transform: bool = False
# Storage width for fully-connected weights. 8 leaves them as-is; 4 and 6
# physically pack them, which is only a size win, not an accuracy one - the
# values are expected to already be clamped to the narrower range.
weight_bits: int = 8


# Return the overload packet for the edge or torch op.
Expand Down
10 changes: 10 additions & 0 deletions backends/cadence/aot/passes.py
Original file line number Diff line number Diff line change
Expand Up @@ -61,6 +61,9 @@
SimplifySliceOpPass,
)
from executorch.backends.cadence.aot.type_dispatch import CompileTimeTypeDispatchPass
from executorch.backends.cadence.aot.weight_packing import (
fold_and_pack_fully_connected_weights,
)
from executorch.exir import EdgeProgramManager
from executorch.exir.pass_base import ExportPass, PassResult
from executorch.exir.pass_manager import PassManager, PassType
Expand Down Expand Up @@ -198,6 +201,13 @@ def apply_exir_ops_passes(
list[Callable[[torch.fx.GraphModule], Optional[PassResult]]], cadence_passes
)
)
# Weight packing runs last and outside the pass list on purpose: it changes
# the shape of a weight constant, so it needs the ExportedProgram rather
# than the GraphModule that `transform` hands to a pass.
config = edge_passes_config or EdgePassesConfig()
fold_and_pack_fully_connected_weights(
cadence_prog_manager.exported_program(), config.weight_bits
)
return cadence_prog_manager


Expand Down
32 changes: 32 additions & 0 deletions backends/cadence/aot/ref_implementations.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
import torch.nn as nn
import torch.nn.functional as F
from executorch.backends.cadence.aot.utils import is_depthwise_conv
from executorch.backends.cadence.aot.weight_packing import unpack_rows
from executorch.exir.scalar_type import ScalarType
from torch.library import impl, Library

Expand Down Expand Up @@ -683,6 +684,37 @@ def quantized_fully_connected_asym8sxasym8s_asym8s_per_tensor() -> torch.Tensor:
def quantized_fully_connected_asym8uxasym8u_asym8u_per_tensor() -> torch.Tensor: ...


@impl_tracked(m, "quantized_fully_connected_packed")
def quantized_fully_connected_packed(
src: torch.Tensor,
weight: torch.Tensor,
bias: torch.Tensor,
in_dim: int,
weight_bits: int,
in_zero_point: int,
weight_zero_point: Optional[torch.Tensor],
out_multiplier: torch.Tensor,
out_shift: torch.Tensor,
out_zero_point: int,
offset: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""Sub-byte weights: unpack, then reuse the ordinary quantized path.

Keeping the arithmetic in `quantized_linear_common` means packing cannot
drift from the unpacked operator, which is exactly what the tests assert.
"""
return quantized_linear_common(
src,
unpack_rows(weight, in_dim, weight_bits),
bias,
in_zero_point,
weight_zero_point,
out_multiplier,
out_shift,
out_zero_point,
)


@impl_tracked(m, "fully_connected")
def fully_connected(
input_tensor: torch.Tensor,
Expand Down
Loading
Loading