3737F = TypeVar ("F" , bound = Callable [..., Any ])
3838PathCGradientProbe = Callable [[Mapping [str , Any ]], None ]
3939PathCTrainingRuntime = Any
40- PATH_C_TRAINING_VALUE_AND_GRAD_CONTRACT = "path_c_direct_fusion_value_and_grad_v1"
40+ PATH_C_DIRECT_FUSION_VALUE_AND_GRAD_CONTRACT = (
41+ "path_c_direct_fusion_value_and_grad_v1"
42+ )
43+ PATH_C_TRAINING_VALUE_AND_GRAD_CONTRACT = (
44+ PATH_C_DIRECT_FUSION_VALUE_AND_GRAD_CONTRACT
45+ )
46+ PATH_C_FUSED_TRAIN_BLOCK_TRAINING_RUNTIME_CONTRACT = (
47+ "path_c_fused_train_block_training_runtime_v1"
48+ )
49+ PATH_C_FUSED_TRAIN_BLOCK_VALUE_AND_GRAD_CONTRACT = (
50+ "path_c_fused_train_block_value_and_grad_v1"
51+ )
4152
4253REGIONAL_COMPILE_TARGETS : Mapping [CompileTarget , bool ] = {
4354 "mamba3_pre" : True ,
@@ -337,12 +348,16 @@ def attach_path_c_training_runtime(self, runtime: PathCTrainingRuntime) -> None:
337348 bind = getattr (runtime , "bind_training_graph" , None )
338349 bound = False
339350 if callable (bind ):
340- bind (
341- owner = "CompiledPretrainingStep" ,
342- uses_direct_chain_runtime = True ,
343- uses_forward_hook = True ,
344- uses_backward_or_vjp_hook = True ,
345- )
351+ binding = {
352+ "owner" : "CompiledPretrainingStep" ,
353+ "uses_forward_hook" : True ,
354+ "uses_backward_or_vjp_hook" : True ,
355+ }
356+ if _path_c_training_runtime_uses_fused_train_block (runtime ):
357+ binding ["uses_fused_train_block_runtime" ] = True
358+ else :
359+ binding ["uses_direct_chain_runtime" ] = True
360+ bind (** binding )
346361 bound = True
347362 try :
348363 value_and_grad_contract = _path_c_training_runtime_value_and_grad_contract (
@@ -575,7 +590,16 @@ def _path_c_training_runtime_value_and_grad_contract(
575590 payload = dict (raw_contract )
576591 contract = str (payload .get ("contract" , "" ))
577592 owner = str (payload .get ("owner" , "" ))
578- uses_runtime = bool (payload .get ("uses_direct_chain_runtime" ))
593+ uses_direct_chain_runtime = bool (payload .get ("uses_direct_chain_runtime" ))
594+ uses_fused_train_block_runtime = bool (
595+ payload .get ("uses_fused_train_block_runtime" )
596+ )
597+ direct_contract = contract == PATH_C_DIRECT_FUSION_VALUE_AND_GRAD_CONTRACT
598+ fused_contract = contract == PATH_C_FUSED_TRAIN_BLOCK_VALUE_AND_GRAD_CONTRACT
599+ uses_runtime = bool (
600+ (direct_contract and uses_direct_chain_runtime )
601+ or (fused_contract and uses_fused_train_block_runtime )
602+ )
579603 uses_forward = bool (payload .get ("uses_forward_hook" ))
580604 uses_reverse = bool (payload .get ("uses_backward_or_vjp_hook" ))
581605 returns_model_grads = bool (payload .get ("returns_model_grads" ))
@@ -586,7 +610,7 @@ def _path_c_training_runtime_value_and_grad_contract(
586610 hidden_packing = bool (payload .get ("hidden_packing_performed" , False ))
587611 status = (
588612 "ok"
589- if contract == PATH_C_TRAINING_VALUE_AND_GRAD_CONTRACT
613+ if ( direct_contract or fused_contract )
590614 and owner == "CompiledPretrainingStep"
591615 and uses_runtime
592616 and uses_forward
@@ -604,7 +628,8 @@ def _path_c_training_runtime_value_and_grad_contract(
604628 "status" : status ,
605629 "contract" : contract or PATH_C_TRAINING_VALUE_AND_GRAD_CONTRACT ,
606630 "owner" : owner or None ,
607- "uses_direct_chain_runtime" : uses_runtime ,
631+ "uses_direct_chain_runtime" : uses_direct_chain_runtime ,
632+ "uses_fused_train_block_runtime" : uses_fused_train_block_runtime ,
608633 "uses_forward_hook" : uses_forward ,
609634 "uses_backward_or_vjp_hook" : uses_reverse ,
610635 "returns_model_grads" : returns_model_grads ,
@@ -616,9 +641,22 @@ def _path_c_training_runtime_value_and_grad_contract(
616641 }
617642
618643
644+ def _path_c_training_runtime_uses_fused_train_block (
645+ runtime : PathCTrainingRuntime ,
646+ ) -> bool :
647+ return bool (
648+ getattr (runtime , "uses_fused_train_block_runtime" , False )
649+ or str (getattr (runtime , "contract" , "" ))
650+ == PATH_C_FUSED_TRAIN_BLOCK_TRAINING_RUNTIME_CONTRACT
651+ )
652+
653+
619654__all__ = [
620655 "CompileTarget" ,
621656 "CompiledPretrainingStep" ,
657+ "PATH_C_DIRECT_FUSION_VALUE_AND_GRAD_CONTRACT" ,
658+ "PATH_C_FUSED_TRAIN_BLOCK_TRAINING_RUNTIME_CONTRACT" ,
659+ "PATH_C_FUSED_TRAIN_BLOCK_VALUE_AND_GRAD_CONTRACT" ,
622660 "PATH_C_TRAINING_VALUE_AND_GRAD_CONTRACT" ,
623661 "PathCGradientBufferCapture" ,
624662 "PathCGradientProbe" ,
0 commit comments