Skip to content

Commit defb915

Browse files
[executorch][native] Fold torchao q4 weight dequantizes into linear (pytorch#23337)
This PR was created by the merge bot to help merge the original PR into the main branch. ghstack PR number: pytorch#23232 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/112/base ghstack PR head: https://github.com/pytorch/executorch/tree/gh/digantdesai/112/head Merge bot PR base: https://github.com/pytorch/executorch/tree/gh/digantdesai/111/orig Merge bot PR head: https://github.com/pytorch/executorch/tree/gh/digantdesai/112/orig Differential Revision: [D121704321](https://our.internmc.facebook.com/intern/diff/D121704321/) @diff-train-skip-merge Co-authored-by: Digant Desai <digantdesai@meta.com>
1 parent c589d6c commit defb915

3 files changed

Lines changed: 203 additions & 14 deletions

File tree

‎backends/native/serialization/graph_serialize.py‎

Lines changed: 59 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -390,8 +390,8 @@ def _node_to_arg_value(
390390
v: torch.fx.Node,
391391
subgraph_map: "dict[str, torch.fx.GraphModule] | None",
392392
) -> ArgumentValue:
393-
# A GGUF weight consumed through dequantize_gguf serializes as the raw packed
394-
# constant (see _fold_gguf_dequant), so redirect the reference to it.
393+
# A packed weight consumed through a folded dequantize serializes as the packed
394+
# constant itself (see _fold_gguf_dequant, _mark_torchao_q4_weights).
395395
redirect = v.meta.get("native_serialize_as")
396396
if redirect is not None:
397397
return TensorArg(name=redirect)
@@ -800,8 +800,8 @@ def _collect_fx_nodes(
800800
):
801801
continue
802802

803-
# Folded GGUF dequantize nodes are not emitted; the consuming op references
804-
# the packed weight directly (see _fold_gguf_dequant).
803+
# Folded dequantize nodes are not emitted; the consuming op references the
804+
# packed weight directly.
805805
if fx_node.meta.get("native_skip_serialize"):
806806
continue
807807

@@ -1054,12 +1054,35 @@ def _fold_gguf_dequant(graph_module: torch.fx.GraphModule) -> None:
10541054
node.meta["native_skip_serialize"] = True
10551055

10561056

1057-
def _mark_torchao_q4_weights(graph_module: torch.fx.GraphModule) -> None:
1057+
def _is_linear_weight_use(dequantize: torch.fx.Node, user: torch.fx.Node) -> bool:
1058+
op = _resolve_op_overload(user.target) if user.op == "call_function" else None
1059+
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:]
1066+
and dequantize not in user.kwargs.values()
1067+
)
1068+
1069+
1070+
def _mark_torchao_q4_weights(
1071+
graph_module: torch.fx.GraphModule, lifted_constants: set[str] | None
1072+
) -> None:
10581073
"""Mark portable torchao q4 weights for packed PTN storage.
10591074
1060-
The q/dq nodes remain in the serialized graph; only the weight storage
1061-
changes, to two four-bit values per byte.
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.
10621082
"""
1083+
if lifted_constants is None:
1084+
return
1085+
folds: dict[torch.fx.Node, list[tuple[torch.fx.Node, dict[str, object]]]] = {}
10631086
for node in graph_module.graph.nodes:
10641087
op = _resolve_op_overload(node.target) if node.op == "call_function" else None
10651088
if (
@@ -1105,12 +1128,35 @@ def _mark_torchao_q4_weights(graph_module: torch.fx.GraphModule) -> None:
11051128
"zero_point_dtype": _scalar_type(zero_point_value.dtype),
11061129
"group_size": block_size[1],
11071130
}
1108-
existing = weight.meta.get("native_packed_quant")
1109-
if existing is not None and existing != packed:
1110-
raise ValueError(
1111-
f"constant {weight.name!r} has incompatible q4 interpretations"
1112-
)
1131+
if (
1132+
{weight.name, scale.name, zero_point.name} <= lifted_constants
1133+
and output.dtype == scale_value.dtype
1134+
and node.users
1135+
and all(_is_linear_weight_use(node, user) for user in node.users)
1136+
):
1137+
folds.setdefault(weight, []).append((node, packed))
1138+
1139+
for weight, weight_folds in folds.items():
1140+
packed = weight_folds[0][1]
1141+
if len(weight.users) != len(weight_folds) or any(
1142+
p != packed for _, p in weight_folds
1143+
):
1144+
continue
11131145
weight.meta["native_packed_quant"] = packed
1146+
for node, _ in weight_folds:
1147+
node.meta["native_serialize_as"] = weight.name
1148+
node.meta["native_skip_serialize"] = True
1149+
1150+
1151+
def _lifted_constant_names(graph_signature: object | None) -> set[str] | None:
1152+
if graph_signature is None:
1153+
return None
1154+
return {
1155+
ispec.arg.name
1156+
for ispec in getattr(graph_signature, "input_specs", []) or []
1157+
if ispec.kind in _INPUT_KIND_MAP
1158+
and isinstance(getattr(ispec, "arg", None), TensorArgument)
1159+
}
11141160

