|
11 | 11 |
|
12 | 12 | from executorch.backends.native.serialization import deserialize_program |
13 | 13 | 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 | +) |
15 | 20 | from executorch.backends.native.test.utils import ( |
16 | 21 | _call_function_targets, |
17 | 22 | _get_delegate_blob, |
18 | 23 | _lower, |
19 | 24 | ) |
| 25 | +from executorch.exir.native import to_native |
20 | 26 | 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 |
21 | 33 |
|
22 | 34 |
|
23 | 35 | class FuseGGUFPassTest(unittest.TestCase): |
@@ -114,3 +126,132 @@ def test_rejects_values_outside_signed_int4_range(self): |
114 | 126 | with self.subTest(value=value): |
115 | 127 | with self.assertRaisesRegex(ValueError, r"\[-8, 7\]"): |
116 | 128 | _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