litert-quantizer-nightly 0.10.0.dev20261008__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (211) hide show
  1. litert_quantizer/__init__.py +19 -0
  2. litert_quantizer/algorithm_manager.py +480 -0
  3. litert_quantizer/algorithm_manager_api.py +440 -0
  4. litert_quantizer/algorithm_manager_api_test.py +281 -0
  5. litert_quantizer/algorithms/__init__.py +15 -0
  6. litert_quantizer/algorithms/nonlinear_quantize/__init__.py +15 -0
  7. litert_quantizer/algorithms/nonlinear_quantize/float_casting.py +340 -0
  8. litert_quantizer/algorithms/nonlinear_quantize/float_casting_test.py +788 -0
  9. litert_quantizer/algorithms/uniform_quantize/__init__.py +15 -0
  10. litert_quantizer/algorithms/uniform_quantize/common_quantize.py +1497 -0
  11. litert_quantizer/algorithms/uniform_quantize/common_quantize_test.py +186 -0
  12. litert_quantizer/algorithms/uniform_quantize/dequantized_weight_recovery.py +362 -0
  13. litert_quantizer/algorithms/uniform_quantize/dequantized_weight_recovery_test.py +449 -0
  14. litert_quantizer/algorithms/uniform_quantize/gptq.py +303 -0
  15. litert_quantizer/algorithms/uniform_quantize/gptq_test.py +402 -0
  16. litert_quantizer/algorithms/uniform_quantize/hadamard_rotation.py +500 -0
  17. litert_quantizer/algorithms/uniform_quantize/hadamard_rotation_test.py +490 -0
  18. litert_quantizer/algorithms/uniform_quantize/mse.py +128 -0
  19. litert_quantizer/algorithms/uniform_quantize/mse_test.py +195 -0
  20. litert_quantizer/algorithms/uniform_quantize/naive_min_max_quantize.py +229 -0
  21. litert_quantizer/algorithms/uniform_quantize/naive_min_max_quantize_test.py +346 -0
  22. litert_quantizer/algorithms/uniform_quantize/octav.py +230 -0
  23. litert_quantizer/algorithms/uniform_quantize/octav_test.py +240 -0
  24. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/add_test.py +166 -0
  25. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/average_pool_2d_test.py +119 -0
  26. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/batch_matmul_test.py +307 -0
  27. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/broadcast_to_test.py +101 -0
  28. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/concatenation_test.py +117 -0
  29. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/conv2d_test.py +214 -0
  30. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/conv2d_transpose_test.py +222 -0
  31. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/depthwise_conv2d_test.py +209 -0
  32. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/div_test.py +99 -0
  33. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/dynamic_update_slice_test.py +113 -0
  34. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/embedding_lookup_test.py +128 -0
  35. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/equal_test.py +100 -0
  36. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/fully_connected_test.py +167 -0
  37. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/gather_nd_test.py +168 -0
  38. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/gather_test.py +101 -0
  39. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/gelu_test.py +115 -0
  40. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/hard_swish_test.py +96 -0
  41. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/input_output_test.py +307 -0
  42. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/logistic_test.py +115 -0
  43. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/max_pool_2d_test.py +102 -0
  44. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/maximum_test.py +100 -0
  45. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/mean_test.py +117 -0
  46. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/mirror_pad_test.py +101 -0
  47. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/mul_test.py +161 -0
  48. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/not_equal_test.py +100 -0
  49. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/pack_test.py +100 -0
  50. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/pad_test.py +104 -0
  51. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/padv2_test.py +106 -0
  52. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/reduce_min_test.py +101 -0
  53. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/relu_test.py +98 -0
  54. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/reshape_test.py +118 -0
  55. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/resize_bilinear_test.py +104 -0
  56. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/resize_nearest_neighbor_test.py +102 -0
  57. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/rsqrt_test.py +108 -0
  58. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/select_test.py +103 -0
  59. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/select_v2_test.py +113 -0
  60. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/slice_test.py +113 -0
  61. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/softmax_test.py +115 -0
  62. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/space_to_depth_test.py +97 -0
  63. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/split_test.py +119 -0
  64. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/sqrt_test.py +98 -0
  65. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/squared_difference_test.py +102 -0
  66. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/strided_slice_test.py +117 -0
  67. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/sub_test.py +159 -0
  68. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/sum_test.py +110 -0
  69. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/tanh_test.py +115 -0
  70. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/test_utils.py +698 -0
  71. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/transpose_test.py +118 -0
  72. litert_quantizer/algorithms/uniform_quantize/op_architecture_tests/unpack_test.py +100 -0
  73. litert_quantizer/algorithms/uniform_quantize/oscar.py +682 -0
  74. litert_quantizer/algorithms/uniform_quantize/oscar_test.py +547 -0
  75. litert_quantizer/algorithms/uniform_quantize/uniform_quantize_tensor.py +645 -0
  76. litert_quantizer/algorithms/uniform_quantize/uniform_quantize_tensor_test.py +519 -0
  77. litert_quantizer/algorithms/utils/__init__.py +15 -0
  78. litert_quantizer/algorithms/utils/common_utils.py +1304 -0
  79. litert_quantizer/algorithms/utils/common_utils_test.py +660 -0
  80. litert_quantizer/calibrator.py +685 -0
  81. litert_quantizer/calibrator_test.py +699 -0
  82. litert_quantizer/conftest.py +22 -0
  83. litert_quantizer/default_policy.py +440 -0
  84. litert_quantizer/examples/mnist/quantize_toy_model.py +613 -0
  85. litert_quantizer/litert_quantizer.py +297 -0
  86. litert_quantizer/litert_quantizer_test.py +182 -0
  87. litert_quantizer/model_modifier.py +394 -0
  88. litert_quantizer/model_modifier_test.py +324 -0
  89. litert_quantizer/model_validator.py +447 -0
  90. litert_quantizer/model_validator_test.py +438 -0
  91. litert_quantizer/params_generator.py +565 -0
  92. litert_quantizer/params_generator_test.py +1166 -0
  93. litert_quantizer/policies/dummy_config_policy.json +17 -0
  94. litert_quantizer/policies/example_config_policy.json +174 -0
  95. litert_quantizer/qtyping.py +709 -0
  96. litert_quantizer/qtyping_test.py +238 -0
  97. litert_quantizer/quantizer.py +620 -0
  98. litert_quantizer/quantizer_test.py +892 -0
  99. litert_quantizer/recipe.py +397 -0
  100. litert_quantizer/recipe_manager.py +414 -0
  101. litert_quantizer/recipe_manager_test.py +1027 -0
  102. litert_quantizer/recipe_test.py +180 -0
  103. litert_quantizer/recipes/default_a16w8_recipe.json +25 -0
  104. litert_quantizer/recipes/default_a8w8_recipe.json +25 -0
  105. litert_quantizer/recipes/default_af32w4float_recipe.json +19 -0
  106. litert_quantizer/recipes/default_af32w8float_recipe.json +19 -0
  107. litert_quantizer/recipes/dynamic_legacy_wi8_afp32_recipe.json +20 -0
  108. litert_quantizer/recipes/dynamic_wi4_afp32_hadamard_recipe.json +39 -0
  109. litert_quantizer/recipes/dynamic_wi8_afp32_hadamard_recipe.json +39 -0
  110. litert_quantizer/recipes/dynamic_wi8_afp32_litertlm_recipe.json +3 -0
  111. litert_quantizer/recipes/dynamic_wi8_afp32_recipe.json +19 -0
  112. litert_quantizer/recipes/sample_advanced_usage_recipe.json +47 -0
  113. litert_quantizer/tests/end_to_end_tests/add_test.py +141 -0
  114. litert_quantizer/tests/end_to_end_tests/batch_matmul_test.py +113 -0
  115. litert_quantizer/tests/end_to_end_tests/broadcast_to_test.py +69 -0
  116. litert_quantizer/tests/end_to_end_tests/composite_test.py +148 -0
  117. litert_quantizer/tests/end_to_end_tests/concatenation_test.py +96 -0
  118. litert_quantizer/tests/end_to_end_tests/conv2d_transpose_test.py +199 -0
  119. litert_quantizer/tests/end_to_end_tests/depthwise_conv2d_test.py +198 -0
  120. litert_quantizer/tests/end_to_end_tests/div_test.py +69 -0
  121. litert_quantizer/tests/end_to_end_tests/dynamic_update_slice_test.py +126 -0
  122. litert_quantizer/tests/end_to_end_tests/embedding_lookup_test.py +153 -0
  123. litert_quantizer/tests/end_to_end_tests/equal_test.py +68 -0
  124. litert_quantizer/tests/end_to_end_tests/fully_connected_test.py +473 -0
  125. litert_quantizer/tests/end_to_end_tests/gather_nd_test.py +96 -0
  126. litert_quantizer/tests/end_to_end_tests/gather_test.py +67 -0
  127. litert_quantizer/tests/end_to_end_tests/gelu_test.py +99 -0
  128. litert_quantizer/tests/end_to_end_tests/hard_swish_test.py +67 -0
  129. litert_quantizer/tests/end_to_end_tests/input_output_test.py +273 -0
  130. litert_quantizer/tests/end_to_end_tests/logistic_test.py +94 -0
  131. litert_quantizer/tests/end_to_end_tests/max_pool_2d_test.py +69 -0
  132. litert_quantizer/tests/end_to_end_tests/maximum_test.py +69 -0
  133. litert_quantizer/tests/end_to_end_tests/mean_test.py +96 -0
  134. litert_quantizer/tests/end_to_end_tests/mirror_pad_test.py +67 -0
  135. litert_quantizer/tests/end_to_end_tests/mul_test.py +177 -0
  136. litert_quantizer/tests/end_to_end_tests/not_equal_test.py +67 -0
  137. litert_quantizer/tests/end_to_end_tests/pack_test.py +68 -0
  138. litert_quantizer/tests/end_to_end_tests/pad_test.py +74 -0
  139. litert_quantizer/tests/end_to_end_tests/padv2_test.py +67 -0
  140. litert_quantizer/tests/end_to_end_tests/reduce_min_test.py +70 -0
  141. litert_quantizer/tests/end_to_end_tests/relu_test.py +67 -0
  142. litert_quantizer/tests/end_to_end_tests/resize_bilinear_test.py +69 -0
  143. litert_quantizer/tests/end_to_end_tests/resize_nearest_neighbor_test.py +70 -0
  144. litert_quantizer/tests/end_to_end_tests/rsqrt_test.py +98 -0
  145. litert_quantizer/tests/end_to_end_tests/select_test.py +69 -0
  146. litert_quantizer/tests/end_to_end_tests/select_v2_test.py +123 -0
  147. litert_quantizer/tests/end_to_end_tests/slice_test.py +123 -0
  148. litert_quantizer/tests/end_to_end_tests/space_to_depth_test.py +65 -0
  149. litert_quantizer/tests/end_to_end_tests/split_test.py +98 -0
  150. litert_quantizer/tests/end_to_end_tests/sqrt_test.py +70 -0
  151. litert_quantizer/tests/end_to_end_tests/squared_difference_test.py +75 -0
  152. litert_quantizer/tests/end_to_end_tests/strided_slice_test.py +96 -0
  153. litert_quantizer/tests/end_to_end_tests/sub_test.py +141 -0
  154. litert_quantizer/tests/end_to_end_tests/sum_test.py +129 -0
  155. litert_quantizer/tests/end_to_end_tests/tanh_test.py +94 -0
  156. litert_quantizer/tests/end_to_end_tests/transpose_test.py +162 -0
  157. litert_quantizer/tests/end_to_end_tests/unpack_test.py +70 -0
  158. litert_quantizer/tests/mnist_test.py +280 -0
  159. litert_quantizer/tests/padv2_inf_max_pool_2d_test.py +68 -0
  160. litert_quantizer/tests/shared_buffer_test.py +447 -0
  161. litert_quantizer/transformation_instruction_generator.py +862 -0
  162. litert_quantizer/transformation_instruction_generator_test.py +1552 -0
  163. litert_quantizer/transformation_performer.py +350 -0
  164. litert_quantizer/transformation_performer_test.py +521 -0
  165. litert_quantizer/transformations/__init__.py +15 -0
  166. litert_quantizer/transformations/dequant_insert.py +86 -0
  167. litert_quantizer/transformations/dequant_insert_test.py +292 -0
  168. litert_quantizer/transformations/duplicate_buffer.py +46 -0
  169. litert_quantizer/transformations/duplicate_buffer_test.py +108 -0
  170. litert_quantizer/transformations/duplicate_tensor.py +62 -0
  171. litert_quantizer/transformations/duplicate_tensor_test.py +133 -0
  172. litert_quantizer/transformations/insert_decomposed_hadamard_rotation.py +265 -0
  173. litert_quantizer/transformations/insert_decomposed_hadamard_rotation_test.py +242 -0
  174. litert_quantizer/transformations/insert_hadamard_rotation.py +156 -0
  175. litert_quantizer/transformations/insert_hadamard_rotation_test.py +196 -0
  176. litert_quantizer/transformations/insert_multiply.py +129 -0
  177. litert_quantizer/transformations/insert_multiply_test.py +196 -0
  178. litert_quantizer/transformations/quant_insert.py +98 -0
  179. litert_quantizer/transformations/quant_insert_test.py +272 -0
  180. litert_quantizer/transformations/quantize_tensor.py +227 -0
  181. litert_quantizer/transformations/quantize_tensor_test.py +387 -0
  182. litert_quantizer/transformations/transformation_utils.py +419 -0
  183. litert_quantizer/transformations/transformation_utils_test.py +278 -0
  184. litert_quantizer/utils/__init__.py +15 -0
  185. litert_quantizer/utils/calibration_utils.py +372 -0
  186. litert_quantizer/utils/calibration_utils_test.py +218 -0
  187. litert_quantizer/utils/constrained_ops_utils.py +112 -0
  188. litert_quantizer/utils/constrained_ops_utils_test.py +56 -0
  189. litert_quantizer/utils/histogram_utils.py +480 -0
  190. litert_quantizer/utils/histogram_utils_test.py +299 -0
  191. litert_quantizer/utils/litertlm_utils.py +283 -0
  192. litert_quantizer/utils/litertlm_utils_test.py +91 -0
  193. litert_quantizer/utils/progress_utils.py +168 -0
  194. litert_quantizer/utils/progress_utils_test.py +134 -0
  195. litert_quantizer/utils/qsv_utils.py +171 -0
  196. litert_quantizer/utils/qsv_utils_test.py +245 -0
  197. litert_quantizer/utils/recipe_utils.py +248 -0
  198. litert_quantizer/utils/recipe_utils_test.py +126 -0
  199. litert_quantizer/utils/test_utils.py +248 -0
  200. litert_quantizer/utils/tfl_flatbuffer_utils.py +417 -0
  201. litert_quantizer/utils/tfl_flatbuffer_utils_test.py +228 -0
  202. litert_quantizer/utils/tfl_interpreter_utils.py +475 -0
  203. litert_quantizer/utils/tfl_interpreter_utils_test.py +436 -0
  204. litert_quantizer/utils/validation_utils.py +255 -0
  205. litert_quantizer/utils/validation_utils_test.py +173 -0
  206. litert_quantizer_nightly-0.10.0.dev20261008.dist-info/METADATA +530 -0
  207. litert_quantizer_nightly-0.10.0.dev20261008.dist-info/RECORD +211 -0
  208. litert_quantizer_nightly-0.10.0.dev20261008.dist-info/WHEEL +5 -0
  209. litert_quantizer_nightly-0.10.0.dev20261008.dist-info/entry_points.txt +2 -0
  210. litert_quantizer_nightly-0.10.0.dev20261008.dist-info/licenses/LICENSE +201 -0
  211. litert_quantizer_nightly-0.10.0.dev20261008.dist-info/top_level.txt +1 -0