11151161

11161162
def _build_graph_body(
@@ -1126,7 +1172,7 @@ def _build_graph_body(
11261172
dict[str, torch.Tensor],
11271173
]:
11281174
_fold_gguf_dequant(graph_module)
1129-
_mark_torchao_q4_weights(graph_module)
1175+
_mark_torchao_q4_weights(graph_module, _lifted_constant_names(graph_signature))
11301176
subgraph_map = _subgraph_map(graph_module)
11311177
nodes, val_by_name, output_names = _collect_fx_nodes(graph_module, subgraph_map)
11321178

‎backends/native/test/BUCK‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -59,7 +59,9 @@ fbcode_target(
5959
"//executorch/backends/native:lib",
6060
"//executorch/backends/native/test:utils",
6161
"//executorch/exir:lib",
62+
"//executorch/exir/native:native",
6263
"//executorch/extension/llm/export:gguf",
64+
"//pytorch/ao:torchao",
6365
],
6466
)
6567

‎backends/native/test/test_quant.py‎

Lines changed: 142 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -11,13 +11,25 @@
1111

1212
from executorch.backends.native.serialization import deserialize_program
1313
from executorch.backends.native.serialization.graph_serialize import _pack_signed_int4
14-
from executorch.backends.native.serialization.schema import PackedQuant, ScalarType
14+
from executorch.backends.native.serialization.schema import (
15+
AffineGroup,
16+
PackedQuant,
17+
ScalarType,
18+
TensorArg,
19+
)
1520
from executorch.backends.native.test.utils import (
1621
_call_function_targets,
1722
_get_delegate_blob,
1823
_lower,
1924
)
25+
from executorch.exir.native import to_native
2026
from executorch.extension.llm.export.gguf import ExportableGGUFTensor
27+
from torchao.quantization import (
28+
Int8DynamicActivationIntxWeightConfig,
29+
IntxWeightOnlyConfig,
30+
quantize_,
31+
)
32+
from torchao.quantization.granularity import PerGroup
2133

2234

