@@ -2075,12 +2075,12 @@ def generate_loop_schedules(kernel, callables_table, debug_args={}):
20752075 callables_table , debug_args = debug_args )
20762076
20772077
2078- def postprocess_schedule (kernel , gen_sched ):
2078+ def postprocess_schedule (kernel , callables_table , gen_sched ):
20792079 from loopy .kernel import KernelState
20802080 gen_sched = convert_barrier_instructions_to_barriers (
20812081 kernel , gen_sched )
20822082
2083- gsize , lsize = kernel .get_grid_size_upper_bounds ()
2083+ gsize , lsize = kernel .get_grid_size_upper_bounds (callables_table )
20842084
20852085 if (gsize or lsize ):
20862086 if not kernel .options .disable_global_barriers :
@@ -2118,7 +2118,7 @@ def generate_loop_schedules_inner(kernel, callables_table, debug_args={}):
21182118
21192119 try :
21202120 gen_sched = generate_loop_schedules_v2 (kernel )
2121- yield postprocess_schedule (kernel , gen_sched )
2121+ yield postprocess_schedule (kernel , callables_table , gen_sched )
21222122 return
21232123 except V2SchedulerNotImplementedException as e :
21242124 from warnings import warn
@@ -2233,7 +2233,7 @@ def print_longest_dead_end():
22332233 sched_state , debug = debug , ** schedule_gen_kwargs ):
22342234 debug .stop ()
22352235
2236- new_kernel = postprocess_schedule (kernel , gen_sched )
2236+ new_kernel = postprocess_schedule (kernel , callables_table , gen_sched )
22372237 yield new_kernel
22382238
22392239 debug .start ()
0 commit comments