(self, expected_result)
| 341 | ) |
| 342 | |
| 343 | def check_marker(self, expected_result): |
| 344 | paddle.framework.set_flags({"FLAGS_trt_min_group_size": 1}) |
| 345 | with paddle.pir_utils.IrGuard(): |
| 346 | main_program, startup_program, fetch_list = ( |
| 347 | self.create_fake_program() |
| 348 | ) |
| 349 | main_program = run_pir_pass( |
| 350 | main_program, |
| 351 | disable_passes=self.disable_passes, |
| 352 | ) |
| 353 | marker_result = False |
| 354 | for op in main_program.global_block().ops: |
| 355 | if op.name() == self.target_marker_op: |
| 356 | marker_result = op.attrs().get("__l_trt__", False) |
| 357 | self.assertEqual(marker_result, expected_result) |
no test coverage detected