| 49 | self.fn = fn |
| 50 | |
| 51 | def _bench(self, *args, config, **meta): |
| 52 | # check for conflicts, i.e. meta-parameters both provided |
| 53 | # as kwargs and by the autotuner |
| 54 | conflicts = meta.keys() & config.kwargs.keys() |
| 55 | if conflicts: |
| 56 | raise ValueError( |
| 57 | f"Conflicting meta-parameters: {', '.join(conflicts)}." |
| 58 | " Make sure that you don't re-define auto-tuned symbols." |
| 59 | ) |
| 60 | # augment meta-parameters with tunable ones |
| 61 | current = dict(meta, **config.kwargs) |
| 62 | |
| 63 | def kernel_call(): |
| 64 | if config.pre_hook: |
| 65 | config.pre_hook(self.nargs) |
| 66 | self.hook(args) |
| 67 | self.fn.run(*args, num_warps=config.num_warps, num_stages=config.num_stages, **current) |
| 68 | try: |
| 69 | # In testings using only 40 reps seems to be close enough and it appears to be what PyTorch uses |
| 70 | # PyTorch also sets fast_flush to True, but I didn't see any speedup so I'll leave the default |
| 71 | return triton.testing.do_bench(kernel_call, rep=40) |
| 72 | except triton.compiler.OutOfResources: |
| 73 | return float('inf') |
| 74 | |
| 75 | def run(self, *args, **kwargs): |
| 76 | self.nargs = dict(zip(self.arg_names, args)) |