MCPcopy Create free account
hub / github.com/pytorch/executorch / QuantizedEmbeddingHandler

Class QuantizedEmbeddingHandler

backends/mlx/patterns.py:947–1074  ·  view source on GitHub ↗

Pattern for quantized embedding: dequantize_affine + embedding.

Source from the content-addressed store, hash-verified

945
946@REGISTRY.register_pattern(name="QUANTIZED_EMBEDDING")
947class 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 )

Callers 1

maybe_createMethod · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected