@@ -113,7 +113,55 @@ def call_pipeline(config, pipeline, prompt, negative_prompt):
113113 return out
114114
115115
116+ def maybe_tune_block_sizes (config ):
117+ """If enable_tile_search, run a fast one-DiT-block tile-size grid search and overwrite
118+ flash_block_sizes' block_q/block_kv/block_kv_compute with the winner IN PLACE.
119+ """
120+ keys = config .get_keys ()
121+ val = keys .get ("enable_tile_search" , False )
122+ if str (val ).lower () not in ("true" , "1" , "yes" ):
123+ return
124+ from maxdiffusion .utils .tile_size_grid_search import grid_search
125+ from maxdiffusion .utils .ltx2_block_benchmark import LTX2BlockBenchmark
126+
127+ vmem = config .flash_block_sizes .get ("vmem_limit_bytes" , None ) if config .flash_block_sizes else None
128+ if vmem is None :
129+ import os , re
130+ m = re .search (r'--xla_tpu_scoped_vmem_limit_kib=(\d+)' , os .environ .get ("LIBTPU_INIT_ARGS" , "" ))
131+ vmem = int (m .group (1 )) * 1024 if m else 32 * 1024 * 1024
132+
133+ mesh = jax .sharding .Mesh (max_utils .create_device_mesh (config ), config .mesh_axes )
134+ bench = LTX2BlockBenchmark .from_config (config , mesh , vmem_limit_bytes = vmem )
135+ max_logging .log (f"[tile-search] tuning block sizes for { bench .label } (vmem={ vmem / 1024 / 1024 :.1f} MB) before inference..." )
136+ result = grid_search (
137+ bench ,
138+ mode = keys .get ("tile_search_mode" , "smart" ),
139+ iters = keys .get ("tile_search_iters" , 10 ),
140+ out_dir = (keys .get ("tile_search_out" , "" ) or None ),
141+ log = max_logging .log ,
142+ )
143+ if result .best is None :
144+ max_logging .log ("[tile-search] no config succeeded; keeping configured flash_block_sizes" )
145+ return
146+ fbs = dict (config .flash_block_sizes ) if config .flash_block_sizes else {}
147+ fbs .update ({
148+ "block_q" : result .best .bq ,
149+ "block_kv" : result .best .bkv ,
150+ "block_kv_compute" : result .best .bkv_compute ,
151+ "block_kv_compute_in" : result .best .bkv_compute ,
152+ "vmem_limit_bytes" : vmem ,
153+ })
154+ config .get_keys ()["flash_block_sizes" ] = fbs
155+ max_logging .log (
156+ f"[tile-search] using block_q={ result .best .bq } block_kv={ result .best .bkv } "
157+ f"(block-bench { result .best .mean_ms :.2f} ms)"
158+ )
159+
160+
116161def run (config , pipeline = None , filename_prefix = "" , commit_hash = None ):
162+ if pipeline is None :
163+ maybe_tune_block_sizes (config )
164+
117165 writer = max_utils .initialize_summary_writer (config )
118166 if jax .process_index () == 0 and writer :
119167 max_logging .log (f"TensorBoard logs will be written to: { config .tensorboard_dir } " )
0 commit comments