@@ -449,8 +449,14 @@ def compile(
449449 if isinstance (seq_len , list ) and len (seq_len ) >= 15 :
450450 warnings .warn ("Recommended: `seq_len` should contain fewer than 15 items." )
451451
452+ _seq_lens = seq_len if isinstance (seq_len , list ) else [seq_len ]
452453 specializations = [
453- {"batch_size" : batch_size , "seq_len" : sl } for sl in (seq_len if isinstance (seq_len , list ) else [seq_len ])
454+ {
455+ "_graph_name" : "Embedding" if len (_seq_lens ) == 1 else f"Embedding_{ i } " ,
456+ "batch_size" : batch_size ,
457+ "seq_len" : sl ,
458+ }
459+ for i , sl in enumerate (_seq_lens )
454460 ]
455461
456462 target_dtype = getattr (self .model .config , "torch_dtype" , torch .float32 )
@@ -794,8 +800,14 @@ def compile(
794800 if isinstance (seq_len , list ) and len (seq_len ) >= 15 :
795801 warnings .warn ("Recommended: `seq_len` should contain fewer than 15 items." )
796802
803+ _seq_lens = seq_len if isinstance (seq_len , list ) else [seq_len ]
797804 specializations = [
798- {"batch_size" : batch_size , "seq_len" : sl } for sl in (seq_len if isinstance (seq_len , list ) else [seq_len ])
805+ {
806+ "_graph_name" : "SeqClassification" if len (_seq_lens ) == 1 else f"SeqClassification_{ i } " ,
807+ "batch_size" : batch_size ,
808+ "seq_len" : sl ,
809+ }
810+ for i , sl in enumerate (_seq_lens )
799811 ]
800812 target_dtype = getattr (self .model .config , "torch_dtype" , torch .float32 )
801813 return self ._compile (
@@ -1582,6 +1594,7 @@ def compile(
15821594 compile_dir = compile_dir ,
15831595 compile_only = True ,
15841596 specializations = specializations ["vision" ],
1597+ specialization_module_name = "Vision" ,
15851598 convert_to_fp16 = (CUSTOM_IO_DTYPE_MAP [target_dtype ] == "float16" ),
15861599 mxfp6_matmul = constants .VISION_MXFP6_MATMUL ,
15871600 mdp_ts_num_devices = num_devices ,
@@ -3256,7 +3269,9 @@ def build_prefill_specialization(
32563269 # TODO: remove this; not required
32573270 if full_batch_size :
32583271 spec ["full_batch_exec_size" ] = exec_batch_size
3259- return {k : v for k , v in spec .items () if v is not None }
3272+ result = {k : v for k , v in spec .items () if v is not None }
3273+ result ["_graph_name" ] = "Prefill"
3274+ return result
32603275
32613276 def build_decode_specialization (
32623277 self ,
@@ -3314,7 +3329,9 @@ def build_decode_specialization(
33143329 spec ["full_batch_size" ] = kv_cache_batch_size
33153330 else :
33163331 spec ["batch_size" ] = kv_cache_batch_size
3317- return {k : v for k , v in spec .items () if v is not None }
3332+ result = {k : v for k , v in spec .items () if v is not None }
3333+ result ["_graph_name" ] = "Decode"
3334+ return result
33183335
33193336 def compile (
33203337 self ,
@@ -4235,8 +4252,10 @@ def compile(
42354252 :str: Path of the compiled ``qpc`` package.
42364253 """
42374254
4255+ _seq_lens = seq_len if isinstance (seq_len , list ) else [seq_len ]
42384256 specializations = [
4239- {"batch_size" : batch_size , "seq_len" : sl } for sl in (seq_len if isinstance (seq_len , list ) else [seq_len ])
4257+ {"_graph_name" : "CTC" if len (_seq_lens ) == 1 else f"CTC_{ i } " , "batch_size" : batch_size , "seq_len" : sl }
4258+ for i , sl in enumerate (_seq_lens )
42404259 ]
42414260
42424261 target_dtype = getattr (self .model .config , "torch_dtype" , torch .float32 )
0 commit comments