@@ -0,0 +1,19 @@
1
+ # Copyright 2024 Google LLC
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+ # ==============================================================================
15
+
16
+ """Init file for the AI Edge quantizer package."""
17
+
18
+ # pylint: disable=unused-import
19
+ from .quantizer import Quantizer
@@ -0,0 +1,480 @@
1
+ # Copyright 2024 Google LLC
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+ # ==============================================================================
15
+
16
+ """Quantizer Algorithm Manager Interface."""
17
+
18
+ import enum
19
+ import functools
20
+ import immutabledict
21
+ from litert_quantizer import algorithm_manager_api
22
+ from litert_quantizer import default_policy
23
+ from litert_quantizer import qtyping
24
+ from litert_quantizer.algorithms.nonlinear_quantize import float_casting
25
+ from litert_quantizer.algorithms.uniform_quantize import common_quantize
26
+ from litert_quantizer.algorithms.uniform_quantize import dequantized_weight_recovery
27
+ from litert_quantizer.algorithms.uniform_quantize import gptq
28
+ from litert_quantizer.algorithms.uniform_quantize import hadamard_rotation
29
+ from litert_quantizer.algorithms.uniform_quantize import mse
30
+ from litert_quantizer.algorithms.uniform_quantize import naive_min_max_quantize
31
+ from litert_quantizer.algorithms.uniform_quantize import octav
32
+ from litert_quantizer.algorithms.uniform_quantize import oscar
33
+ from litert_quantizer.utils import qsv_utils
34
+
35
+
36
+ immutabledict = immutabledict.immutabledict
37
+
38
+
39
+ # TODO: b/399775701 - Clean up this file.
40
+
41
+ _TFLOpName = qtyping.TFLOperationName
42
+
43
+ _alg_manager_instance = algorithm_manager_api.AlgorithmManagerApi()
44
+
45
+ # Expose instance functions.
46
+ get_quantization_func = _alg_manager_instance.get_quantization_func
47
+ get_supported_ops = _alg_manager_instance.get_supported_ops
48
+ get_update_qsv_func = _alg_manager_instance.get_update_qsv_func
49
+ get_init_qsv_func = _alg_manager_instance.get_init_qsv_func
50
+ register_op_quant_config_validation_func = (
51
+ _alg_manager_instance.register_op_quant_config_validation_func
52
+ )
53
+ register_config_check_policy_func = (
54
+ _alg_manager_instance.register_config_check_policy
55
+ )
56
+ register_quantized_op = _alg_manager_instance.register_quantized_op
57
+ is_op_registered = _alg_manager_instance.is_op_registered
58
+ is_algorithm_registered = _alg_manager_instance.is_algorithm_registered
59
+ check_op_quantization_config = (
60
+ _alg_manager_instance.check_op_quantization_config
61
+ )
62
+
63
+
64
+ # Quantization algorithms.
65
+ class AlgorithmName(str, enum.Enum):
66
+ NO_QUANTIZE = "no_quantize"
67
+ MIN_MAX_UNIFORM_QUANT = naive_min_max_quantize.ALGORITHM_KEY
68
+ FLOAT_CASTING = float_casting.ALGORITHM_KEY
69
+ DEQUANTIZED_WEIGHT_RECOVERY = dequantized_weight_recovery.ALGORITHM_KEY
70
+ OCTAV = octav.ALGORITHM_KEY
71
+ HADAMARD_ROTATION = hadamard_rotation.CUSTOM_OP_ALGORITHM_KEY
72
+ DECOMPOSED_HADAMARD_ROTATION = hadamard_rotation.DECOMPOSED_ALGORITHM_KEY
73
+ MSE = mse.ALGORITHM_KEY
74
+ GPTQ = gptq.ALGORITHM_KEY
75
+ OSCAR = oscar.ALGORITHM_KEY
76
+
77
+
78
+ ### MIN/MAX_UNIFORM_QUANT ###
79
+
80
+ # Register MIN_MAX_UNIFORM_QUANT algorithm.
81
+ register_op_quant_config_validation_func(
82
+ AlgorithmName.MIN_MAX_UNIFORM_QUANT,
83
+ common_quantize.check_op_quantization_config,
84
+ )
85
+
86
+ # Register a config check policy for MIN_MAX_UNIFORM_QUANT algorithm.
87
+ register_config_check_policy_func(
88
+ AlgorithmName.MIN_MAX_UNIFORM_QUANT,
89
+ default_policy.DEFAULT_CONFIG_CHECK_POLICY,
90
+ )
91
+
92
+ MIN_MAX_OP_NAME_MATERIALIZE_FUNC_DICT = {
93
+ _TFLOpName.INPUT: common_quantize.materialize_input,
94
+ _TFLOpName.OUTPUT: common_quantize.materialize_output,
95
+ _TFLOpName.FULLY_CONNECTED: common_quantize.materialize_fc_conv,
96
+ _TFLOpName.BATCH_MATMUL: common_quantize.materialize_batch_matmul,
97
+ _TFLOpName.CONV_2D: common_quantize.materialize_fc_conv,
98
+ _TFLOpName.DEPTHWISE_CONV_2D: common_quantize.materialize_fc_conv,
99
+ _TFLOpName.CONV_2D_TRANSPOSE: common_quantize.materialize_conv2d_transpose,
100
+ _TFLOpName.RESHAPE: common_quantize.materialize_reshape,
101
+ _TFLOpName.AVERAGE_POOL_2D: common_quantize.materialize_average_pool_2d,
102
+ _TFLOpName.EMBEDDING_LOOKUP: common_quantize.materialize_embedding_lookup,
103
+ _TFLOpName.SOFTMAX: common_quantize.materialize_softmax_and_logistic,
104
+ _TFLOpName.TANH: common_quantize.materialize_tanh,
105
+ _TFLOpName.TRANSPOSE: common_quantize.materialize_transpose,
106
+ _TFLOpName.GELU: common_quantize.materialize_gelu,
107
+ _TFLOpName.ADD: common_quantize.materialize_add,
108
+ _TFLOpName.SUB: common_quantize.materialize_sub,
109
+ _TFLOpName.MUL: common_quantize.materialize_mul,
110
+ _TFLOpName.MEAN: common_quantize.materialize_mean,
111
+ _TFLOpName.RSQRT: common_quantize.materialize_rsqrt,
112
+ _TFLOpName.CONCATENATION: common_quantize.materialize_concatenation,
113
+ _TFLOpName.STRIDED_SLICE: common_quantize.materialize_strided_slice,
114
+ _TFLOpName.SPLIT: common_quantize.materialize_split,
115
+ _TFLOpName.LOGISTIC: common_quantize.materialize_softmax_and_logistic,
116
+ _TFLOpName.SLICE: common_quantize.materialize_slice,
117
+ _TFLOpName.SUM: common_quantize.materialize_sum,
118
+ _TFLOpName.SELECT: common_quantize.materialize_select,
119
+ _TFLOpName.SELECT_V2: common_quantize.materialize_select_v2,
120
+ _TFLOpName.DYNAMIC_UPDATE_SLICE: (
121
+ common_quantize.materialize_dynamic_update_slice
122
+ ),
123
+ _TFLOpName.STABLEHLO_COMPOSITE: common_quantize.materialize_composite,
124
+ _TFLOpName.PAD: common_quantize.materialize_pad,
125
+ _TFLOpName.SQUARED_DIFFERENCE: (
126
+ common_quantize.materialize_squared_difference
127
+ ),
128
+ _TFLOpName.MAX_POOL_2D: common_quantize.materialize_max_pool_2d,
129
+ _TFLOpName.RESIZE_BILINEAR: common_quantize.materialize_resize_bilinear,
130
+ _TFLOpName.RESIZE_NEAREST_NEIGHBOR: (
131
+ common_quantize.materialize_resize_nearest_neighbor
132
+ ),
133
+ _TFLOpName.GATHER_ND: common_quantize.materialize_gather_nd,
134
+ _TFLOpName.PACK: common_quantize.materialize_pack,
135
+ _TFLOpName.UNPACK: common_quantize.materialize_unpack,
136
+ _TFLOpName.DIV: common_quantize.materialize_div,
137
+ _TFLOpName.BROADCAST_TO: common_quantize.materialize_broadcast_to,
138
+ _TFLOpName.SQRT: common_quantize.materialize_sqrt,
139
+ _TFLOpName.GATHER: common_quantize.materialize_gather,
140
+ _TFLOpName.HARD_SWISH: common_quantize.materialize_hard_swish,
141
+ _TFLOpName.MAXIMUM: common_quantize.materialize_maximum,
142
+ _TFLOpName.PADV2: common_quantize.materialize_padv2,
143
+ _TFLOpName.REDUCE_MIN: common_quantize.materialize_reduce_min,
144
+ _TFLOpName.EQUAL: common_quantize.materialize_equal,
145
+ _TFLOpName.NOT_EQUAL: common_quantize.materialize_not_equal,
146
+ _TFLOpName.MIRROR_PAD: common_quantize.materialize_mirror_pad,
147
+ _TFLOpName.SPACE_TO_DEPTH: common_quantize.materialize_space_to_depth,
148
+ _TFLOpName.RELU: common_quantize.materialize_relu,
149
+ }
150
+ for op_name, materialize_func in MIN_MAX_OP_NAME_MATERIALIZE_FUNC_DICT.items():
151
+ register_quantized_op(
152
+ AlgorithmName.MIN_MAX_UNIFORM_QUANT,
153
+ op_name,
154
+ naive_min_max_quantize.init_qsvs,
155
+ calibration_func=naive_min_max_quantize.min_max_calibrate,
156
+ # Most of the materialize op functions are common for all algorithms
157
+ # except for the function to get scale and zero point, i.e.,
158
+ # get_tensor_quant_params. So we use functools.partial here to pass in the
159
+ # common utility function and thealgorithm-specific function.
160
+ materialize_func=functools.partial(
161
+ materialize_func,
162
+ naive_min_max_quantize.get_tensor_quant_params,
163
+ ),
164
+ )
165
+
166
+ ### FLOAT_CASTING ###
167
+ register_op_quant_config_validation_func(
168
+ AlgorithmName.FLOAT_CASTING,
169
+ float_casting.check_op_quantization_config,
170
+ )
171
+
172
+ # Register a config check policy for FLOAT_CASTING algorithm.
173
+ # TODO: b/353780772 - Replace an empty policy for FLOAT_CASTING algorithm.
174
+ register_config_check_policy_func(
175
+ AlgorithmName.FLOAT_CASTING, qtyping.ConfigCheckPolicyDict()
176
+ )
177
+
178
+ for op_name, materialize_func in zip(
179
+ (
180
+ _TFLOpName.FULLY_CONNECTED,
181
+ _TFLOpName.CONV_2D,
182
+ _TFLOpName.DEPTHWISE_CONV_2D,
183
+ _TFLOpName.CONV_2D_TRANSPOSE,
184
+ _TFLOpName.EMBEDDING_LOOKUP,
185
+ ),
186
+ (
187
+ float_casting.materialize_fc_conv,
188
+ float_casting.materialize_fc_conv,
189
+ float_casting.materialize_fc_conv,
190
+ float_casting.materialize_conv2d_transpose,
191
+ float_casting.materialize_embedding_lookup,
192
+ ),
193
+ ):
194
+ register_quantized_op(
195
+ AlgorithmName.FLOAT_CASTING,
196
+ op_name,
197
+ float_casting.init_qsvs, # pyrefly: ignore[bad-argument-type]
198
+ calibration_func=float_casting.calibrate, # pyrefly: ignore[bad-argument-type]
199
+ materialize_func=materialize_func, # pyrefly: ignore[bad-argument-type]
200
+ )
201
+
202
+ ### DEQUANTIZED_WEIGHT_RECOVERY ###
203
+ register_op_quant_config_validation_func(
204
+ AlgorithmName.DEQUANTIZED_WEIGHT_RECOVERY,
205
+ common_quantize.check_op_quantization_config,
206
+ )
207
+
208
+ register_config_check_policy_func(
209
+ AlgorithmName.DEQUANTIZED_WEIGHT_RECOVERY,
210
+ default_policy.DEFAULT_CONFIG_CHECK_POLICY,
211
+ )
212
+
213
+ DEQUANTIZED_WEIGHT_RECOVERY_OP_NAME_MATERIALIZE_FUNC_DICT = {
214
+ _TFLOpName.FULLY_CONNECTED: common_quantize.materialize_fc_conv,
215
+ _TFLOpName.CONV_2D: common_quantize.materialize_fc_conv,
216
+ _TFLOpName.EMBEDDING_LOOKUP: common_quantize.materialize_embedding_lookup,
217
+ }
218
+
219
+ for (
220
+ op_name,
221
+ materialize_func,
222
+ ) in DEQUANTIZED_WEIGHT_RECOVERY_OP_NAME_MATERIALIZE_FUNC_DICT.items():
223
+ register_quantized_op(
224
+ algorithm_key=AlgorithmName.DEQUANTIZED_WEIGHT_RECOVERY,
225
+ tfl_op_name=op_name,
226
+ init_qsv_func=dequantized_weight_recovery.init_qsvs, # pyrefly: ignore[bad-argument-type]
227
+ calibration_func=dequantized_weight_recovery.calibrate, # pyrefly: ignore[bad-argument-type]
228
+ # Most of the materialize op functions are common for all algorithms
229
+ # except for the function to get scale and zero point, i.e.,
230
+ # get_tensor_quant_params. So we use functools.partial here to pass in the
231
+ # common utility function and the algorithm-specific function.
232
+ materialize_func=functools.partial(
233
+ materialize_func,
234
+ dequantized_weight_recovery.get_tensor_quant_params,
235
+ ),
236
+ )
237
+
238
+
239
+ # Register OCTAV algorithm.
240
+ register_op_quant_config_validation_func(
241
+ AlgorithmName.OCTAV,
242
+ common_quantize.check_op_quantization_config,
243
+ )
244
+
245
+ # Register a config check policy for OCTAV algorithm.
246
+ register_config_check_policy_func(
247
+ AlgorithmName.OCTAV,
248
+ default_policy.DEFAULT_CONFIG_CHECK_POLICY,
249
+ )
250
+
251
+ _OCTAV_OP_NAME_MATERIALIZE_FUNC_DICT = immutabledict({
252
+ _TFLOpName.INPUT: common_quantize.materialize_input,
253
+ _TFLOpName.OUTPUT: common_quantize.materialize_output,
254
+ _TFLOpName.FULLY_CONNECTED: common_quantize.materialize_fc_conv,
255
+ _TFLOpName.BATCH_MATMUL: common_quantize.materialize_batch_matmul,
256
+ _TFLOpName.CONV_2D: common_quantize.materialize_fc_conv,
257
+ _TFLOpName.DEPTHWISE_CONV_2D: common_quantize.materialize_fc_conv,
258
+ _TFLOpName.CONV_2D_TRANSPOSE: common_quantize.materialize_conv2d_transpose,
259
+ _TFLOpName.RESHAPE: common_quantize.materialize_reshape,
260
+ _TFLOpName.AVERAGE_POOL_2D: common_quantize.materialize_average_pool_2d,
261
+ _TFLOpName.EMBEDDING_LOOKUP: common_quantize.materialize_embedding_lookup,
262
+ _TFLOpName.SOFTMAX: common_quantize.materialize_softmax_and_logistic,
263
+ _TFLOpName.TANH: common_quantize.materialize_tanh,
264
+ _TFLOpName.TRANSPOSE: common_quantize.materialize_transpose,
265
+ _TFLOpName.GELU: common_quantize.materialize_gelu,
266
+ _TFLOpName.ADD: common_quantize.materialize_add,
267
+ _TFLOpName.SUB: common_quantize.materialize_sub,
268
+ _TFLOpName.MUL: common_quantize.materialize_mul,
269
+ _TFLOpName.MEAN: common_quantize.materialize_mean,
270
+ _TFLOpName.RSQRT: common_quantize.materialize_rsqrt,
271
+ _TFLOpName.CONCATENATION: common_quantize.materialize_concatenation,
272
+ _TFLOpName.STRIDED_SLICE: common_quantize.materialize_strided_slice,
273
+ _TFLOpName.SPLIT: common_quantize.materialize_split,
274
+ _TFLOpName.LOGISTIC: common_quantize.materialize_softmax_and_logistic,
275
+ _TFLOpName.SLICE: common_quantize.materialize_slice,
276
+ _TFLOpName.SUM: common_quantize.materialize_sum,
277
+ _TFLOpName.SELECT: common_quantize.materialize_select,
278
+ _TFLOpName.SELECT_V2: common_quantize.materialize_select_v2,
279
+ _TFLOpName.DYNAMIC_UPDATE_SLICE: (
280
+ common_quantize.materialize_dynamic_update_slice
281
+ ),
282
+ _TFLOpName.STABLEHLO_COMPOSITE: common_quantize.materialize_composite,
283
+ _TFLOpName.PAD: common_quantize.materialize_pad,
284
+ _TFLOpName.SQUARED_DIFFERENCE: (
285
+ common_quantize.materialize_squared_difference
286
+ ),
287
+ _TFLOpName.MAX_POOL_2D: common_quantize.materialize_max_pool_2d,
288
+ _TFLOpName.RESIZE_BILINEAR: common_quantize.materialize_resize_bilinear,
289
+ _TFLOpName.RESIZE_NEAREST_NEIGHBOR: (
290
+ common_quantize.materialize_resize_nearest_neighbor
291
+ ),
292
+ _TFLOpName.GATHER_ND: common_quantize.materialize_gather_nd,
293
+ _TFLOpName.PACK: common_quantize.materialize_pack,
294
+ _TFLOpName.UNPACK: common_quantize.materialize_unpack,
295
+ _TFLOpName.DIV: common_quantize.materialize_div,
296
+ _TFLOpName.BROADCAST_TO: common_quantize.materialize_broadcast_to,
297
+ _TFLOpName.SQRT: common_quantize.materialize_sqrt,
298
+ _TFLOpName.GATHER: common_quantize.materialize_gather,
299
+ _TFLOpName.HARD_SWISH: common_quantize.materialize_hard_swish,
300
+ _TFLOpName.MAXIMUM: common_quantize.materialize_maximum,
301
+ _TFLOpName.PADV2: common_quantize.materialize_padv2,
302
+ _TFLOpName.REDUCE_MIN: common_quantize.materialize_reduce_min,
303
+ _TFLOpName.EQUAL: common_quantize.materialize_equal,
304
+ _TFLOpName.NOT_EQUAL: common_quantize.materialize_not_equal,
305
+ _TFLOpName.MIRROR_PAD: common_quantize.materialize_mirror_pad,
306
+ _TFLOpName.SPACE_TO_DEPTH: common_quantize.materialize_space_to_depth,
307
+ _TFLOpName.RELU: common_quantize.materialize_relu,
308
+ })
309
+
310
+ for op_name, materialize_func in _OCTAV_OP_NAME_MATERIALIZE_FUNC_DICT.items():
311
+ register_quantized_op(
312
+ AlgorithmName.OCTAV,
313
+ op_name,
314
+ naive_min_max_quantize.init_qsvs,
315
+ calibration_func=naive_min_max_quantize.min_max_calibrate,
316
+ materialize_func=functools.partial(
317
+ materialize_func,
318
+ octav.get_tensor_quant_params,
319
+ ),
320
+ )
321
+
322
+ # Register the Hadamard Rotation algorithm.
323
+ register_op_quant_config_validation_func(
324
+ AlgorithmName.HADAMARD_ROTATION,
325
+ common_quantize.check_op_quantization_config,
326
+ )
327
+
328
+ # Register a config check policy for the Hadamard Rotation algorithm.
329
+ register_config_check_policy_func(
330
+ AlgorithmName.HADAMARD_ROTATION,
331
+ default_policy.DEFAULT_CONFIG_CHECK_POLICY,
332
+ )
333
+
334
+ # Register specialized hadamard rotation materialize functions.
335
+ _HADAMARD_ROTATION_OP_NAME_MATERIALIZE_FUNC_DICT = immutabledict({
336
+ _TFLOpName.FULLY_CONNECTED: (
337
+ hadamard_rotation.materialize_fully_connected_custom_op
338
+ ),
339
+ _TFLOpName.EMBEDDING_LOOKUP: (
340
+ hadamard_rotation.materialize_embedding_lookup_custom_op
341
+ ),
342
+ })
343
+ for (
344
+ op_name,
345
+ materialize_func,
346
+ ) in _HADAMARD_ROTATION_OP_NAME_MATERIALIZE_FUNC_DICT.items():
347
+ register_quantized_op(
348
+ AlgorithmName.HADAMARD_ROTATION,
349
+ op_name,
350
+ naive_min_max_quantize.init_qsvs,
351
+ calibration_func=naive_min_max_quantize.min_max_calibrate,
352
+ materialize_func=materialize_func, # pyrefly: ignore[bad-argument-type]
353
+ )
354
+
355
+ register_op_quant_config_validation_func(
356
+ AlgorithmName.DECOMPOSED_HADAMARD_ROTATION,
357
+ common_quantize.check_op_quantization_config,
358
+ )
359
+
360
+ register_config_check_policy_func(
361
+ AlgorithmName.DECOMPOSED_HADAMARD_ROTATION,
362
+ default_policy.DEFAULT_CONFIG_CHECK_POLICY,
363
+ )
364
+
365
+ _DECOMPOSED_HADAMARD_ROTATION_OP_NAME_MATERIALIZE_FUNC_DICT = immutabledict({
366
+ _TFLOpName.FULLY_CONNECTED: (
367
+ hadamard_rotation.materialize_fully_connected_decomposed
368
+ ),
369
+ _TFLOpName.EMBEDDING_LOOKUP: (
370
+ hadamard_rotation.materialize_embedding_lookup_decomposed
371
+ ),
372
+ })
373
+ for (
374
+ op_name,
375
+ materialize_func,
376
+ ) in _DECOMPOSED_HADAMARD_ROTATION_OP_NAME_MATERIALIZE_FUNC_DICT.items():
377
+ register_quantized_op(
378
+ AlgorithmName.DECOMPOSED_HADAMARD_ROTATION,
379
+ op_name,
380
+ naive_min_max_quantize.init_qsvs,
381
+ calibration_func=naive_min_max_quantize.min_max_calibrate,
382
+ materialize_func=materialize_func, # pyrefly: ignore[bad-argument-type]
383
+ )
384
+
385
+
386
+ # Register the MSE algorithm.
387
+ register_op_quant_config_validation_func(
388
+ AlgorithmName.MSE,
389
+ common_quantize.check_op_quantization_config,
390
+ )
391
+
392
+ # Register a config check policy for the MSE algorithm.
393
+ register_config_check_policy_func(
394
+ AlgorithmName.MSE,
395
+ default_policy.DEFAULT_CONFIG_CHECK_POLICY,
396
+ )
397
+
398
+ # Register specialized MSE materialize functions.
399
+ _MSE_OP_NAME_MATERIALIZE_FUNC_DICT = immutabledict({
400
+ _TFLOpName.FULLY_CONNECTED: common_quantize.materialize_fc_conv,
401
+ _TFLOpName.EMBEDDING_LOOKUP: common_quantize.materialize_embedding_lookup,
402
+ _TFLOpName.CONV_2D: common_quantize.materialize_fc_conv,
403
+ _TFLOpName.DEPTHWISE_CONV_2D: common_quantize.materialize_fc_conv,
404
+ _TFLOpName.CONV_2D_TRANSPOSE: common_quantize.materialize_conv2d_transpose,
405
+ })
406
+ for (
407
+ op_name,
408
+ materialize_func,
409
+ ) in _MSE_OP_NAME_MATERIALIZE_FUNC_DICT.items():
410
+ register_quantized_op(
411
+ AlgorithmName.MSE,
412
+ op_name,
413
+ naive_min_max_quantize.init_qsvs,
414
+ calibration_func=naive_min_max_quantize.min_max_calibrate,
415
+ materialize_func=functools.partial(
416
+ materialize_func,
417
+ mse.get_tensor_quant_params,
418
+ ),
419
+ )
420
+
421
+ # Register the GPTQ algorithm.
422
+ register_op_quant_config_validation_func(
423
+ AlgorithmName.GPTQ,
424
+ common_quantize.check_op_quantization_config,
425
+ )
426
+
427
+ # Register a config check policy for the GPTQ algorithm.
428
+ register_config_check_policy_func(
429
+ AlgorithmName.GPTQ,
430
+ default_policy.DEFAULT_CONFIG_CHECK_POLICY,
431
+ )
432
+
433
+ # Register specialized GPTQ materialize functions.
434
+ _GPTQ_OP_NAME_MATERIALIZE_FUNC_DICT = immutabledict({
435
+ _TFLOpName.FULLY_CONNECTED: common_quantize.materialize_fc_conv,
436
+ })
437
+ for (
438
+ op_name,
439
+ materialize_func,
440
+ ) in _GPTQ_OP_NAME_MATERIALIZE_FUNC_DICT.items():
441
+ register_quantized_op(
442
+ AlgorithmName.GPTQ,
443
+ op_name,
444
+ naive_min_max_quantize.init_qsvs,
445
+ calibration_func=gptq.calibrate, # pyrefly: ignore[bad-argument-type]
446
+ materialize_func=functools.partial(
447
+ materialize_func,
448
+ gptq.get_tensor_quant_params,
449
+ ),
450
+ update_qsv_func=qsv_utils.gptq_and_moving_average_update, # pyrefly: ignore[bad-argument-type]
451
+ )
452
+
453
+ # Register the OSCAR algorithm.
454
+ register_op_quant_config_validation_func(
455
+ AlgorithmName.OSCAR,
456
+ common_quantize.check_op_quantization_config,
457
+ )
458
+
459
+ # Register a config check policy for the OSCAR algorithm.
460
+ register_config_check_policy_func(
461
+ AlgorithmName.OSCAR,
462
+ default_policy.DEFAULT_CONFIG_CHECK_POLICY,
463
+ )
464
+
465
+ # Register specialized OSCAR materialize functions.
466
+ _OSCAR_OP_NAME_MATERIALIZE_FUNC_DICT = immutabledict({
467
+ _TFLOpName.FULLY_CONNECTED: oscar.materialize_fully_connected,
468
+ })
469
+ for (
470
+ op_name,
471
+ materialize_func,
472
+ ) in _OSCAR_OP_NAME_MATERIALIZE_FUNC_DICT.items():
473
+ register_quantized_op(
474
+ AlgorithmName.OSCAR,
475
+ op_name,
476
+ naive_min_max_quantize.init_qsvs,
477
+ calibration_func=oscar.calibrate, # pyrefly: ignore[bad-argument-type]
478
+ materialize_func=materialize_func, # pyrefly: ignore[bad-argument-type]
479
+ update_qsv_func=qsv_utils.oscar_and_moving_average_update, # pyrefly: ignore[bad-argument-type]
480
+ )