7171from thunder .extend import FUEL_LEVEL , FusionExecutor , register_executor
7272from thunder .executors .nvfuserex import nvfuser_version
7373
74+
75+ DTENSOR_SUPPORTED_VERSION = LooseVersion ("0.2.28" )
76+ if nvfuser_version () >= DTENSOR_SUPPORTED_VERSION :
77+ import nvfuser_direct as nvfd
78+ from nvfuser_direct import FusionDefinition as DirectFusionDefinition
79+
7480# NOTE This impl file is here because nvFuser may not be available, so it's imported conditionally
7581# by nvfuserex.py when nvFuser is available.
7682import nvfuser
@@ -241,41 +247,34 @@ def get_translator(bsym: BoundSymbol) -> Callable:
241247 return _translation_map [bsym .sym .id ]
242248
243249
244- class MultiDeviceFusionDefinition (FusionDefinition ):
245- def __init__ (self , define_fn : Callable [[FusionDefinition ], None ], in_dtensors : list [DTensorProxy ], max_length : int ):
246- super ().__init__ (max_length = max_length )
247- self ._in_dtensors = in_dtensors
248- self ._define_fn = define_fn
250+ def register_dtensor_supported (prim_id : int , fn : Callable , checker_fn : Callable ) -> None :
251+ if nvfuser_version () < DTENSOR_SUPPORTED_VERSION :
252+ # Only register dtensor ops if supported version is available.
253+ return
249254
250- def definition (self ) -> None :
251- self ._define_fn (self )
255+ register_supported (prim_id , fn , checker_fn )
252256
253- def _find_tensor_by_index (self , index : int ) -> nvfuser .Tensor :
254- for t in self .sched .tensors ():
255- if t .index == index :
256- return t
257- return None
258257
259- def multidevice_schedule (self ) -> None :
260- for in_tensor_index , in_dtensor in zip (self .inputs (), self ._in_dtensors ):
261- in_tensor = self ._find_tensor_by_index (in_tensor_index )
258+ def multidevice_schedule (fd : FusionDefinition , in_dtensors : list [Proxy ]) -> None :
259+ for in_tv , in_dtensor in zip (fd .fusion .inputs (), in_dtensors ):
260+ assert isinstance (in_dtensor , DTensorProxy )
261+ # Set the device mesh.
262+ assert in_dtensor .device_mesh .ndim == 1 , "nvFuser's Python API only supports 1D meshes."
263+ mesh = nvfd .multidevice .DeviceMesh (in_dtensor .device_mesh .mesh .tolist ())
262264
263- # Set the device mesh.
264- utils .check (in_dtensor .device_mesh .ndim == 1 , lambda : "nvFuser's Python API only supports 1D meshes." )
265- mesh = nvfuser .DeviceMesh (in_dtensor .device_mesh .mesh .tolist ())
265+ in_tv .set_device_mesh (mesh )
266266
267- self . sched . _set_device_mesh ( in_tensor , mesh )
267+ assert len ( in_dtensor . placements ) == 1 , "nvFuser's Python API only supports 1D meshes."
268268
269- # Split and parallelize.
270- utils .check (len (in_dtensor .placements ) == 1 , lambda : "nvFuser's Python API only supports 1D meshes." )
271- # When the mesh is multi-dimensional, iterate through the
272- # placements in descending order of Placement.dim.
273- placement : Placement = in_dtensor .placements [0 ]
274- if placement .is_shard ():
275- dim = cast (Shard , placement ).dim
276- self .sched .split (in_tensor , dim , mesh .size , False )
277- self .sched .parallelize (in_tensor , dim , nvfuser .ParallelType .mesh_x )
278- self .sched .set_allocation_as_loop (in_tensor )
269+ # Split and parallelize.
270+ # When the mesh is multi-dimensional, iterate through the
271+ # placements in descending order of Placement.dim.
272+ placement : Placement = in_dtensor .placements [0 ]
273+ if placement .is_shard ():
274+ dim = cast (Shard , placement ).dim
275+ in_tv .split (dim , mesh .size , inner_split = False )
276+ in_tv .axis (dim ).parallelize (nvfd .ParallelType .mesh_x )
277+ in_tv .set_allocation_domain (in_tv .get_loop_domain (), new_contiguity = True )
279278
280279
281280def create_fd (
@@ -376,10 +375,13 @@ def check_dtensor_tracing_and_runtime_metadata(inp):
376375 lambda : "nvfuser: Expected runtime and tracing metadata to be the same for DTensor." ,
377376 )
378377
379- fd = MultiDeviceFusionDefinition ( definition , sorted_unique_inputs , max_length = MAX_LENGTH )
378+ fd = DirectFusionDefinition ( )
380379 # Device may be set in one of the "factory" methods like full, iota, or uniform
381380 # NOTE: This should be called before defining because a factory method may look-up at `_selected_device` while being defined.
382381 fd ._selected_device = None
382+ with fd :
383+ definition (fd )
384+ multidevice_schedule (fd , sorted_unique_inputs )
383385 else :
384386 # NOTE nvFuser's default max length is 1024 operations at the time of this writing
385387 # This arbitrarily increases it to 9999
@@ -535,28 +537,10 @@ def __call__(self, *args):
535537 if self .store_inputs :
536538 self .last_inputs = args
537539
538- if hasattr ( fd , "multidevice_schedule" ):
540+ if dist . is_available () and any ( isinstance ( t , torch . distributed . tensor . DTensor ) for t in args ):
539541 with annotate_for_profile (self .name ):
540- in_tensors = [in_dtensor .to_local () for in_dtensor in args ]
541- out_tensors , out_shardings = fd .execute (
542- in_tensors ,
543- device = fd ._selected_device ,
544- save_repro_inputs = self .save_fake_inputs ,
545- _enable_options = self .enable_options ,
546- _disable_options = self .disable_options ,
547- )
548-
549- assert len (out_tensors ) == len (out_shardings )
550- out_dtensors : list [DTensor ] = []
551- for out_tensor , out_sharding in zip (out_tensors , out_shardings ):
552- mesh = dist .device_mesh .init_device_mesh ("cuda" , (out_sharding .mesh .size ,))
553- placements : list [Placement ] = []
554- for parallel_type in [nvfuser .ParallelType .mesh_x ]:
555- axis : int = out_sharding .axis_sharded_on (parallel_type )
556- placements .append (Replicate () if axis == - 1 else Shard (axis ))
557- out_dtensors .append (DTensor .from_local (out_tensor , mesh , placements ))
558-
559- return out_dtensors
542+ output = nvfd .execute_with_dtensors (fd , args )
543+ return output
560544 else :
561545 with annotate_for_profile (self .name ):
562546 return fd .execute (
@@ -1906,6 +1890,9 @@ def le(a: TensorProxy | Number, b: TensorProxy | Number, *, fd: FusionDefinition
19061890 return fd .ops .le (nva , nvb )
19071891
19081892
1893+ register_supported (PrimIDs .LE , le , _elementwise_binary_check )
1894+
1895+
19091896def lt (a : TensorProxy | Number , b : TensorProxy | Number , * , fd : FusionDefinition , lc_to_nv_map : dict ) -> Any :
19101897 nva = getnv (a , fd , lc_to_nv_map )
19111898 nvb = getnv (b , fd , lc_to_nv_map )
@@ -1924,7 +1911,7 @@ def mul(a: TensorProxy | Number, b: TensorProxy | Number, *, fd: FusionDefinitio
19241911
19251912
19261913register_supported (PrimIDs .MUL , mul , _elementwise_binary_check )
1927- register_supported (dtensor_mul_prim .id , mul , _elementwise_binary_check )
1914+ register_dtensor_supported (dtensor_mul_prim .id , mul , _elementwise_binary_check )
19281915
19291916
19301917def ne (a : TensorProxy | Number , b : TensorProxy | Number , * , fd : FusionDefinition , lc_to_nv_map : dict ) -> Any :
0 commit comments