2626 CUSTOM_SHADER_DOMAIN_NAME ,
2727 decode_payload ,
2828 grid_sampler_2d_operator_name ,
29+ GRID_SAMPLER_2D_QUANTIZED_GRID_VK_FORMAT ,
2930 GRID_SAMPLER_2D_SAMPLER_INT8_VK_FORMAT ,
3031 GRID_SAMPLER_2D_SAMPLER_VK_FORMAT ,
3132 GRID_SAMPLER_2D_SHADER_ENTRY_POINT ,
3435)
3536from executorch .exir import to_edge
3637from executorch .exir .dialects ._ops import ops as exir_ops
38+ from executorch .exir .pass_base import ExportedProgramPassBase , ExportedProgramPassResult
3739from torch .export import export
3840from torchao .quantization .pt2e .quantize_pt2e import convert_pt2e , prepare_pt2e
3941
@@ -56,6 +58,43 @@ def forward(self, x, grid):
5658 )
5759
5860
61+ class _RewriteGridSamplerToTosaCustomExportPass (ExportedProgramPassBase ):
62+ # The quantized grid-sampler rewrite materializes exported constant
63+ # placeholders for grid scale/zero-point, so it needs ExportedProgram
64+ # context. Production VGF lowering injects that context via the Arm pass
65+ # manager adapter, but this unit test drives the pass through the generic
66+ # EXIR transform path instead. Wrap the graph pass so the test exercises
67+ # the same rewrite logic without depending on the Arm-specific adapter.
68+ def call (self , exported_program ):
69+ rewrite_pass = RewriteGridSamplerToTosaCustomPass (exported_program )
70+ result = rewrite_pass (exported_program .graph_module )
71+ exported_program ._graph_module = result .graph_module
72+ return ExportedProgramPassResult (exported_program , result .modified )
73+
74+
75+ def test_get_first_user_input_placeholder_accepts_renamed_placeholder_node ():
76+ model = GridSampler2d ()
77+ example_inputs = (
78+ torch .randn (1 , 3 , 8 , 8 ),
79+ torch .randn (1 , 4 , 4 , 2 ),
80+ )
81+
82+ exported_program = to_edge (export (model , example_inputs )).exported_program ()
83+ first_placeholder = next (
84+ node for node in exported_program .graph .nodes if node .op == "placeholder"
85+ )
86+ original_target = first_placeholder .target
87+ first_placeholder .name = f"{ original_target } _renamed"
88+
89+ rewrite_pass = RewriteGridSamplerToTosaCustomPass (exported_program )
90+
91+ assert (
92+ rewrite_pass ._get_first_user_input_placeholder (exported_program .graph )
93+ is first_placeholder
94+ )
95+ assert first_placeholder .target == original_target
96+
97+
5998def test_rewrite_grid_sampler_to_tosa_custom_vgf_no_target ():
6099 model = GridSampler2d ()
61100 example_inputs = (
@@ -123,7 +162,8 @@ def test_rewrite_grid_sampler_to_tosa_custom_sampler_dispatch_rounds_up_output()
123162 edge_model = to_edge (export (model , example_inputs ))
124163 with TosaLoweringContext (TosaSpecification .create_from_string ("TOSA-1.0+FP" )):
125164 edge_model = edge_model .transform ([RewriteGridSamplerToTosaCustomPass ()])
126- nodes = list (edge_model .exported_program ().graph .nodes )
165+ exported_program = edge_model .exported_program ()
166+ nodes = list (exported_program .graph .nodes )
127167
128168 custom_node = next (
129169 node for node in nodes if node .target == exir_ops .backend .tosa .CUSTOM .default
@@ -145,7 +185,8 @@ def test_rewrite_grid_sampler_to_tosa_custom_no_target_uses_sampler_for_c4():
145185 edge_model = to_edge (export (model , example_inputs ))
146186 with TosaLoweringContext (TosaSpecification .create_from_string ("TOSA-1.0+FP" )):
147187 edge_model = edge_model .transform ([RewriteGridSamplerToTosaCustomPass ()])
148- nodes = list (edge_model .exported_program ().graph .nodes )
188+ exported_program = edge_model .exported_program ()
189+ nodes = list (exported_program .graph .nodes )
149190
150191 custom_node = next (
151192 node for node in nodes if node .target == exir_ops .backend .tosa .CUSTOM .default
@@ -187,29 +228,62 @@ def test_quantized_grid_sampler_uses_int8_sampler_payload(
187228
188229 edge_model = to_edge (export (converted , example_inputs , strict = True ))
189230 with TosaLoweringContext (TosaSpecification .create_from_string ("TOSA-1.0+FP+INT" )):
231+ edge_model = edge_model .transform ([FoldAndAnnotateQParamsPass ()])
232+ grid_sampler_node = next (
233+ node
234+ for node in edge_model .exported_program ().graph .nodes
235+ if node .target == exir_ops .edge .aten .grid_sampler_2d .default
236+ )
237+ expected_grid_qparams = grid_sampler_node .meta ["input_qparams" ][1 ]
238+ expected_grid_scale = torch .tensor (
239+ [expected_grid_qparams .get_scale_per_tensor ()], dtype = torch .float32
240+ )
241+ expected_grid_zero_point = torch .tensor (
242+ [expected_grid_qparams .get_zp_per_tensor ()], dtype = torch .int32
243+ )
190244 edge_model = edge_model .transform (
191245 [
192- FoldAndAnnotateQParamsPass (),
193246 InsertGridSamplerGridDequantPass (),
194- RewriteGridSamplerToTosaCustomPass (),
247+ _RewriteGridSamplerToTosaCustomExportPass (),
195248 ]
196249 )
197- nodes = list (edge_model .exported_program ().graph .nodes )
250+ exported_program = edge_model .exported_program ()
251+ nodes = list (exported_program .graph .nodes )
198252
199253 custom_node = next (
200254 node for node in nodes if node .target == exir_ops .backend .tosa .CUSTOM .default
201255 )
202256 payload = decode_payload (custom_node .kwargs ["implementation_attrs" ])
203257 grid_input = custom_node .args [0 ][1 ]
204-
258+ grid_scale_input = custom_node .args [0 ][2 ]
259+ grid_zero_point_input = custom_node .args [0 ][3 ]
205260 assert payload ["input_0_type" ] == "Image"
206261 assert payload ["input_0_vkformat" ] == GRID_SAMPLER_2D_SAMPLER_INT8_VK_FORMAT
207262 assert payload ["input_1_type" ] == "Tensor"
208- assert payload ["input_1_vkformat" ] == GRID_SAMPLER_2D_VK_FORMAT
263+ assert payload ["input_1_vkformat" ] == GRID_SAMPLER_2D_QUANTIZED_GRID_VK_FORMAT
264+ assert payload ["input_2_type" ] == "Tensor"
265+ assert payload ["input_2_vkformat" ] == "VK_FORMAT_R32_SFLOAT"
266+ assert payload ["input_2_binding" ] == 3
267+ assert payload ["input_3_type" ] == "Tensor"
268+ assert payload ["input_3_vkformat" ] == "VK_FORMAT_R32_SINT"
269+ assert payload ["input_3_binding" ] == 4
209270 assert payload ["output_0_type" ] == "Image"
210271 assert payload ["output_0_vkformat" ] == GRID_SAMPLER_2D_SAMPLER_INT8_VK_FORMAT
211- assert grid_input .meta ["val" ].dtype == torch .float32
212- assert grid_input .target in (
272+ assert payload ["output_0_binding" ] == 2
273+ assert grid_input .meta ["val" ].dtype == torch .int8
274+ assert grid_scale_input .op == "placeholder"
275+ assert grid_scale_input .meta ["val" ].dtype == torch .float32
276+ assert grid_scale_input .meta ["val" ].shape == expected_grid_scale .shape
277+ assert torch .equal (
278+ exported_program .constants [grid_scale_input .name ], expected_grid_scale
279+ )
280+ assert grid_zero_point_input .op == "placeholder"
281+ assert grid_zero_point_input .meta ["val" ].dtype == torch .int32
282+ assert grid_zero_point_input .meta ["val" ].shape == expected_grid_zero_point .shape
283+ assert torch .equal (
284+ exported_program .constants [grid_zero_point_input .name ], expected_grid_zero_point
285+ )
286+ assert grid_input .target not in (
213287 exir_ops .edge .quantized_decomposed .dequantize_per_tensor .default ,
214288 exir_ops .edge .quantized_decomposed .dequantize_per_channel .default ,
215289 )
@@ -220,6 +294,42 @@ def test_quantized_grid_sampler_uses_int8_sampler_payload(
220294 assert next (iter (custom_node .meta ["output_qparams" ].values ())).qmax == 127
221295
222296
297+ def test_quantized_grid_sampler_dequantizes_grid_for_non_sampler_path ():
298+ model = GridSampler2d ().eval ()
299+ model .interpolation_mode_ = 2
300+ example_inputs = (
301+ torch .randn (1 , 4 , 8 , 8 ),
302+ torch .rand (1 , 4 , 4 , 2 ),
303+ )
304+ quantizer = VgfQuantizer (VgfCompileSpec ("TOSA-1.0+INT" ))
305+ quantizer .set_global (get_symmetric_quantization_config (is_per_channel = False ))
306+
307+ exported = export (model , example_inputs , strict = True )
308+ prepared = prepare_pt2e (exported .module (), quantizer )
309+ prepared (* example_inputs )
310+ converted = convert_pt2e (prepared )
311+
312+ edge_model = to_edge (export (converted , example_inputs , strict = True ))
313+ with TosaLoweringContext (TosaSpecification .create_from_string ("TOSA-1.0+FP+INT" )):
314+ edge_model = edge_model .transform (
315+ [FoldAndAnnotateQParamsPass (), InsertGridSamplerGridDequantPass ()]
316+ )
317+
318+ grid_sampler_node = next (
319+ node
320+ for node in edge_model .exported_program ().graph .nodes
321+ if node .target == exir_ops .edge .aten .grid_sampler_2d .default
322+ )
323+ grid_input = grid_sampler_node .args [1 ]
324+
325+ assert (
326+ grid_input .target
327+ == exir_ops .edge .quantized_decomposed .dequantize_per_tensor .default
328+ )
329+ assert grid_input .meta ["val" ].dtype == torch .float32
330+ assert 1 not in grid_sampler_node .meta ["input_qparams" ]
331+
332+
223333def test_quantized_grid_sampler_rejects_dequantized_grid_with_int8_image_payload ():
224334 model = GridSampler2d ().eval ()
225335 example_inputs = (
0 commit comments