Skip to content

Commit 3ebf896

Browse files
[executorch][native] Fold torchao q4 weight dequantizes into embedding (pytorch#23338)
This PR was created by the merge bot to help merge the original PR into the main branch. ghstack PR number: pytorch#23233 by @digantdesai ^ Please use this as the source of truth for the PR details, comments, and reviews ghstack PR base: https://github.com/pytorch/executorch/tree/gh/digantdesai/113/base ghstack PR head: https://github.com/pytorch/executorch/tree/gh/digantdesai/113/head Merge bot PR base: https://github.com/pytorch/executorch/tree/gh/digantdesai/112/orig Merge bot PR head: https://github.com/pytorch/executorch/tree/gh/digantdesai/113/orig Differential Revision: [D121864324](https://our.internmc.facebook.com/intern/diff/D121864324/) @diff-train-skip-merge Co-authored-by: Digant Desai <digantdesai@meta.com>
1 parent defb915 commit 3ebf896

2 files changed

Lines changed: 51 additions & 19 deletions

File tree

‎backends/native/serialization/graph_serialize.py‎

Lines changed: 18 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -1054,15 +1054,18 @@ def _fold_gguf_dequant(graph_module: torch.fx.GraphModule) -> None:
10541054
node.meta["native_skip_serialize"] = True
10551055

10561056

1057-
def _is_linear_weight_use(dequantize: torch.fx.Node, user: torch.fx.Node) -> bool:
1057+
# The argument through which each op reads a foldable weight.
1058+
_WEIGHT_ARG_INDEX = {"aten::linear": 1, "aten::embedding": 0}
1059+
1060+
1061+
def _is_weight_use(dequantize: torch.fx.Node, user: torch.fx.Node) -> bool:
10581062
op = _resolve_op_overload(user.target) if user.op == "call_function" else None
1063+
index = _WEIGHT_ARG_INDEX.get(op.name()) if op is not None else None
10591064
return (
1060-
op is not None
1061-
and op.name() == "aten::linear"
1062-
and len(user.args) >= 2
1063-
and user.args[1] is dequantize
1064-
and user.args[0] is not dequantize
1065-
and dequantize not in user.args[2:]
1065+
index is not None
1066+
and len(user.args) > index
1067+
and user.args[index] is dequantize
1068+
and sum(arg is dequantize for arg in user.args) == 1
10661069
and dequantize not in user.kwargs.values()
10671070
)
10681071

@@ -1072,13 +1075,13 @@ def _mark_torchao_q4_weights(
10721075
) -> None:
10731076
"""Mark portable torchao q4 weights for packed PTN storage.
10741077
1075-
A weight dequantize read only as the weight of `aten.linear`, whose output
1076-
dtype matches its scales and whose operands are serialized state inputs, is
1077-
folded: the linear reads the packed `AffineGroup` weight directly and the
1078-
dequantize is not serialized. A weight is packed only if every one of its
1079-
readers is such a dequantize with the same parameters. Any other dequantize
1080-
stays in the graph over the plain int8 weight, with torchao semantics.
1081-
Activation q/dq nodes are never folded.
1078+
A weight dequantize read only as the weight of `aten.linear` or
1079+
`aten.embedding`, whose output dtype matches its scales and whose operands
1080+
are serialized state inputs, is folded: the op reads the packed `AffineGroup`
1081+
weight directly and the dequantize is not serialized. A weight is packed only
1082+
if every one of its readers is such a dequantize with the same parameters.
1083+
Any other dequantize stays in the graph over the plain int8 weight, with
1084+
torchao semantics. Activation q/dq nodes are never folded.
10821085
"""
10831086
if lifted_constants is None:
10841087
return
@@ -1132,7 +1135,7 @@ def _mark_torchao_q4_weights(
11321135
{weight.name, scale.name, zero_point.name} <= lifted_constants
11331136
and output.dtype == scale_value.dtype
11341137
and node.users
1135-
and all(_is_linear_weight_use(node, user) for user in node.users)
1138+
and all(_is_weight_use(node, user) for user in node.users)
11361139
):
11371140
folds.setdefault(weight, []).append((node, packed))
11381141

‎backends/native/test/test_quant.py‎

Lines changed: 33 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -128,9 +128,9 @@ def test_rejects_values_outside_signed_int4_range(self):
128128
_pack_signed_int4(torch.tensor([[value, 0]], dtype=torch.int8))
129129

130130

131-
def _linear_weight(method):
132-
[linear] = [n for n in method.graph.nodes if n.target and "aten.linear" in n.target]
133-
[weight] = [a.arg.value for a in linear.inputs if a.name == "weight"]
131+
def _weight_of(method, op="aten.linear"):
132+
[node] = [n for n in method.graph.nodes if n.target and op in n.target]
133+
[weight] = [a.arg.value for a in node.inputs if a.name == "weight"]
134134
assert isinstance(weight, TensorArg)
135135
return weight.name
136136

@@ -148,7 +148,7 @@ def _assert_linear_reads_packed_weight(self, method, constants):
148148
[weight] = [c for c in method.constants if c.meta.quant is not None]
149149
scheme = weight.meta.quant.scheme
150150
self.assertIsInstance(scheme, AffineGroup)
151-
self.assertEqual(_linear_weight(method), weight.name)
151+
self.assertEqual(_weight_of(method), weight.name)
152152
self.assertEqual(weight.meta.dtype, ScalarType.BYTE)
153153
self.assertEqual([d.max for d in weight.meta.sizes], [self.N, self.K])
154154
self.assertEqual(
@@ -185,6 +185,35 @@ def test_4w_linear_reads_packed_weight(self):
185185
self._assert_linear_reads_packed_weight(method, constants)
186186

187187

188+
class TorchaoQ4EmbeddingTest(unittest.TestCase):
189+
def test_embedding_reads_packed_weight(self):
190+
rows, cols, group_size = 16, 64, 32
191+
model = nn.Sequential(nn.Embedding(rows, cols))
192+
quantize_(
193+
model,
194+
IntxWeightOnlyConfig(
195+
weight_dtype=torch.int4, granularity=PerGroup(group_size)
196+
),
197+
filter_fn=lambda m, _: isinstance(m, nn.Embedding),
198+
)
199+
indices = torch.tensor([[1, 5, 7]])
200+
manager = to_native(torch.export.export(model, (indices,)))
201+
method = deserialize_program(manager._ptg).methods[0]
202+
203+
self.assertFalse(
204+
any("dequantize_affine" in t for t in _call_function_targets(method.graph))
205+
)
206+
[weight] = [c for c in method.constants if c.meta.quant is not None]
207+
scheme = weight.meta.quant.scheme
208+
self.assertEqual(_weight_of(method, "aten.embedding"), weight.name)
209+
self.assertEqual([d.max for d in weight.meta.sizes], [rows, cols])
210+
self.assertEqual(
211+
(scheme.quant_min, scheme.quant_max, scheme.group_size),
212+
(-8, 7, group_size),
213+
)
214+
self.assertEqual(manager._constants[weight.data_key].numel(), rows * cols // 2)
215+
216+
188217
class _DequantizedLinear(nn.Module):
189218
"""linear(x, dequantize_affine(weight)) over an int4 [8, 16] weight in groups
190219
of 8, optionally with the weight as an input or the dequantize read twice."""

0 commit comments

Comments
 (0)