Pattern for quantized embedding: dequantize_affine + embedding.
| 945 | |
| 946 | @REGISTRY.register_pattern(name="QUANTIZED_EMBEDDING") |
| 947 | class QuantizedEmbeddingHandler(PatternHandler): |
| 948 | """ |
| 949 | Pattern for quantized embedding: dequantize_affine + embedding. |
| 950 | """ |
| 951 | |
| 952 | def __init__( |
| 953 | self, |
| 954 | head: Node, |
| 955 | body: List[Node], |
| 956 | qdata: Node, |
| 957 | scale: Node, |
| 958 | zero_point: Node, |
| 959 | group_size: int, |
| 960 | bits: int, |
| 961 | out_dtype: torch.dtype, |
| 962 | ): |
| 963 | super().__init__(head, body) |
| 964 | self.qdata = qdata |
| 965 | self.scale = scale |
| 966 | self.zero_point = zero_point |
| 967 | self.group_size = group_size |
| 968 | self.bits = bits |
| 969 | self.out_dtype = out_dtype |
| 970 | |
| 971 | @classmethod |
| 972 | def maybe_create( |
| 973 | cls, ep: ExportedProgram, head: Node |
| 974 | ) -> Optional["QuantizedEmbeddingHandler"]: |
| 975 | embedding_node = head |
| 976 | if not match_target(embedding_node, torch.ops.aten.embedding.default): |
| 977 | return None |
| 978 | |
| 979 | w, x = embedding_node.args[0:2] |
| 980 | |
| 981 | dequant_node = w |
| 982 | if not match_target(dequant_node, torch.ops.torchao.dequantize_affine.default): |
| 983 | return None |
| 984 | if not has_single_user(dequant_node): |
| 985 | return None |
| 986 | |
| 987 | parsed = parse_dequant_node(dequant_node) |
| 988 | if parsed is None: |
| 989 | return None |
| 990 | qdata, scale, zero_point, group_size, bits, out_dtype, _quantized_dim = parsed |
| 991 | out_dtype = scale.meta["val"].dtype if out_dtype is None else out_dtype |
| 992 | |
| 993 | head = embedding_node |
| 994 | body = [dequant_node] |
| 995 | return QuantizedEmbeddingHandler( |
| 996 | head, |
| 997 | body, |
| 998 | qdata=qdata, |
| 999 | scale=scale, |
| 1000 | zero_point=zero_point, |
| 1001 | group_size=group_size, |
| 1002 | bits=bits, |
| 1003 | out_dtype=out_dtype, |
| 1004 | ) |