2335
class FuseGGUFPassTest(unittest.TestCase):
@@ -114,3 +126,132 @@ def test_rejects_values_outside_signed_int4_range(self):
114126
with self.subTest(value=value):
115127
with self.assertRaisesRegex(ValueError, r"\[-8, 7\]"):
116128
_pack_signed_int4(torch.tensor([[value, 0]], dtype=torch.int8))
129+
130+
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"]
134+
assert isinstance(weight, TensorArg)
135+
return weight.name
136+
137+
138+
class TorchaoQ4LinearTest(unittest.TestCase):
139+
N, K, GROUP_SIZE = 16, 64, 32
140+
141+
def _serialize(self, config):
142+
model = nn.Sequential(nn.Linear(self.K, self.N, bias=False))
143+
quantize_(model, config)
144+
manager = to_native(torch.export.export(model, (torch.randn(2, self.K),)))
145+
return deserialize_program(manager._ptg).methods[0], manager._constants
146+
147+
def _assert_linear_reads_packed_weight(self, method, constants):
148+
[weight] = [c for c in method.constants if c.meta.quant is not None]
149+
scheme = weight.meta.quant.scheme
150+
self.assertIsInstance(scheme, AffineGroup)
151+
self.assertEqual(_linear_weight(method), weight.name)
152+
self.assertEqual(weight.meta.dtype, ScalarType.BYTE)
153+
self.assertEqual([d.max for d in weight.meta.sizes], [self.N, self.K])
154+
self.assertEqual(
155+
(scheme.quant_min, scheme.quant_max, scheme.group_size),
156+
(-8, 7, self.GROUP_SIZE),
157+
)
158+
self.assertEqual(scheme.scale_dtype, ScalarType.FLOAT)
159+
self.assertEqual(constants[weight.data_key].dtype, torch.uint8)
160+
self.assertEqual(constants[weight.data_key].numel(), self.N * self.K // 2)
161+
self.assertIn(scheme.scale_data_key, constants)
162+
self.assertIn(scheme.zero_point_data_key, constants)
163+
164+
def test_8da4w_linear_reads_packed_weight(self):
165+
method, constants = self._serialize(
166+
Int8DynamicActivationIntxWeightConfig(
167+
weight_dtype=torch.int4, weight_granularity=PerGroup(self.GROUP_SIZE)
168+
)
169+
)
170+
targets = _call_function_targets(method.graph)
171+
# Only the activation keeps its choose_qparams -> quantize -> dequantize.
172+
self.assertEqual(sum("dequantize_affine" in t for t in targets), 1)
173+
self.assertEqual(sum("quantize_affine" in t for t in targets), 2)
174+
self.assertTrue(any("choose_qparams_affine" in t for t in targets))
175+
self._assert_linear_reads_packed_weight(method, constants)
176+
177+
def test_4w_linear_reads_packed_weight(self):
178+
method, constants = self._serialize(
179+
IntxWeightOnlyConfig(
180+
weight_dtype=torch.int4, granularity=PerGroup(self.GROUP_SIZE)
181+
)
182+
)
183+
targets = _call_function_targets(method.graph)
184+
self.assertFalse(any("quantize_affine" in t for t in targets))
185+
self._assert_linear_reads_packed_weight(method, constants)
186+
187+
188+
class _DequantizedLinear(nn.Module):
189+
"""linear(x, dequantize_affine(weight)) over an int4 [8, 16] weight in groups
190+
of 8, optionally with the weight as an input or the dequantize read twice."""
191+
192+
def __init__(self, output_dtype=torch.float32, weight_input=False, reuse=False):
193+
super().__init__()
194+
self.output_dtype = output_dtype
195+
self.weight_input = weight_input
196+
self.reuse = reuse
197+
self.register_buffer("weight", torch.randint(-8, 8, (8, 16), dtype=torch.int8))
198+
self.register_buffer("scale", torch.rand(8, 2) + 0.5)
199+
self.register_buffer("zero_point", torch.zeros(8, 2, dtype=torch.int8))
200+
201+
def example_inputs(self):
202+
x = torch.randn(2, 16, dtype=self.output_dtype)
203+
return (x, self.weight.clone()) if self.weight_input else (x,)
204+
205+
def forward(self, x, weight=None):
206+
weight = self.weight if weight is None else weight
207+
dequantized = torch.ops.torchao.dequantize_affine(
208+
weight,
209+
[1, 8],
210+
self.scale,
211+
self.zero_point,
212+
torch.int8,
213+
-8,
214+
7,
215+
output_dtype=self.output_dtype,
216+
)
217+
out = torch.nn.functional.linear(x, dequantized)
218+
return out + dequantized.sum() if self.reuse else out
219+
220+
221+
class FoldTorchaoQ4DequantizeTest(unittest.TestCase):
222+
def _method(self, model):
223+
return deserialize_program(
224+
_get_delegate_blob(_lower(model, model.example_inputs()))
225+
).methods[0]
226+
227+
def _weight_dequantizes(self, model):
228+
targets = _call_function_targets(self._method(model).graph)
229+
return sum("dequantize_affine" in t for t in targets)
230+
231+
def _packed_constants(self, model):
232+
return [c for c in self._method(model).constants if c.meta.quant is not None]
233+
234+
def test_folds_constant_weight_read_by_linear(self):
235+
self.assertEqual(self._weight_dequantizes(_DequantizedLinear()), 0)
236+
self.assertEqual(len(self._packed_constants(_DequantizedLinear())), 1)
237+
238+
def test_keeps_unfolded_weight_unpacked(self):
239+
for model in (
240+
_DequantizedLinear(reuse=True),
241+
_DequantizedLinear(output_dtype=torch.float16),
242+
):
243+
self.assertEqual(self._packed_constants(model), [])
244+
245+
def test_keeps_dequantize_of_runtime_weight(self):
246+
self.assertEqual(
247+
self._weight_dequantizes(_DequantizedLinear(weight_input=True)), 1
248+
)
249+
250+
def test_keeps_dequantize_read_by_another_op(self):
251+
self.assertEqual(self._weight_dequantizes(_DequantizedLinear(reuse=True)), 1)
252+
253+
def test_keeps_dequantize_to_another_dtype_than_scales(self):
254+
self.assertEqual(
255+
self._weight_dequantizes(_DequantizedLinear(output_dtype=torch.float16)),
256+
1,
257+
)

0 commit comments

Comments
 (0)