已合并
refactor: OpDispatcher 重构 #508
refactor: OpDispatcher 重构 #508
已合并
hedongdong创建于 4月8日
共 37 个文件变更+637-614
@@ -131,23 +131,33 @@ def get_dtensor_dispatch():
131 131 
132 132 
133class LayoutCacheKey:133class LayoutCacheKey:
134- """134+ """Immutable layout cache key."""
135- Layout cache key135+ __slots__ = ('_tuple', '_hash')
136- """136+ 
137 def __init__(self, layout_ids: List[str]):137 def __init__(self, layout_ids: List[str]):
138- self.layout_ids = layout_ids138+ self._tuple = tuple(layout_ids)
139+ self._hash = hash(self._tuple)
140+ 
141+ @classmethod
142+ def from_cache_values(cls, cache_values):
143+ key_values = []
144+ for v in cache_values:
145+ if hasattr(v, 'compact_str'):
146+ key_values.append(str(v.compact_str))
147+ else:
148+ key_values.append(str(v))
149+ return cls(key_values)
139 150 
140 def __eq__(self, other):151 def __eq__(self, other):
141 if not isinstance(other, LayoutCacheKey):152 if not isinstance(other, LayoutCacheKey):
142 return False153 return False
143- return self.layout_ids == other.layout_ids154+ return self._tuple == other._tuple
144 155 
145 def __hash__(self):156 def __hash__(self):
146- seed = 0157+ return self._hash
147- for id_str in self.layout_ids:158+ 
148- h = hash(id_str)159+ def __repr__(self):
149- seed ^= h + 0x9e3779b9 + (seed << 6) + (seed >> 2)160+ return f"LayoutCacheKey({self._tuple})"
150- return seed
151 161 
152class LayoutCacheManager:162class LayoutCacheManager:
153 """163 """
@@ -339,17 +349,15 @@ class OpDispatcher:
339 349 
340 @staticmethod350 @staticmethod
341 def _process_args_and_kwargs(351 def _process_args_and_kwargs(
342- args, kwargs, cache_key: "LayoutCacheKey"352+ args, kwargs
343- ) -> tuple[list, list, list, dict]:353+ ) -> tuple[list, list, list, dict, list]:
344 """_process_args_and_kwargs"""354 """_process_args_and_kwargs"""
345- # input_layouts contain prarmeters which have layout, extra_args contain other parameters
346 input_layouts = []355 input_layouts = []
347 extra_args = []356 extra_args = []
348- # input_args are position prarmeters, input_kwargs are keyword parameters
349 input_args = []357 input_args = []
350 input_kwargs = kwargs.copy()358 input_kwargs = kwargs.copy()
359+ cache_key_values = []
351 360 
352- # Normal ops pass real inputs directly (e.g. SumExt: args = (dtensor, axis: list, keep_dims: bool, dtype: None)).
353 for arg in args:361 for arg in args:
354 if arg is None:362 if arg is None:
355 input_layouts.append(None)363 input_layouts.append(None)
@@ -360,14 +368,14 @@ class OpDispatcher:
360 id_str = "scalar"368 id_str = "scalar"
361 if not isinstance(arg, Tensor):369 if not isinstance(arg, Tensor):
362 id_str = str(arg)370 id_str = str(arg)
363- cache_key.layout_ids.append(id_str)371+ cache_key_values.append(id_str)
364 extra_args.append(arg)372 extra_args.append(arg)
365 input_layouts.append(None)373 input_layouts.append(None)
366 input_args.append(arg)374 input_args.append(arg)
367 else:375 else:
368 layout = arg.layout376 layout = arg.layout
369 layout_id = layout.compact_str377 layout_id = layout.compact_str
370- cache_key.layout_ids.append(str(layout_id))378+ cache_key_values.append(str(layout_id))
371 input_layouts.append(layout)379 input_layouts.append(layout)
372 if isinstance(arg, DTensor):380 if isinstance(arg, DTensor):
373 input_args.append(arg.to_local())381 input_args.append(arg.to_local())
@@ -382,33 +390,31 @@ class OpDispatcher:
382 id_str = "scalar"390 id_str = "scalar"
383 if not isinstance(val, Tensor):391 if not isinstance(val, Tensor):
384 id_str = str(val)392 id_str = str(val)
385- cache_key.layout_ids.append(id_str)393+ cache_key_values.append(id_str)
386 extra_args.append(val)394 extra_args.append(val)
387 input_layouts.append(None)395 input_layouts.append(None)
388 else:396 else:
389 layout = val.layout397 layout = val.layout
390 layout_id = layout.compact_str398 layout_id = layout.compact_str
391- cache_key.layout_ids.append(str(layout_id))399+ cache_key_values.append(str(layout_id))
392 input_layouts.append(layout)400 input_layouts.append(layout)
393 if isinstance(val, DTensor):401 if isinstance(val, DTensor):
394 input_kwargs[k] = val.to_local()402 input_kwargs[k] = val.to_local()
395 403 
396- return input_layouts, extra_args, input_args, input_kwargs404+ return input_layouts, extra_args, input_args, input_kwargs, cache_key_values
397 405 
398 def _with_layout_infer(self, func: callable, *args, **kwargs) -> Tensor:406 def _with_layout_infer(self, func: callable, *args, **kwargs) -> Tensor:
399 """_with_layout_infer"""407 """_with_layout_infer"""
400 func_name = platform.get_op_name(func)408 func_name = platform.get_op_name(func)
401 packed_call = None409 packed_call = None
402- # Ops in unpack_ops use packed fallback args (e.g. ScatterUpdate: (prim_obj, op_name: str, (input_x, indices, updates))).
403 if(func_name in self.unpack_ops and len(args) == 3 and410 if(func_name in self.unpack_ops and len(args) == 3 and
404 isinstance(args[1], str) and isinstance(args[2],(tuple,list))):411 isinstance(args[1], str) and isinstance(args[2],(tuple,list))):
405 packed_call = (args[0], args[1])412 packed_call = (args[0], args[1])
406 args = tuple(args[2])413 args = tuple(args[2])
407 414 
408- cache_key = LayoutCacheKey([])415+ input_layouts, extra_args, input_args, input_kwargs, cache_key_values = \
409- input_layouts, extra_args, input_args, input_kwargs = OpDispatcher._process_args_and_kwargs(416+ OpDispatcher._process_args_and_kwargs(args, kwargs)
410- args, kwargs, cache_key417+ cache_key = LayoutCacheKey(cache_key_values)
411- )
412 cache_manager = LayoutCacheManager.get_instance()418 cache_manager = LayoutCacheManager.get_instance()
413 layout_cache = cache_manager.get_layout_cache()419 layout_cache = cache_manager.get_layout_cache()
414 if func_name not in layout_cache:420 if func_name not in layout_cache:
@@ -422,9 +428,7 @@ class OpDispatcher:
422 else:428 else:
423 all_args = (input_layouts, extra_args)429 all_args = (input_layouts, extra_args)
424 output_layout = distribute_op.infer_layout(*all_args)430 output_layout = distribute_op.infer_layout(*all_args)
425- op_impl = getattr(431+ op_impl = distribute_op.get_expand_impl(func, output_layout, input_layouts, extra_args)
426- distribute_op, "get_expand_impl", lambda *args, **kwargs: None
427- )(func, output_layout, input_layouts, extra_args)
428 op_layout_cache[cache_key] = (output_layout, op_impl)432 op_layout_cache[cache_key] = (output_layout, op_impl)
429 433 
430 if op_impl is None:434 if op_impl is None:
@@ -453,21 +457,27 @@ class OpDispatcher:
453 return DTensor.from_local(457 return DTensor.from_local(
454 py_output, output_layout.mesh, output_layout.alias_placements)458 py_output, output_layout.mesh, output_layout.alias_placements)
455 459 
456- def _extract_single_arg_layout(self, arg, cache_key, extra_args, input_layouts):460+ def _extract_single_arg_layout(self, expanded_args, kwargs_value):
457 """Helper to extract layout and cache info for a single argument."""461 """Helper to extract layout and cache info for a single argument."""
458- if arg is None:462+ cache_key_values = []
459- input_layouts.append(None)463+ input_layouts = []
460- return464+ extra_args = []
461 465 
462- if not hasattr(arg, "_layout"):466+ for arg in chain(expanded_args, kwargs_value):
463- id_str = "scalar" if isinstance(arg, Tensor) else str(arg)467+ if arg is None:
464- cache_key.layout_ids.append(id_str)468+ input_layouts.append(None)
465- extra_args.append(arg)469+ continue
466- input_layouts.append(None)470+ 
467- else:471+ if not hasattr(arg, "_layout"):
468- layout = arg.layout472+ id_str = "scalar" if isinstance(arg, Tensor) else str(arg)
469- cache_key.layout_ids.append(str(layout.compact_str))473+ cache_key_values.append(id_str)
470- input_layouts.append(layout)474+ extra_args.append(arg)
475+ input_layouts.append(None)
476+ else:
477+ layout = arg.layout
478+ cache_key_values.append(str(layout.compact_str))
479+ input_layouts.append(layout)
480+ return cache_key_values, input_layouts, extra_args
471 481 
472 def _pack_infer_output(self, py_output, output_layout):482 def _pack_infer_output(self, py_output, output_layout):
473 """Helper to pack py_output into DTensors using output_layout."""483 """Helper to pack py_output into DTensors using output_layout."""
@@ -501,17 +511,10 @@ class OpDispatcher:
501 # Process kwargs into local tensors511 # Process kwargs into local tensors
502 input_kwargs = {k: (v.to_local() if isinstance(v, DTensor) else v) for k, v in kwargs.items()}512 input_kwargs = {k: (v.to_local() if isinstance(v, DTensor) else v) for k, v in kwargs.items()}
503 513 
504- cache_key = LayoutCacheKey([])
505- input_layouts = []
506- extra_args = []
507- 
508 # Extract layouts for positional args514 # Extract layouts for positional args
509- for arg in expanded_args:515+ cache_key_values, input_layouts, extra_args = self._extract_single_arg_layout(expanded_args, kwargs.values())
510- self._extract_single_arg_layout(arg, cache_key, extra_args, input_layouts)
511 516 
512- # Extract layouts for keyword args517+ cache_key = LayoutCacheKey(cache_key_values)
513- for val in kwargs.values():
514- self._extract_single_arg_layout(val, cache_key, extra_args, input_layouts)
515 518 
516 cache_manager = LayoutCacheManager.get_instance()519 cache_manager = LayoutCacheManager.get_instance()
517 layout_cache = cache_manager.get_layout_cache()520 layout_cache = cache_manager.get_layout_cache()
@@ -542,21 +545,14 @@ class OpDispatcher:
542 """_with_layout_infer_reshape"""545 """_with_layout_infer_reshape"""
543 input_tensor = args[0]546 input_tensor = args[0]
544 shape = args[1]547 shape = args[1]
545- cache_key = LayoutCacheKey([])
546- input_layouts = []
547 548 
548 layout = input_tensor.layout549 layout = input_tensor.layout
549- input_layouts.append(layout)550+ input_layouts = [layout]
550- layout_id = layout.compact_str
551- cache_key.layout_ids.append(str(layout_id))
552 551 
553- extra_args = []552+ extra_args = [shape, input_tensor.shape]
554- extra_args.append(shape)
555- cache_key.layout_ids.append(str(shape))
556 553 
557- input_shape = input_tensor.shape554+ cache_key_values = [str(layout.compact_str), str(shape), str(input_tensor.shape)]
558- extra_args.append(input_shape)555+ cache_key = LayoutCacheKey(cache_key_values)
559- cache_key.layout_ids.append(str(input_shape))
560 556 
561 cache_manager = LayoutCacheManager.get_instance()557 cache_manager = LayoutCacheManager.get_instance()
562 layout_cache = cache_manager.get_layout_cache()558 layout_cache = cache_manager.get_layout_cache()
@@ -586,15 +582,23 @@ class OpDispatcher:
586 return DTensor.from_local(py_output, infer_output_tuple[0].mesh, infer_output_tuple[0].alias_placements)582 return DTensor.from_local(py_output, infer_output_tuple[0].mesh, infer_output_tuple[0].alias_placements)
587 583 
588 @staticmethod584 @staticmethod
589- def _process_args_and_kwargs_with_shape(585+ def _process_args_and_kwargs_with_shape(args, kwargs):
590- args, kwargs, cache_key: "LayoutCacheKey"586+ """Process args and kwargs with input shapes for WithShape suffix operators.
591- ) -> tuple[list, list, list, list, dict]:587+ 
592- """_process_args_and_kwargs_with_shape"""588+ Args:
589+ args: Positional arguments from dispatch.
590+ kwargs: Keyword arguments from dispatch.
591+ 
592+ Returns:
593+ tuple: (input_layouts, input_shapes, extra_args, input_args, input_kwargs, cache_key_values)
594+ """
593 input_layouts = []595 input_layouts = []
594 extra_args = []596 extra_args = []
595 input_shapes = []597 input_shapes = []
596 input_args = []598 input_args = []
597 input_kwargs = kwargs.copy()599 input_kwargs = kwargs.copy()
600+ cache_key_values = []
601+ 
598 for arg in args:602 for arg in args:
599 if arg is None:603 if arg is None:
600 input_layouts.append(None)604 input_layouts.append(None)
@@ -606,14 +610,14 @@ class OpDispatcher:
606 id_str = "scalar"610 id_str = "scalar"
607 if not isinstance(arg, Tensor):611 if not isinstance(arg, Tensor):
608 id_str = str(arg)612 id_str = str(arg)
609- cache_key.layout_ids.append(id_str)613+ cache_key_values.append(id_str)
610 extra_args.append(arg)614 extra_args.append(arg)
611 input_layouts.append(None)615 input_layouts.append(None)
612 input_args.append(arg)616 input_args.append(arg)
613 else:617 else:
614 layout = arg.layout618 layout = arg.layout
615 layout_id = layout.compact_str619 layout_id = layout.compact_str
616- cache_key.layout_ids.append(str(layout_id))620+ cache_key_values.append(str(layout_id))
617 input_layouts.append(layout)621 input_layouts.append(layout)
618 if isinstance(arg, DTensor):622 if isinstance(arg, DTensor):
619 input_args.append(arg.to_local())623 input_args.append(arg.to_local())
@@ -625,7 +629,7 @@ class OpDispatcher:
625 else:629 else:
626 input_shape = arg.shape630 input_shape = arg.shape
627 input_shapes.append(input_shape)631 input_shapes.append(input_shape)
628- cache_key.layout_ids.append(str(input_shape))632+ cache_key_values.append(str(input_shape))
629 633 
630 for k, val in kwargs.items():634 for k, val in kwargs.items():
631 if val is None:635 if val is None:
@@ -635,13 +639,13 @@ class OpDispatcher:
635 id_str = "scalar"639 id_str = "scalar"
636 if not isinstance(val, Tensor):640 if not isinstance(val, Tensor):
637 id_str = str(val)641 id_str = str(val)
638- cache_key.layout_ids.append(id_str)642+ cache_key_values.append(id_str)
639 extra_args.append(val)643 extra_args.append(val)
640 input_layouts.append(None)644 input_layouts.append(None)
641 else:645 else:
642 layout = val.layout646 layout = val.layout
643 layout_id = layout.compact_str647 layout_id = layout.compact_str
644- cache_key.layout_ids.append(str(layout_id))648+ cache_key_values.append(str(layout_id))
645 input_layouts.append(layout)649 input_layouts.append(layout)
646 if isinstance(val, DTensor):650 if isinstance(val, DTensor):
647 input_kwargs[k] = val.to_local()651 input_kwargs[k] = val.to_local()
@@ -651,9 +655,9 @@ class OpDispatcher:
651 else:655 else:
652 input_shape = val.shape656 input_shape = val.shape
653 input_shapes.append(input_shape)657 input_shapes.append(input_shape)
654- cache_key.layout_ids.append(str(input_shape))658+ cache_key_values.append(str(input_shape))
655 659 
656- return input_layouts, input_shapes, extra_args, input_args, input_kwargs660+ return input_layouts, input_shapes, extra_args, input_args, input_kwargs, cache_key_values
657 661 
658 def _with_layout_infer_with_shape(self, func: callable, *args, **kwargs) -> Tensor:662 def _with_layout_infer_with_shape(self, func: callable, *args, **kwargs) -> Tensor:
659 """_with_layout_infer_with_shape"""663 """_with_layout_infer_with_shape"""
@@ -665,9 +669,9 @@ class OpDispatcher:
665 packed_call = (args[0], args[1])669 packed_call = (args[0], args[1])
666 args = tuple(args[2])670 args = tuple(args[2])
667 671 
668- cache_key = LayoutCacheKey([])672+ (input_layouts, input_shapes, extra_args, input_args,
669- input_layouts, input_shapes, extra_args, input_args, input_kwargs = \673+ input_kwargs, cache_key_values) = OpDispatcher._process_args_and_kwargs_with_shape(args, kwargs)
670- OpDispatcher._process_args_and_kwargs_with_shape(args, kwargs, cache_key)674+ cache_key = LayoutCacheKey(cache_key_values)
671 675 
672 cache_manager = LayoutCacheManager.get_instance()676 cache_manager = LayoutCacheManager.get_instance()
673 layout_cache = cache_manager.get_layout_cache()677 layout_cache = cache_manager.get_layout_cache()
@@ -720,22 +724,19 @@ class OpDispatcher:
720 end = args[2]724 end = args[2]
721 725 
722 # input layout726 # input layout
723- cache_key = LayoutCacheKey([])
724 input_layouts = []727 input_layouts = []
725 728 
726 layout = input_tensor.layout729 layout = input_tensor.layout
727 global_shape = input_tensor.shape730 global_shape = input_tensor.shape
728 input_layouts.append(layout)731 input_layouts.append(layout)
729 layout_id = layout.compact_str732 layout_id = layout.compact_str
730- cache_key.layout_ids.append(str(layout_id))
731 733 
732 extra_args = []734 extra_args = []
733 extra_args.append(begin)735 extra_args.append(begin)
734 extra_args.append(end)736 extra_args.append(end)
735 extra_args.append(global_shape)737 extra_args.append(global_shape)
736- cache_key.layout_ids.append(str(begin))738+ cache_key_values = [str(layout_id), str(begin), str(end), str(global_shape)]
737- cache_key.layout_ids.append(str(end))739+ cache_key = LayoutCacheKey(cache_key_values)
738- cache_key.layout_ids.append(str(global_shape))
739 740 
740 cache_manager = LayoutCacheManager.get_instance()741 cache_manager = LayoutCacheManager.get_instance()
741 layout_cache = cache_manager.get_layout_cache()742 layout_cache = cache_manager.get_layout_cache()
@@ -918,6 +919,13 @@ class OpDispatcher:
918 if op_name not in self.layout_infer_ops:919 if op_name not in self.layout_infer_ops:
919 raise RuntimeError(f"Operator {op_name} dose not contain parallel layout infer func.")920 raise RuntimeError(f"Operator {op_name} dose not contain parallel layout infer func.")
920 921 
922+ cache_manager = LayoutCacheManager.get_instance()
923+ distribute_op = cache_manager.distributed_op(op_name)
924+ 
925+ result = distribute_op.preprocess(args, kwargs)
926+ if result is not None:
927+ return self._dispatch_new(op_call, distribute_op, result)
928+ 
921 suffix = self.layout_infer_ops[op_name].get('infer_layout_suffix', '')929 suffix = self.layout_infer_ops[op_name].get('infer_layout_suffix', '')
922 if not suffix:930 if not suffix:
923 return self._with_layout_infer(op_call, *args, **kwargs)931 return self._with_layout_infer(op_call, *args, **kwargs)
@@ -927,7 +935,53 @@ class OpDispatcher:
927 raise RuntimeError(f"Operator {op_name} specified wrong suffix in parallel yaml.")935 raise RuntimeError(f"Operator {op_name} specified wrong suffix in parallel yaml.")
928 return getattr(self, handler_name)(op_call, *args, **kwargs)936 return getattr(self, handler_name)(op_call, *args, **kwargs)
929 937 
930- def dispatch(self, op_call: callable, args: tuple, kwargs: dict):938+ def _wrap_output(self, py_output, output_layouts) -> Tensor:
939+ if isinstance(py_output, (tuple, list)):
940+ if len(py_output) != len(output_layouts):
941+ raise RuntimeError(
942+ f"Output tuple size ({len(py_output)}) "
943+ f"does not match layout tuple size ({len(output_layouts)})")
944+ return tuple(
945+ DTensor.from_local(item, layout.mesh, layout.alias_placements)
946+ for item, layout in zip(py_output, output_layouts))
947+ return DTensor.from_local(
948+ py_output, output_layouts[0].mesh, output_layouts[0].alias_placements)
949+ 
950+ def _dispatch_new(self, func, distribute_op, result) -> Tensor:
951+ """New dispatch flow using preprocess result.
952+ 
953+ Args:
954+ func: Original function.
955+ distribute_op: Distributed operation instance.
956+ result: Preprocessed result (local_args, local_kwargs, cache_values).
957+ 
958+ Returns:
959+ Tensor: Dispatched result as DTensor.
960+ """
961+ local_args, local_kwargs, cache_values = result
962+ cache_key = LayoutCacheKey.from_cache_values(cache_values)
963+ func_name = platform.get_op_name(func)
964+ cache_manager = LayoutCacheManager.get_instance()
965+ layout_cache = cache_manager.get_layout_cache()
966+ if func_name not in layout_cache:
967+ layout_cache[func_name] = {}
968+ op_layout_cache = layout_cache[func_name]
969+ if cache_key in op_layout_cache:
970+ infer_result, op_impl = op_layout_cache[cache_key]
971+ else:
972+ infer_result = distribute_op.infer_layout(cache_values)
973+ op_impl = distribute_op.get_expand_impl(func, infer_result, cache_values)
974+ op_layout_cache[cache_key] = (infer_result, op_impl)
975+ output_layouts, extra_info = infer_result
976+ if op_impl is None:
977+ op_impl = func
978+ if extra_info is not None:
979+ py_output = op_impl(*local_args, *extra_info, **local_kwargs)
980+ else:
981+ py_output = op_impl(*local_args, **local_kwargs)
982+ return self._wrap_output(py_output, output_layouts)
983+ 
984+ def dispatch(self, op_call: callable, args: tuple[object, ...], kwargs: dict[str, object]):
931 """Route an op call through the appropriate DTensor dispatch path.985 """Route an op call through the appropriate DTensor dispatch path.
932 986 
933 Args:987 Args:
@@ -26,7 +26,7 @@ class ActivationWithAxisDistributedOp(DistributedOp):
26 Inherits from DistributedOp and provides activation-with-axis specific implementations.26 Inherits from DistributedOp and provides activation-with-axis specific implementations.
27 """27 """
28 28 
29- def infer_layout(self, layouts, extra_args):29+ def infer_layout(self, layouts, extra_args=None):
30 """30 """
31 Infer output layouts for activation-with-axis operations.31 Infer output layouts for activation-with-axis operations.
32 32 
@@ -22,7 +22,7 @@ from .parallel_ops import DistributedOp
22class ArgMaxWithValueDistributedOp(DistributedOp):22class ArgMaxWithValueDistributedOp(DistributedOp):
23 """Distributed implementation for ArgMaxWithValue operator."""23 """Distributed implementation for ArgMaxWithValue operator."""
24 24 
25- def infer_layout(self, layouts, extra_args):25+ def infer_layout(self, layouts, extra_args=None):
26 """26 """
27 Infer output layout for ArgMaxWithValue operator.27 Infer output layout for ArgMaxWithValue operator.
28 Args:28 Args:
@@ -22,7 +22,7 @@ from .parallel_ops import DistributedOp
22class ConcatDistributedOp(DistributedOp):22class ConcatDistributedOp(DistributedOp):
23 """Distributed implementation for Concat."""23 """Distributed implementation for Concat."""
24 24 
25- def infer_layout(self, layouts, extra_args):25+ def infer_layout(self, layouts, extra_args=None):
26 """26 """
27 Infer output layout for Concat and normalize the concatenation dimension.27 Infer output layout for Concat and normalize the concatenation dimension.
28 Raises an error if the specified concatenation dimension is sharded.28 Raises an error if the specified concatenation dimension is sharded.
@@ -132,7 +132,7 @@ class Conv3dDistributedOp(DistributedOp):
132 132 
133 return output_layout133 return output_layout
134 134 
135- def get_expand_impl(self, func, output_layout, layouts, extra_args):135+ def get_expand_impl(self, func, infer_result, layouts, extra_args=None):
136 """136 """
137 Get expand implementation for the operator.137 Get expand implementation for the operator.
138 Intercepts the execution to handle Grouped Convolution with Column Parallelism.138 Intercepts the execution to handle Grouped Convolution with Column Parallelism.
@@ -31,7 +31,7 @@ class ElementWiseDistributedOp(DistributedOp):
31 op_name (str): Name of the operator to register.31 op_name (str): Name of the operator to register.
32 """32 """
33 33 
34- def infer_layout(self, layouts, extra_args):34+ def infer_layout(self, layouts, extra_args=None):
35 """35 """
36 Infer output layouts for element-wise operations with broadcasting support.36 Infer output layouts for element-wise operations with broadcasting support.
37 37 
@@ -621,8 +621,7 @@ class AddDistributedOp(ElementWiseWithPartialDistributedOp):
621 which is useful for operations like gradient accumulation where partial621 which is useful for operations like gradient accumulation where partial
622 results need to be preserved through the computation graph.622 results need to be preserved through the computation graph.
623 """623 """
624- 624+ def get_expand_impl(self, func, infer_result, layouts, extra_args=None):
625- def get_expand_impl(self, func, output_layout, layouts, extra_args):
626 """625 """
627 Get expand implementation for the operator626 Get expand implementation for the operator
628 """627 """
@@ -633,9 +632,9 @@ class AddDistributedOp(ElementWiseWithPartialDistributedOp):
633 632 
634 if x1_partial != x2_partial:633 if x1_partial != x2_partial:
635 scaling_factor = 1634 scaling_factor = 1
636- for i, partial_type in enumerate(output_layout.partial):635+ for i, partial_type in enumerate(infer_result.partial):
637 if partial_type == "sum":636 if partial_type == "sum":
638- scaling_factor *= output_layout.mesh_shape[i]637+ scaling_factor *= infer_result.mesh_shape[i]
639 elif partial_type is not None:638 elif partial_type is not None:
640 raise ValueError(639 raise ValueError(
641 f"For {self.op_name}, inputs partial status should be 'sum' or None, "640 f"For {self.op_name}, inputs partial status should be 'sum' or None, "
@@ -121,7 +121,7 @@ class EmbeddingDistributedOp(DistributedOp):
121 121 
122 return local_input, mask_int122 return local_input, mask_int
123 123 
124- def get_expand_impl(self, func, output_layout, layouts, extra_args):124+ def get_expand_impl(self, func, infer_result, layouts, extra_args=None):
125 """125 """
126 Returns the execution implementation wrapper.126 Returns the execution implementation wrapper.
127 Helper functions are used to keep Cyclomatic Complexity (CCN) low.127 Helper functions are used to keep Cyclomatic Complexity (CCN) low.
@@ -22,7 +22,7 @@ from .parallel_ops import DistributedOp
22class ExpandDimsDistributedOp(DistributedOp):22class ExpandDimsDistributedOp(DistributedOp):
23 """Distributed implementation for ExpandDims operator."""23 """Distributed implementation for ExpandDims operator."""
24 24 
25- def infer_layout(self, layouts, extra_args):25+ def infer_layout(self, layouts, extra_args=None):
26 """26 """
27 Infer output layout for ExpandDims.27 Infer output layout for ExpandDims.
28 28 
@@ -24,7 +24,7 @@ from .parallel_ops import DistributedOp
24class IndexSelectDistributedOp(DistributedOp):24class IndexSelectDistributedOp(DistributedOp):
25 """Distributed implementation for Index Select operator."""25 """Distributed implementation for Index Select operator."""
26 26 
27- def infer_layout(self, layouts, extra_args):27+ def infer_layout(self, layouts, extra_args=None):
28 """28 """
29 Infer output layouts for Index Select operations.29 Infer output layouts for Index Select operations.
30 30 
@@ -99,8 +99,7 @@ class IndexSelectDistributedOp(DistributedOp):
99 99 
100 return output_layout100 return output_layout
101 101 
102- 102+ def get_expand_impl(self, func, infer_result, layouts, extra_args=None):
103- def get_expand_impl(self, func, output_layout, layouts, extra_args):
104 """103 """
105 Get the expanded execution implementation for Index Select.104 Get the expanded execution implementation for Index Select.
106 """105 """
@@ -179,7 +178,7 @@ class GatherDDistributedOp(DistributedOp):
179 - Output inherits the sharding pattern of the input tensor178 - Output inherits the sharding pattern of the input tensor
180 """179 """
181 180 
182- def infer_layout(self, layouts, extra_args):181+ def infer_layout(self, layouts, extra_args=None):
183 """182 """
184 Infer output layouts for GatherD operations.183 Infer output layouts for GatherD operations.
185 Args:184 Args:
@@ -267,7 +266,7 @@ class GatherDDistributedOp(DistributedOp):
267 output_layout.update_compact_str()266 output_layout.update_compact_str()
268 return output_layout267 return output_layout
269 268 
270- def get_expand_impl(self, func, output_layout, layouts, extra_args):269+ def get_expand_impl(self, func, infer_result, layouts, extra_args=None):
271 """270 """
272 Returns the execution implementation wrapper for distributed GatherD.271 Returns the execution implementation wrapper for distributed GatherD.
273 272
@@ -337,7 +336,7 @@ class GatherDDistributedOp(DistributedOp):
337class GatherNdDistributedOp(DistributedOp):336class GatherNdDistributedOp(DistributedOp):
338 """Distributed implementation for GatherNd operator."""337 """Distributed implementation for GatherNd operator."""
339 338 
340- def infer_layout(self, layouts, extra_args):339+ def infer_layout(self, layouts, extra_args=None):
341 """340 """
342 Infer output layout for GatherNd.341 Infer output layout for GatherNd.
343 342 
@@ -22,7 +22,7 @@ from .parallel_ops import DistributedOp
22 22 
23class MatMulExtDistributedOp(DistributedOp):23class MatMulExtDistributedOp(DistributedOp):
24 """Distributed implementation for MatMul operator."""24 """Distributed implementation for MatMul operator."""
25- def infer_layout(self, layouts, extra_args):25+ def infer_layout(self, layouts, extra_args=None):
26 """26 """
27 Infer output layout for MatMul operator.27 Infer output layout for MatMul operator.
28 28 
@@ -82,7 +82,7 @@ class MatMulExtDistributedOp(DistributedOp):
82 82 
83class MatMulDistributedOp(DistributedOp):83class MatMulDistributedOp(DistributedOp):
84 """Distributed implementation for MatMul operator."""84 """Distributed implementation for MatMul operator."""
85- def infer_layout(self, layouts, extra_args):85+ def infer_layout(self, layouts, extra_args=None):
86 """86 """
87 Infer output layout for MatMul operator.87 Infer output layout for MatMul operator.
88 88 
@@ -273,7 +273,7 @@ class BatchMatMulExtDistributedOp(BaseBatchMatMulDistributedOp):
273class BatchMatMulDistributedOp(BaseBatchMatMulDistributedOp):273class BatchMatMulDistributedOp(BaseBatchMatMulDistributedOp):
274 """Distributed implementation for BatchMatMul operator."""274 """Distributed implementation for BatchMatMul operator."""
275 275 
276- def infer_layout(self, layouts, extra_args):276+ def infer_layout(self, layouts, extra_args=None):
277 """277 """
278 Infer output layout for BatchMatMul operator. Inputs shape are x=[b, n, m] and w=[b, m, p].278 Infer output layout for BatchMatMul operator. Inputs shape are x=[b, n, m] and w=[b, m, p].
279 279 
@@ -341,9 +341,34 @@ class BatchMatMulDistributedOp(BaseBatchMatMulDistributedOp):
341 return self._build_output_layout(x_layout, merged_batch, x_n, w_p, x_contract)341 return self._build_output_layout(x_layout, merged_batch, x_n, w_p, x_contract)
342 342 
343 343 
344+def _normalize_linear_args(x, weight, bias=None):
345+ return (x, weight), {'bias': bias}
346+ 
347+ 
344class LinearDistributedOp(DistributedOp):348class LinearDistributedOp(DistributedOp):
345 """Distributed implementation for Linear operator."""349 """Distributed implementation for Linear operator."""
346- def infer_layout(self, layouts, extra_args):350+ def preprocess(self, args, kwargs):
351+ """
352+ Preprocess arguments for Linear operator.
353+ 
354+ Args:
355+ args (tuple): Input arguments (x, w)
356+ kwargs (dict): Keyword arguments (bias)
357+ 
358+ Returns:
359+ tuple: Preprocessed arguments (x_local, w_local), local kwargs, cache values
360+ """
361+ args, kwargs = _normalize_linear_args(*args, **kwargs)
362+ x_tensor = args[0]
363+ w_tensor = args[1]
364+ bias = kwargs['bias']
365+ local_args = (x_tensor.to_local(), w_tensor.to_local())
366+ local_kwargs = {'bias': bias.to_local() if hasattr(bias, '_layout') else bias}
367+ cache_values = [x_tensor.layout, w_tensor.layout,
368+ bias.layout if hasattr(bias, '_layout') else None]
369+ return local_args, local_kwargs, cache_values
370+ 
371+ def infer_layout(self, cache_values):
347 """372 """
348 Infer output layout for MatMul operator.373 Infer output layout for MatMul operator.
349 374 
@@ -357,15 +382,16 @@ class LinearDistributedOp(DistributedOp):
357 Args:382 Args:
358 x_layout (Layout): Layout of input x383 x_layout (Layout): Layout of input x
359 w_layout (Layout): Layout of input w384 w_layout (Layout): Layout of input w
385+ bias_layout (Layout): Layout of input bias
360 386 
361 Returns:387 Returns:
362 tuple: Layout for output tensor388 tuple: Layout for output tensor
363 """389 """
364- if len(layouts) != 3:390+ if len(cache_values) != 3:
365- raise ValueError(f"Linear layout length is not 3, but {len(layouts)}")391+ raise ValueError(f"Linear cache_values length is not 3, but {len(cache_values)}")
366- x_layout = layouts[0]392+ x_layout = cache_values[0]
367- w_layout = layouts[1]393+ w_layout = cache_values[1]
368- bias_layout = layouts[2]394+ bias_layout = cache_values[2]
369 if not x_layout or not w_layout:395 if not x_layout or not w_layout:
370 raise ValueError(f"x_layout : {x_layout}, w_layout : {w_layout}")396 raise ValueError(f"x_layout : {x_layout}, w_layout : {w_layout}")
371 x_mesh_shape = x_layout.mesh_shape397 x_mesh_shape = x_layout.mesh_shape
@@ -403,14 +429,16 @@ class LinearDistributedOp(DistributedOp):
403 else:429 else:
404 out_layout.set_partial_by_dev_axis(x_map[x_contract_dim], 'sum')430 out_layout.set_partial_by_dev_axis(x_map[x_contract_dim], 'sum')
405 431 
406- return out_layout432+ return ((out_layout,), None)
407 433 
408- def get_expand_impl(self, func, output_layout, layouts, extra_args):434+ def get_expand_impl(self, func, infer_result, cache_values):
409 """435 """
410 Get expand implementation for the operator436 Get expand implementation for the operator
411 """437 """
412- x_layout = layouts[0]438+ output_layouts, _ = infer_result
413- bias_layout = layouts[2]439+ output_layout = output_layouts[0]
440+ x_layout = cache_values[0]
441+ bias_layout = cache_values[2]
414 x_map = x_layout.alias_tensor_map442 x_map = x_layout.alias_tensor_map
415 x_contract_dim = len(x_map) - 1443 x_contract_dim = len(x_map) - 1
416 scaling_factor = 1444 scaling_factor = 1
@@ -19,13 +19,13 @@
19import copy19import copy
20import warnings20import warnings
21 21 
22-from typing import List, Tuple, Optional, Any22+from typing import List, Tuple, Optional
23from hyper_parallel.core.shard.ops.parallel_npu_flash_attention_score import ( # pylint: disable=C041523from hyper_parallel.core.shard.ops.parallel_npu_flash_attention_score import ( # pylint: disable=C0415
24 _get_lb_override,24 _get_lb_override,
25)25)
26from hyper_parallel.core.dtensor.layout import Layout26from hyper_parallel.core.dtensor.layout import Layout
27from hyper_parallel.core.dtensor.placement_types import Shard, Replicate27from hyper_parallel.core.dtensor.placement_types import Shard, Replicate
28-from hyper_parallel.core.shard.ops.parallel_ops_register import register_distributed_op28+from hyper_parallel.core.shard.ops.parallel_ops import DistributedOp
29from hyper_parallel.platform import get_platform29from hyper_parallel.platform import get_platform
30 30 
31platform = get_platform()31platform = get_platform()
@@ -85,13 +85,11 @@ def _resolve_input_layout(input_layout) -> str:
85 return str(input_layout)85 return str(input_layout)
86 86 
87 87 
88-class FlashAttentionScoreDistributedOp:88+class FlashAttentionScoreDistributedOp(DistributedOp):
89 """Distributed operator for mindspore.ops.flash_attention_score."""89 """Distributed operator for mindspore.ops.flash_attention_score."""
90 90 
91 def __init__(self, op_name: str):91 def __init__(self, op_name: str):
92- self.op_name = op_name92+ super().__init__(op_name)
93- register_distributed_op(op_name, self)
94- 
95 self._layout_dims = {93 self._layout_dims = {
96 "BSH": {"batch": 0, "seq": 1, "hidden": 2},94 "BSH": {"batch": 0, "seq": 1, "hidden": 2},
97 "BNSD": {"batch": 0, "head": 1, "seq": 2, "dim": 3},95 "BNSD": {"batch": 0, "head": 1, "seq": 2, "dim": 3},
@@ -396,10 +394,10 @@ class FlashAttentionScoreDistributedOp:
396 warnings.warn("Detected Alibi positional encoding compression scenario")394 warnings.warn("Detected Alibi positional encoding compression scenario")
397 395 
398 def infer_layout(396 def infer_layout(
399- self, input_layouts: List[Optional[Layout]], extra_args: List[Any]397+ self, layouts: List[Optional[Layout]], extra_args: Optional[dict] = None
400 ) -> Tuple[Layout, ...]:398 ) -> Tuple[Layout, ...]:
401 """Infer output layouts from input layouts and scalar params in extra_args."""399 """Infer output layouts from input layouts and scalar params in extra_args."""
402- query_layout = input_layouts[0]400+ query_layout = layouts[0]
403 if query_layout is None:401 if query_layout is None:
404 raise ValueError("Query layout cannot be None")402 raise ValueError("Query layout cannot be None")
405 403 
@@ -793,13 +791,7 @@ class FlashAttentionScoreDistributedOp:
793 return (sparse_mode, pre_tokens, next_tokens,791 return (sparse_mode, pre_tokens, next_tokens,
794 adjusted_actual_seq_qlen, adjusted_actual_seq_kvlen)792 adjusted_actual_seq_qlen, adjusted_actual_seq_kvlen)
795 793 
796- def get_expand_impl(794+ def get_expand_impl(self, func, infer_result, layouts, extra_args=None):
797- self,
798- func: callable,
799- output_layouts: Tuple[Layout, ...],
800- input_layouts: List[Optional[Layout]],
801- extra_args: List[Any],
802- ) -> Optional[callable]:
803 """Create expanded implementation.795 """Create expanded implementation.
804 796 
805 extra_args contains 8 scalar params extracted by _process_args_and_kwargs:797 extra_args contains 8 scalar params extracted by _process_args_and_kwargs:
@@ -808,13 +800,13 @@ class FlashAttentionScoreDistributedOp:
808 The returned expanded_impl receives 18 positional args matching the800 The returned expanded_impl receives 18 positional args matching the
809 pyboost flat arg list (10 tensor args + 8 scalar args).801 pyboost flat arg list (10 tensor args + 8 scalar args).
810 """802 """
811- query_layout = input_layouts[0]803+ query_layout = layouts[0]
812 if query_layout is None:804 if query_layout is None:
813 return None805 return None
814 806 
815- if len(input_layouts) >= 3:807+ if len(layouts) >= 3:
816- key_layout_check = input_layouts[1]808+ key_layout_check = layouts[1]
817- value_layout_check = input_layouts[2]809+ value_layout_check = layouts[2]
818 810 
819 if (key_layout_check is not None and value_layout_check is not None and811 if (key_layout_check is not None and value_layout_check is not None and
820 hasattr(key_layout_check, 'tensor_map') and812 hasattr(key_layout_check, 'tensor_map') and
@@ -842,7 +834,7 @@ class FlashAttentionScoreDistributedOp:
842 p_pre_tokens, p_next_tokens, p_inner_precise,834 p_pre_tokens, p_next_tokens, p_inner_precise,
843 p_input_layout, p_sparse_mode,835 p_input_layout, p_sparse_mode,
844 ):836 ):
845- key_layout = input_layouts[1]837+ key_layout = layouts[1]
846 self._validate_sharding_consistency(838 self._validate_sharding_consistency(
847 query_layout, key_layout, input_layout839 query_layout, key_layout, input_layout
848 )840 )
@@ -23,7 +23,7 @@ from .parallel_ops import DistributedOp
23class MultinomialDistributedOp(DistributedOp):23class MultinomialDistributedOp(DistributedOp):
24 """Distributed implementation for Multinomial operator."""24 """Distributed implementation for Multinomial operator."""
25 25 
26- def infer_layout(self, layouts, extra_args):26+ def infer_layout(self, layouts, extra_args=None):
27 """27 """
28 Infer output layout for Multinomial operator.28 Infer output layout for Multinomial operator.
29 29 
@@ -24,7 +24,7 @@ from .parallel_ops import DistributedOp
24class NewOnesDistributedOp(DistributedOp):24class NewOnesDistributedOp(DistributedOp):
25 """Distributed implementation for new_ones operator."""25 """Distributed implementation for new_ones operator."""
26 26 
27- def infer_layout(self, layouts, extra_args):27+ def infer_layout(self, layouts, extra_args=None):
28 """28 """
29 Infer output layout for new_ones operator.29 Infer output layout for new_ones operator.
30 30 
@@ -20,10 +20,10 @@ import copy
20import threading20import threading
21import warnings21import warnings
22 22 
23-from typing import List, Tuple, Optional, Any23+from typing import List, Tuple, Optional
24from hyper_parallel.core.dtensor.layout import Layout24from hyper_parallel.core.dtensor.layout import Layout
25from hyper_parallel.core.dtensor.placement_types import Shard, Replicate25from hyper_parallel.core.dtensor.placement_types import Shard, Replicate
26-from hyper_parallel.core.shard.ops.parallel_ops_register import register_distributed_op26+from hyper_parallel.core.shard.ops.parallel_ops import DistributedOp
27from hyper_parallel.platform import get_platform27from hyper_parallel.platform import get_platform
28 28 
29platform = get_platform()29platform = get_platform()
@@ -70,13 +70,11 @@ def _get_lb_override() -> Tuple[Optional[int], Optional[int]]:
70 return getattr(_LB_OVERRIDE, 'split_id', None), getattr(_LB_OVERRIDE, 'split_num', None)70 return getattr(_LB_OVERRIDE, 'split_id', None), getattr(_LB_OVERRIDE, 'split_num', None)
71 71 
72 72 
73-class FlashAttentionScoreDistributedOp:73+class NPUFlashAttentionScoreDistributedOp(DistributedOp):
74 """Distributed operator for torch_npu.npu_fusion_attention."""74 """Distributed operator for torch_npu.npu_fusion_attention."""
75 75 
76 def __init__(self, op_name: str):76 def __init__(self, op_name: str):
77- self.op_name = op_name77+ super().__init__(op_name)
78- register_distributed_op(op_name, self)
79- 
80 self._layout_dims = {78 self._layout_dims = {
81 "BSH": {"batch": 0, "seq": 1, "hidden": 2},79 "BSH": {"batch": 0, "seq": 1, "hidden": 2},
82 "BNSD": {"batch": 0, "head": 1, "seq": 2, "dim": 3},80 "BNSD": {"batch": 0, "head": 1, "seq": 2, "dim": 3},
@@ -458,16 +456,16 @@ class FlashAttentionScoreDistributedOp:
458 warnings.warn("Detected Alibi positional encoding compression scenario")456 warnings.warn("Detected Alibi positional encoding compression scenario")
459 457 
460 def infer_layout(458 def infer_layout(
461- self, input_layouts: List[Optional[Layout]], extra_args: List[Any]459+ self, layouts: List[Optional[Layout]], extra_args: Optional[dict] = None
462 ) -> Tuple[Layout, ...]:460 ) -> Tuple[Layout, ...]:
463 """Infer output layouts."""461 """Infer output layouts."""
464- query_layout = input_layouts[0]462+ query_layout = layouts[0]
465 if query_layout is None:463 if query_layout is None:
466 raise ValueError("Query layout cannot be None")464 raise ValueError("Query layout cannot be None")
467 465 
468 attention_out_layout = copy.deepcopy(query_layout)466 attention_out_layout = copy.deepcopy(query_layout)
469 if attention_out_layout.placements is None and attention_out_layout.tensor_map is not None:467 if attention_out_layout.placements is None and attention_out_layout.tensor_map is not None:
470- attention_out_placements = FlashAttentionScoreDistributedOp._tensor_map_to_placements(468+ attention_out_placements = NPUFlashAttentionScoreDistributedOp._tensor_map_to_placements(
471 attention_out_layout, attention_out_layout.tensor_map469 attention_out_layout, attention_out_layout.tensor_map
472 )470 )
473 attention_out_layout.set_placements(attention_out_placements)471 attention_out_layout.set_placements(attention_out_placements)
@@ -504,7 +502,7 @@ class FlashAttentionScoreDistributedOp:
504 softmax_sum_layout = copy.deepcopy(softmax_layout)502 softmax_sum_layout = copy.deepcopy(softmax_layout)
505 softmax_out_layout = self._create_replicated_scalar_layout(query_layout)503 softmax_out_layout = self._create_replicated_scalar_layout(query_layout)
506 if softmax_out_layout.placements is None and softmax_out_layout.tensor_map is not None:504 if softmax_out_layout.placements is None and softmax_out_layout.tensor_map is not None:
507- softmax_out_placements = FlashAttentionScoreDistributedOp._tensor_map_to_placements(505+ softmax_out_placements = NPUFlashAttentionScoreDistributedOp._tensor_map_to_placements(
508 softmax_out_layout, softmax_out_layout.tensor_map506 softmax_out_layout, softmax_out_layout.tensor_map
509 )507 )
510 softmax_out_layout.set_placements(softmax_out_placements)508 softmax_out_layout.set_placements(softmax_out_placements)
@@ -540,7 +538,7 @@ class FlashAttentionScoreDistributedOp:
540 )538 )
541 539 
542 softmax_layout.set_tensor_map(softmax_tensor_map)540 softmax_layout.set_tensor_map(softmax_tensor_map)
543- softmax_placements = FlashAttentionScoreDistributedOp._tensor_map_to_placements(softmax_layout, softmax_tensor_map)541+ softmax_placements = NPUFlashAttentionScoreDistributedOp._tensor_map_to_placements(softmax_layout, softmax_tensor_map)
544 softmax_layout.set_placements(softmax_placements)542 softmax_layout.set_placements(softmax_placements)
545 543 
546 return softmax_layout544 return softmax_layout
@@ -560,7 +558,7 @@ class FlashAttentionScoreDistributedOp:
560 558 
561 softmax_layout = Layout.from_device_mesh(query_layout.mesh)559 softmax_layout = Layout.from_device_mesh(query_layout.mesh)
562 softmax_layout.set_tensor_map(softmax_tensor_map)560 softmax_layout.set_tensor_map(softmax_tensor_map)
563- softmax_placements = FlashAttentionScoreDistributedOp._tensor_map_to_placements(softmax_layout, softmax_tensor_map)561+ softmax_placements = NPUFlashAttentionScoreDistributedOp._tensor_map_to_placements(softmax_layout, softmax_tensor_map)
564 softmax_layout.set_placements(softmax_placements)562 softmax_layout.set_placements(softmax_placements)
565 563 
566 return softmax_layout564 return softmax_layout
@@ -644,8 +642,8 @@ class FlashAttentionScoreDistributedOp:
644 if batch_idx >= len(q_tm) or batch_idx >= len(k_tm):642 if batch_idx >= len(q_tm) or batch_idx >= len(k_tm):
645 return643 return
646 644 
647- q_batch_shard = FlashAttentionScoreDistributedOp._normalize_dim_map(q_tm[batch_idx])645+ q_batch_shard = NPUFlashAttentionScoreDistributedOp._normalize_dim_map(q_tm[batch_idx])
648- k_batch_shard = FlashAttentionScoreDistributedOp._normalize_dim_map(k_tm[batch_idx])646+ k_batch_shard = NPUFlashAttentionScoreDistributedOp._normalize_dim_map(k_tm[batch_idx])
649 647 
650 if q_batch_shard != k_batch_shard:648 if q_batch_shard != k_batch_shard:
651 raise ValueError(649 raise ValueError(
@@ -666,8 +664,8 @@ class FlashAttentionScoreDistributedOp:
666 if hidden_idx >= len(q_tm) or hidden_idx >= len(k_tm):664 if hidden_idx >= len(q_tm) or hidden_idx >= len(k_tm):
667 return665 return
668 666 
669- q_hidden_shard = FlashAttentionScoreDistributedOp._normalize_dim_map(q_tm[hidden_idx])667+ q_hidden_shard = NPUFlashAttentionScoreDistributedOp._normalize_dim_map(q_tm[hidden_idx])
670- k_hidden_shard = FlashAttentionScoreDistributedOp._normalize_dim_map(k_tm[hidden_idx])668+ k_hidden_shard = NPUFlashAttentionScoreDistributedOp._normalize_dim_map(k_tm[hidden_idx])
671 669 
672 if q_hidden_shard != k_hidden_shard:670 if q_hidden_shard != k_hidden_shard:
673 raise ValueError(671 raise ValueError(
@@ -690,8 +688,8 @@ class FlashAttentionScoreDistributedOp:
690 if dim_idx >= len(q_tm) or dim_idx >= len(k_tm):688 if dim_idx >= len(q_tm) or dim_idx >= len(k_tm):
691 return689 return
692 690 
693- q_dim_shard = FlashAttentionScoreDistributedOp._normalize_dim_map(q_tm[dim_idx])691+ q_dim_shard = NPUFlashAttentionScoreDistributedOp._normalize_dim_map(q_tm[dim_idx])
694- k_dim_shard = FlashAttentionScoreDistributedOp._normalize_dim_map(k_tm[dim_idx])692+ k_dim_shard = NPUFlashAttentionScoreDistributedOp._normalize_dim_map(k_tm[dim_idx])
695 693 
696 if q_dim_shard != k_dim_shard:694 if q_dim_shard != k_dim_shard:
697 raise ValueError(695 raise ValueError(
@@ -738,8 +736,8 @@ class FlashAttentionScoreDistributedOp:
738 f" - KV sequence sharding (requires Ring Attention)"736 f" - KV sequence sharding (requires Ring Attention)"
739 )737 )
740 738 
741- q_seq_shard = FlashAttentionScoreDistributedOp._normalize_dim_map(q_tm[seq_dim_idx])739+ q_seq_shard = NPUFlashAttentionScoreDistributedOp._normalize_dim_map(q_tm[seq_dim_idx])
742- k_seq_shard = FlashAttentionScoreDistributedOp._normalize_dim_map(k_tm[seq_dim_idx])740+ k_seq_shard = NPUFlashAttentionScoreDistributedOp._normalize_dim_map(k_tm[seq_dim_idx])
743 741 
744 if q_seq_shard != k_seq_shard:742 if q_seq_shard != k_seq_shard:
745 if input_layout == "TND":743 if input_layout == "TND":
@@ -881,21 +879,15 @@ class FlashAttentionScoreDistributedOp:
881 return (sparse_mode, pre_tockens, next_tockens,879 return (sparse_mode, pre_tockens, next_tockens,
882 adjusted_actual_seq_qlen, adjusted_actual_seq_kvlen)880 adjusted_actual_seq_qlen, adjusted_actual_seq_kvlen)
883 881 
884- def get_expand_impl(882+ def get_expand_impl(self, func, infer_result, layouts, extra_args=None):
885- self,
886- func: callable,
887- output_layouts: Tuple[Layout, ...],
888- input_layouts: List[Optional[Layout]],
889- extra_args: List[Any],
890- ) -> Optional[callable]:
891 """Create expanded implementation."""883 """Create expanded implementation."""
892- query_layout = input_layouts[0]884+ query_layout = layouts[0]
893 if query_layout is None:885 if query_layout is None:
894 return None886 return None
895 887 
896- if len(input_layouts) >= 3:888+ if len(layouts) >= 3:
897- key_layout = input_layouts[1]889+ key_layout = layouts[1]
898- value_layout = input_layouts[2]890+ value_layout = layouts[2]
899 891 
900 if (key_layout is not None and value_layout is not None and892 if (key_layout is not None and value_layout is not None and
901 hasattr(key_layout, 'tensor_map') and hasattr(value_layout, 'tensor_map')):893 hasattr(key_layout, 'tensor_map') and hasattr(value_layout, 'tensor_map')):
@@ -927,7 +919,7 @@ class FlashAttentionScoreDistributedOp:
927 gen_mask_parallel=True,919 gen_mask_parallel=True,
928 sync=False920 sync=False
929 ):921 ):
930- key_layout = input_layouts[1]922+ key_layout = layouts[1]
931 self._validate_sharding_consistency(query_layout, key_layout, input_layout)923 self._validate_sharding_consistency(query_layout, key_layout, input_layout)
932 924 
933 is_varlen = input_layout == "TND" and actual_seq_qlen is not None925 is_varlen = input_layout == "TND" and actual_seq_qlen is not None
@@ -951,7 +943,7 @@ class FlashAttentionScoreDistributedOp:
951 prefix, actual_seq_qlen, actual_seq_kvlen,943 prefix, actual_seq_qlen, actual_seq_kvlen,
952 sparse_mode, gen_mask_parallel, sync944 sparse_mode, gen_mask_parallel, sync
953 )945 )
954- return FlashAttentionScoreDistributedOp._truncate_result(result)946+ return NPUFlashAttentionScoreDistributedOp._truncate_result(result)
955 947 
956 adjusted_head_num = self._adjust_head_num(head_num, head_split_num)948 adjusted_head_num = self._adjust_head_num(head_num, head_split_num)
957 949 
@@ -979,7 +971,7 @@ class FlashAttentionScoreDistributedOp:
979 gen_mask_parallel, sync971 gen_mask_parallel, sync
980 )972 )
981 973 
982- return FlashAttentionScoreDistributedOp._truncate_result(result)974+ return NPUFlashAttentionScoreDistributedOp._truncate_result(result)
983 975 
984 return expanded_impl976 return expanded_impl
985 977 
@@ -1140,7 +1132,7 @@ class FlashAttentionScoreDistributedOp:
1140 if dim_idx >= len(layout.alias_tensor_map):1132 if dim_idx >= len(layout.alias_tensor_map):
1141 return 11133 return 1
1142 1134 
1143- dim_map = FlashAttentionScoreDistributedOp._normalize_dim_map(layout.alias_tensor_map[dim_idx])1135+ dim_map = NPUFlashAttentionScoreDistributedOp._normalize_dim_map(layout.alias_tensor_map[dim_idx])
1144 1136 
1145 if dim_map == "None":1137 if dim_map == "None":
1146 return 11138 return 1
@@ -1151,7 +1143,7 @@ class FlashAttentionScoreDistributedOp:
1151 if isinstance(dim_map, tuple):1143 if isinstance(dim_map, tuple):
1152 total = 11144 total = 1
1153 for axis_name in dim_map:1145 for axis_name in dim_map:
1154- axis_name = FlashAttentionScoreDistributedOp._normalize_dim_map(axis_name)1146+ axis_name = NPUFlashAttentionScoreDistributedOp._normalize_dim_map(axis_name)
1155 if axis_name != "None":1147 if axis_name != "None":
1156 total *= layout.mesh.get_device_num_along_axis(axis_name)1148 total *= layout.mesh.get_device_num_along_axis(axis_name)
1157 return total1149 return total
@@ -1169,7 +1161,7 @@ class FlashAttentionScoreDistributedOp:
1169 if seq_dim_idx >= len(layout.alias_tensor_map):1161 if seq_dim_idx >= len(layout.alias_tensor_map):
1170 return 01162 return 0
1171 1163 
1172- dim_map = FlashAttentionScoreDistributedOp._normalize_dim_map(layout.alias_tensor_map[seq_dim_idx])1164+ dim_map = NPUFlashAttentionScoreDistributedOp._normalize_dim_map(layout.alias_tensor_map[seq_dim_idx])
1173 1165 
1174 if dim_map == "None":1166 if dim_map == "None":
1175 return 01167 return 0
@@ -1183,7 +1175,7 @@ class FlashAttentionScoreDistributedOp:
1183 1175 
1184 if isinstance(dim_map, tuple):1176 if isinstance(dim_map, tuple):
1185 non_none_axes = [1177 non_none_axes = [
1186- ax for ax in dim_map if FlashAttentionScoreDistributedOp._normalize_dim_map(ax) != "None"1178+ ax for ax in dim_map if NPUFlashAttentionScoreDistributedOp._normalize_dim_map(ax) != "None"
1187 ]1179 ]
1188 if len(non_none_axes) == 0:1180 if len(non_none_axes) == 0:
1189 return 01181 return 0
@@ -28,7 +28,7 @@ platform = get_platform()
28class OneHotExtDistributedOp(DistributedOp):28class OneHotExtDistributedOp(DistributedOp):
29 """Distributed implementation for OneHotExt operator."""29 """Distributed implementation for OneHotExt operator."""
30 30 
31- def infer_layout(self, layouts, extra_args):31+ def infer_layout(self, layouts, extra_args=None):
32 """32 """
33 Infer output layout for OneHotExt.33 Infer output layout for OneHotExt.
34 34 
@@ -72,12 +72,12 @@ class OneHotExtDistributedOp(DistributedOp):
72 72 
73 return out_layout73 return out_layout
74 74 
75- def get_expand_impl(self, func, output_layout, layouts, extra_args):75+ def get_expand_impl(self, func, infer_result, layouts, extra_args=None):
76 """Get expanded implementation for OneHotExt operator."""76 """Get expanded implementation for OneHotExt operator."""
77 import mindspore as ms77 import mindspore as ms
78 from mindspore import ops, Tensor78 from mindspore import ops, Tensor
79 79 
80- del output_layout80+ del infer_result
81 81 
82 indices_layout = layouts[0] if layouts else None82 indices_layout = layouts[0] if layouts else None
83 if indices_layout is None:83 if indices_layout is None:
@@ -56,7 +56,24 @@ class DistributedOp:
56 )56 )
57 57 
58 # pylint: disable=W061358 # pylint: disable=W0613
59- def infer_layout(self, layouts, extra_args):59+ def preprocess(self, args: tuple, kwargs: dict):
60+ """
61+ Unified preprocessing: parameter parsing + to_local + cache_values construction.
62+ 
63+ Subclasses override this to participate in the new dispatch flow.
64+ 
65+ Returns:
66+ None: Fall back to legacy dispatch (default).
67+ tuple: (local_args, local_kwargs, cache_values)
68+ - local_args: Local tensor positional arguments (DTensors already to_local'd).
69+ - local_kwargs: Local tensor keyword arguments (DTensors already to_local'd).
70+ - cache_values: Values affecting layout inference (fixed order).
71+ Contains Layout objects (with compact_str) and raw values (int, bool, tuple, etc.).
72+ """
73+ return None
74+ 
75+ # pylint: disable=W0613
76+ def infer_layout(self, layouts, extra_args=None):
60 """77 """
61 Infer output layouts based on input layouts.78 Infer output layouts based on input layouts.
62 79 
@@ -78,8 +95,8 @@ class DistributedOp:
78 return (layouts[0],)95 return (layouts[0],)
79 return None96 return None
80 97 
81- @staticmethod98+ # pylint: disable=W0613
82- def get_expand_impl(func, output_layout, layouts, extra_args):99+ def get_expand_impl(self, func, infer_result, layouts, extra_args=None):
83 """100 """
84 Get expand implementation for the operator101 Get expand implementation for the operator
85 """102 """
@@ -22,7 +22,7 @@ from .parallel_ops import DistributedOp
22class PadDistributedOp(DistributedOp):22class PadDistributedOp(DistributedOp):
23 """Distributed implementation for Pad operator."""23 """Distributed implementation for Pad operator."""
24 24 
25- def infer_layout(self, layouts, extra_args):25+ def infer_layout(self, layouts, extra_args=None):
26 """26 """
27 Infer output layout for Pad operator.27 Infer output layout for Pad operator.
28 28 
@@ -43,7 +43,7 @@ class ReduceExtDistributedOpBase(DistributedOp):
43 partial_type = ["sum"]43 partial_type = ["sum"]
44 self.partial_type = partial_type44 self.partial_type = partial_type
45 45 
46- def infer_layout(self, layouts, extra_args):46+ def infer_layout(self, layouts, extra_args=None):
47 """47 """
48 Infer output layout for reduce operator.48 Infer output layout for reduce operator.
49 49 
@@ -264,7 +264,7 @@ class MaxDistributedOp(ReduceExtDistributedOpBase):
264 def __init__(self, op_name="max"):264 def __init__(self, op_name="max"):
265 super().__init__(op_name, partial_type=["max"])265 super().__init__(op_name, partial_type=["max"])
266 266 
267- def infer_layout(self, layouts, extra_args):267+ def infer_layout(self, layouts, extra_args=None):
268 """268 """
269 Infer output layouts for torch.max.269 Infer output layouts for torch.max.
270 """270 """
@@ -23,7 +23,7 @@ from .parallel_ops import DistributedOp
23class RepeatInterleaveDistributedOp(DistributedOp):23class RepeatInterleaveDistributedOp(DistributedOp):
24 """Distributed implementation for torch.repeat_interleave."""24 """Distributed implementation for torch.repeat_interleave."""
25 25 
26- def infer_layout(self, layouts, extra_args):26+ def infer_layout(self, layouts, extra_args=None):
27 """27 """
28 Infer output layout for RepeatInterleave operator.28 Infer output layout for RepeatInterleave operator.
29 29 
@@ -262,7 +262,7 @@ class ReshapeDistributedOp(DistributedOp):
262 if partial_op is not None and i < len(out_layout.alias_name):262 if partial_op is not None and i < len(out_layout.alias_name):
263 out_layout.set_partial_by_dev_axis(out_layout.alias_name[i], partial_op)263 out_layout.set_partial_by_dev_axis(out_layout.alias_name[i], partial_op)
264 264 
265- def infer_layout(self, layouts, extra_args):265+ def infer_layout(self, layouts, extra_args=None):
266 """266 """
267 Infer output layout for reshape operator.267 Infer output layout for reshape operator.
268 268 
@@ -19,20 +19,20 @@
19import copy19import copy
20import warnings20import warnings
21 21 
22-from typing import List, Tuple, Optional, Any22+from typing import List, Tuple, Optional
23from hyper_parallel.core.shard.ops.parallel_npu_flash_attention_score import ( # pylint: disable=C041523from hyper_parallel.core.shard.ops.parallel_npu_flash_attention_score import ( # pylint: disable=C0415
24 _get_lb_override,24 _get_lb_override,
25)25)
26from hyper_parallel.core.dtensor.layout import Layout26from hyper_parallel.core.dtensor.layout import Layout
27from hyper_parallel.core.dtensor.placement_types import Shard, Replicate27from hyper_parallel.core.dtensor.placement_types import Shard, Replicate
28-from hyper_parallel.core.shard.ops.parallel_ops_register import register_distributed_op28+from hyper_parallel.core.shard.ops.parallel_ops import DistributedOp
29from hyper_parallel.platform import get_platform29from hyper_parallel.platform import get_platform
30 30 
31platform = get_platform()31platform = get_platform()
32Tensor = platform.Tensor32Tensor = platform.Tensor
33 33 
34 34 
35-class ScaledDotProductAttentionDistributedOp:35+class ScaledDotProductAttentionDistributedOp(DistributedOp):
36 """Distributed operator for torch.nn.functional.scaled_dot_product_attention.36 """Distributed operator for torch.nn.functional.scaled_dot_product_attention.
37 37 
38 Input shape: [B, N, S, D] (4D) or [N, S, D] (3D).38 Input shape: [B, N, S, D] (4D) or [N, S, D] (3D).
@@ -45,10 +45,6 @@ class ScaledDotProductAttentionDistributedOp:
45 - Combinations: DP+MP, SP+MP, DP+SP+MP45 - Combinations: DP+MP, SP+MP, DP+SP+MP
46 """46 """
47 47 
48- def __init__(self, op_name: str):
49- self.op_name = op_name
50- register_distributed_op(op_name, self)
51- 
52 @staticmethod48 @staticmethod
53 def _tensor_map_to_placements(base_layout: Layout, tensor_map: tuple) -> tuple:49 def _tensor_map_to_placements(base_layout: Layout, tensor_map: tuple) -> tuple:
54 """Convert tensor_map to placements."""50 """Convert tensor_map to placements."""
@@ -281,10 +277,10 @@ class ScaledDotProductAttentionDistributedOp:
281 return attn_mask, is_causal, key, value277 return attn_mask, is_causal, key, value
282 278 
283 def infer_layout(279 def infer_layout(
284- self, input_layouts: List[Optional[Layout]], extra_args: List[Any]280+ self, layouts: List[Optional[Layout]], extra_args: Optional[dict] = None
285 ) -> Tuple[Layout, ...]:281 ) -> Tuple[Layout, ...]:
286 """Infer output layout. Output has the same layout as query."""282 """Infer output layout. Output has the same layout as query."""
287- query_layout = input_layouts[0]283+ query_layout = layouts[0]
288 if query_layout is None:284 if query_layout is None:
289 raise ValueError("Query layout cannot be None")285 raise ValueError("Query layout cannot be None")
290 286 
@@ -297,21 +293,15 @@ class ScaledDotProductAttentionDistributedOp:
297 293 
298 return attention_out_layout294 return attention_out_layout
299 295 
300- def get_expand_impl(296+ def get_expand_impl(self, func, infer_result, layouts, extra_args=None):
301- self,
302- func: callable,
303- output_layouts: Tuple[Layout, ...],
304- input_layouts: List[Optional[Layout]],
305- extra_args: List[Any],
306- ) -> Optional[callable]:
307 """Create expanded implementation."""297 """Create expanded implementation."""
308- query_layout = input_layouts[0]298+ query_layout = layouts[0]
309 if query_layout is None:299 if query_layout is None:
310 return None300 return None
311 301 
312- if len(input_layouts) >= 3:302+ if len(layouts) >= 3:
313- key_layout = input_layouts[1]303+ key_layout = layouts[1]
314- value_layout = input_layouts[2]304+ value_layout = layouts[2]
315 305 
316 if (key_layout is not None and value_layout is not None and306 if (key_layout is not None and value_layout is not None and
317 hasattr(key_layout, 'tensor_map') and hasattr(value_layout, 'tensor_map')):307 hasattr(key_layout, 'tensor_map') and hasattr(value_layout, 'tensor_map')):
@@ -334,7 +324,7 @@ class ScaledDotProductAttentionDistributedOp:
334 scale=None,324 scale=None,
335 enable_gqa=False,325 enable_gqa=False,
336 ):326 ):
337- key_layout = input_layouts[1] if len(input_layouts) > 1 else None327+ key_layout = layouts[1] if len(layouts) > 1 else None
338 self._validate_sharding_consistency(query_layout, key_layout, dims)328 self._validate_sharding_consistency(query_layout, key_layout, dims)
339 329 
340 split_info = self._get_split_info(query_layout, dims)330 split_info = self._get_split_info(query_layout, dims)
@@ -22,7 +22,7 @@ from .parallel_ops import DistributedOp
22class ScatterUpdateDistributedOp(DistributedOp):22class ScatterUpdateDistributedOp(DistributedOp):
23 """Distributed implementation for ScatterUpdate operator."""23 """Distributed implementation for ScatterUpdate operator."""
24 24 
25- def infer_layout(self, layouts, extra_args):25+ def infer_layout(self, layouts, extra_args=None):
26 """26 """
27 Infer output layout for ScatterUpdate.27 Infer output layout for ScatterUpdate.
28 28 
@@ -54,7 +54,7 @@ class SliceDistributedOp(DistributedOp):
54 f"the begin is {begin}, the end is {end}, the shape is {shape}, layout is {layout.to_dict()}")54 f"the begin is {begin}, the end is {end}, the shape is {shape}, layout is {layout.to_dict()}")
55 return shard_dim55 return shard_dim
56 56 
57- def infer_layout(self, layouts, extra_args):57+ def infer_layout(self, layouts, extra_args=None):
58 """58 """
59 Infer output layout for slice operator. The shard dim must be fully fetched.59 Infer output layout for slice operator. The shard dim must be fully fetched.
60 60 
@@ -22,7 +22,7 @@ from .parallel_ops import DistributedOp
22class SliceExtDistributedOp(DistributedOp):22class SliceExtDistributedOp(DistributedOp):
23 """Distributed implementation for SliceExt operator."""23 """Distributed implementation for SliceExt operator."""
24 24 
25- def infer_layout(self, layouts, extra_args):25+ def infer_layout(self, layouts, extra_args=None):
26 """26 """
27 Infer output layouts for Split operator.27 Infer output layouts for Split operator.
28 28 
@@ -19,32 +19,27 @@ Distributed implementation for Sort operator.
19from .parallel_ops import DistributedOp19from .parallel_ops import DistributedOp
20 20 
21 21 
22+def _normalize_sort_args(x, dim=-1, descending=False, stable=None):
23+ return (x,), {'dim': dim, 'descending': descending, 'stable': stable}
24+ 
25+ 
22class SortDistributedOp(DistributedOp):26class SortDistributedOp(DistributedOp):
23 """Distributed implementation for Sort operator."""27 """Distributed implementation for Sort operator."""
24 28 
25- def infer_layout(self, layouts, extra_args):29+ def preprocess(self, args, kwargs):
26- """30+ args, kwargs = _normalize_sort_args(*args, **kwargs)
27- Infer output layout for Sort operator.31+ input_tensor = args[0]
32+ dim = kwargs['dim']
33+ descending = kwargs['descending']
34+ stable = kwargs['stable']
35+ local_args = (input_tensor.to_local(),)
36+ local_kwargs = {'dim': dim, 'descending': descending, 'stable': stable}
37+ cache_values = [input_tensor.layout, dim]
38+ return local_args, local_kwargs, cache_values
28 39 
29- The sort operator expects the sorting dimension to be fully available on each device40+ def infer_layout(self, cache_values):
30- (i.e., not sharded). If the dimension is sharded, a global sort cannot be performed41+ layout = cache_values[0]
31- locally without redistribution.42+ dim = cache_values[1]
32- 
33- Args:
34- layouts (tuple): Layouts of input tensor.
35- extra_args (tuple): Arguments for the operator. Expected: (dim, descending, stable).
36- If empty, dim defaults to -1.
37- 
38- Returns:
39- tuple: (Layout, Layout) representing the layouts for (values, indices).
40- """
41- layout = layouts[0]
42- 
43- # Parse dim from extra_args if available, otherwise default to -1
44- dim = -1
45- if extra_args:
46- # extra_args[0] corresponds to 'dim' in torch.sort(input, dim, ...)
47- dim = extra_args[0]
48 43 
49 if not isinstance(dim, int):44 if not isinstance(dim, int):
50 raise TypeError(f"For 'sort', dimension must be int, but got {type(dim)}")45 raise TypeError(f"For 'sort', dimension must be int, but got {type(dim)}")
@@ -78,5 +73,4 @@ class SortDistributedOp(DistributedOp):
78 f"Please redistribute the tensor to Replicate status on this dimension before sorting."73 f"Please redistribute the tensor to Replicate status on this dimension before sorting."
79 )74 )
80 75 
81- # The output layouts for 'values' and 'indices' are the same as the input layout76+ return ((layout, layout), None)
82- return (layout, layout)
@@ -23,7 +23,7 @@ from .parallel_ops import DistributedOp
23class SplitWithSizeDistributedOp(DistributedOp):23class SplitWithSizeDistributedOp(DistributedOp):
24 """Distributed implementation for SplitWithSize operator."""24 """Distributed implementation for SplitWithSize operator."""
25 25 
26- def infer_layout(self, layouts, extra_args):26+ def infer_layout(self, layouts, extra_args=None):
27 """27 """
28 Infer output layouts for Split operator.28 Infer output layouts for Split operator.
29 29 
@@ -54,7 +54,7 @@ class SplitWithSizeDistributedOp(DistributedOp):
54class SplitWithSizeViewDistributedOp(DistributedOp):54class SplitWithSizeViewDistributedOp(DistributedOp):
55 """Distributed implementation for SplitWithSizeView operator."""55 """Distributed implementation for SplitWithSizeView operator."""
56 56 
57- def infer_layout(self, layouts, extra_args):57+ def infer_layout(self, layouts, extra_args=None):
58 """58 """
59 Infer output layouts for SplitWithSizeView operator.59 Infer output layouts for SplitWithSizeView operator.
60 60 
@@ -85,7 +85,7 @@ class SplitWithSizeViewDistributedOp(DistributedOp):
85class SplitDistributedOp(DistributedOp):85class SplitDistributedOp(DistributedOp):
86 """Distributed implementation for Split operator."""86 """Distributed implementation for Split operator."""
87 87 
88- def infer_layout(self, layouts, extra_args):88+ def infer_layout(self, layouts, extra_args=None):
89 """89 """
90 Infer output layouts for Split operator.90 Infer output layouts for Split operator.
91 91 
@@ -143,7 +143,7 @@ class SplitDistributedOp(DistributedOp):
143class SplitTensorDistributedOp(DistributedOp):143class SplitTensorDistributedOp(DistributedOp):
144 """Distributed implementation for SplitTensor operator."""144 """Distributed implementation for SplitTensor operator."""
145 145 
146- def infer_layout(self, layouts, extra_args):146+ def infer_layout(self, layouts, extra_args=None):
147 """147 """
148 Infer output layouts for Split operator.148 Infer output layouts for Split operator.
149 149 
@@ -178,7 +178,7 @@ class SplitTensorDistributedOp(DistributedOp):
178class SplitTensorViewDistributedOp(DistributedOp):178class SplitTensorViewDistributedOp(DistributedOp):
179 """Distributed implementation for SplitTensorView operator."""179 """Distributed implementation for SplitTensorView operator."""
180 180 
181- def infer_layout(self, layouts, extra_args):181+ def infer_layout(self, layouts, extra_args=None):
182 """182 """
183 Infer output layouts for SplitTensorView operator.183 Infer output layouts for SplitTensorView operator.
184 184 
@@ -213,7 +213,7 @@ class SplitTensorViewDistributedOp(DistributedOp):
213class TensorSplitDistributedOp(DistributedOp):213class TensorSplitDistributedOp(DistributedOp):
214 """Distributed implementation for tensor_split operator."""214 """Distributed implementation for tensor_split operator."""
215 215 
216- def infer_layout(self, layouts, extra_args):216+ def infer_layout(self, layouts, extra_args=None):
217 """217 """
218 Infer output layouts for tensor_split operator.218 Infer output layouts for tensor_split operator.
219 219 
@@ -22,7 +22,7 @@ from .parallel_ops import DistributedOp
22class SqueezeDistributedOp(DistributedOp):22class SqueezeDistributedOp(DistributedOp):
23 """Distributed implementation for Squeeze operator."""23 """Distributed implementation for Squeeze operator."""
24 24 
25- def infer_layout(self, layouts, extra_args):25+ def infer_layout(self, layouts, extra_args=None):
26 """26 """
27 Infer output layout for Squeeze.27 Infer output layout for Squeeze.
28 28 
@@ -23,7 +23,7 @@ from .parallel_ops import DistributedOp
23class TransposeDistributedOp(DistributedOp):23class TransposeDistributedOp(DistributedOp):
24 """Distributed implementation for Transpose operator."""24 """Distributed implementation for Transpose operator."""
25 25 
26- def infer_layout(self, layouts, extra_args):26+ def infer_layout(self, layouts, extra_args=None):
27 """27 """
28 Infer output layout for Transpose operator.28 Infer output layout for Transpose operator.
29 29 
@@ -25,7 +25,7 @@ class TupleElementWiseDistributedOp(DistributedOp):
25 25 
26 Inherits from DistributedOp and provides element-wise specific implementations.26 Inherits from DistributedOp and provides element-wise specific implementations.
27 """27 """
28- def infer_layout(self, layouts, extra_args):28+ def infer_layout(self, layouts, extra_args=None):
29 """29 """
30 Infer output layouts for element-wise operations.30 Infer output layouts for element-wise operations.
31 31 
@@ -23,7 +23,7 @@ from .parallel_ops import DistributedOp
23class UnbindDistributedOp(DistributedOp):23class UnbindDistributedOp(DistributedOp):
24 """Distributed implementation for Unbind operator."""24 """Distributed implementation for Unbind operator."""
25 25 
26- def infer_layout(self, layouts, extra_args):26+ def infer_layout(self, layouts, extra_args=None):
27 """27 """
28 Infer output layouts for Unbind operator.28 Infer output layouts for Unbind operator.
29 29 
@@ -1,4 +1,4 @@
1npu_fusion_attention:1npu_fusion_attention:
2 dist_op_name: _torch_npu_fusion_attention_dist_op2 dist_op_name: _torch_npu_fusion_attention_dist_op
3- distributed_op_class: FlashAttentionScoreDistributedOp3+ distributed_op_class: NPUFlashAttentionScoreDistributedOp
4 distributed_op_file: parallel_npu_flash_attention_score4 distributed_op_file: parallel_npu_flash_attention_score
@@ -35,6 +35,7 @@ from dataclasses import dataclass, field
35from contextlib import ExitStack35from contextlib import ExitStack
36import torch36import torch
37import torch.distributed as dist37import torch.distributed as dist
38+from torch.distributed import Work
38from hyper_parallel.core.fully_shard.utils import (39from hyper_parallel.core.fully_shard.utils import (
39 MixedPrecisionPolicy,40 MixedPrecisionPolicy,
40 FSDPMeshInfo,41 FSDPMeshInfo,
@@ -138,7 +139,7 @@ class AllGatherResult(NamedTuple):
138 """139 """
139 all_gather_output: torch.Tensor140 all_gather_output: torch.Tensor
140 metadata: AllGatherMetadata141 metadata: AllGatherMetadata
141- handle: dist.distributed_c10d.Work | None142+ handle: Optional[Work]
142 143 
143 144 
144@dataclass145@dataclass
@@ -159,8 +160,8 @@ class CommContext:
159 Layer N reduce_scatter ↔ Layer N-1 backward compute160 Layer N reduce_scatter ↔ Layer N-1 backward compute
160 Layer N all_reduce ↔ Layer N-1 reduce_scatter161 Layer N all_reduce ↔ Layer N-1 reduce_scatter
161 """162 """
162- comm_handle: dist.distributed_c10d.Work | None = None163+ comm_handle: Optional[Work] = None
163- all_reduce_handle: dist.distributed_c10d.Work | None = None164+ all_reduce_handle: Optional[Work] = None
164 pre_param_group = None165 pre_param_group = None
165 # Param group whose all_reduce has been issued but grad not yet applied166 # Param group whose all_reduce has been issued but grad not yet applied
166 all_reduce_param_group = None167 all_reduce_param_group = None
@@ -181,7 +181,7 @@ class BsndFaCausalLeftupAttn(nn.Module):
181 181 
182 Requires a fixed 2048×2048 compressed causal mask.182 Requires a fixed 2048×2048 compressed causal mask.
183 Upper-right triangle = True (mask future tokens), lower-left = False (attend past).183 Upper-right triangle = True (mask future tokens), lower-left = False (attend past).
184- In Colossal CP, the dispatcher (FlashAttentionScoreDistributedOp._compute_sparse_params)184+ In Colossal CP, the dispatcher (NPUFlashAttentionScoreDistributedOp._compute_sparse_params)
185 adjusts pre_tockens/next_tockens per rank to achieve globally-correct causal attention.185 adjusts pre_tockens/next_tockens per rank to achieve globally-correct causal attention.
186 """186 """
187 187 
@@ -829,7 +829,7 @@ def test_colossal_bsnd_fa_noncausal():
829def test_colossal_bsnd_fa_causal_leftup():829def test_colossal_bsnd_fa_causal_leftup():
830 """C3b: Pure Colossal CP=2, BSND, npu_fusion_attention causal (sparse_mode=2 leftUpCausal).830 """C3b: Pure Colossal CP=2, BSND, npu_fusion_attention causal (sparse_mode=2 leftUpCausal).
831 831 
832- Core test: FlashAttentionScoreDistributedOp._compute_sparse_params converts832+ Core test: NPUFlashAttentionScoreDistributedOp._compute_sparse_params converts
833 sparse_mode=2 to BAND (mode=4) and adjusts pre/next_tockens per rank so that833 sparse_mode=2 to BAND (mode=4) and adjusts pre/next_tockens per rank so that
834 each rank's local Q only attends to the globally-correct KV window.834 each rank's local Q only attends to the globally-correct KV window.
835 835 
@@ -15,14 +15,14 @@
15"""parallel_linear test"""15"""parallel_linear test"""
16import os16import os
17import unittest17import unittest
18-from unittest.mock import patch18+from unittest.mock import MagicMock, patch
19+ 
19import numpy as np20import numpy as np
20-os.environ["HYPER_PARALLEL_PLATFORM"] = "mindspore"21+os.environ["HYPER_PARALLEL_PLATFORM"] = "torch"
21 22 
22from hyper_parallel.core.dtensor.dtensor import _build_layout23from hyper_parallel.core.dtensor.dtensor import _build_layout
23from hyper_parallel.core.dtensor.placement_types import Shard, Replicate24from hyper_parallel.core.dtensor.placement_types import Shard, Replicate
24from hyper_parallel.core.shard.ops.parallel_matmul import LinearDistributedOp25from hyper_parallel.core.shard.ops.parallel_matmul import LinearDistributedOp
25-from hyper_parallel.platform import get_platform
26from hyper_parallel.core.dtensor.device_mesh import (26from hyper_parallel.core.dtensor.device_mesh import (
27 init_device_mesh,27 init_device_mesh,
28 _DEVICE_MESH_MAP28 _DEVICE_MESH_MAP
@@ -33,99 +33,67 @@ op = LinearDistributedOp("Linear")
33 33 
34 34 
35class TestParallelLinear(unittest.TestCase):35class TestParallelLinear(unittest.TestCase):
36- """Unit tests for LinearDistributedOp."""36+ """Test Parallel Linear Distributed Operator."""
37 def setUp(self):37 def setUp(self):
38- """Set up test fixtures before each test method.
39- 
40- Clears global caches to ensure test isolation and initializes
41- the platform for testing.
42- """
43 EXISTING_COMM_GROUPS.clear()38 EXISTING_COMM_GROUPS.clear()
44 _DEVICE_MESH_MAP.clear()39 _DEVICE_MESH_MAP.clear()
45- self.platform = get_platform()
46 40 
47 def tearDown(self):41 def tearDown(self):
48- """Clean up after each test method."""
49 EXISTING_COMM_GROUPS.clear()42 EXISTING_COMM_GROUPS.clear()
50 _DEVICE_MESH_MAP.clear()43 _DEVICE_MESH_MAP.clear()
51 44 
52- def _setup_mock_platform(self, mock_platform, platform_type=None, world_size=8):
53- """Configure common mock-platform attributes used across tests.
54- 
55- Args:
56- mock_platform: The MagicMock object injected by @patch.
57- platform_type: Optional PlatformType to set on the mock.
58- world_size: Value returned by mock_platform.get_world_size().
59- """
60- if platform_type is not None:
61- mock_platform.platform_type = platform_type
62- mock_platform.get_rank.return_value = 0
63- mock_platform.get_world_size.return_value = world_size
64- mock_platform.tensor_to_numpy.side_effect = (
65- lambda t: t.numpy() if hasattr(t, "numpy") else np.array(t)
66- )
67- 
68 def _make_2x4_mesh(self, mock_platform):45 def _make_2x4_mesh(self, mock_platform):
69- """Set up mock and return a standard 2x4 (dp, mp) mesh via init_device_mesh."""46+ """Mock a 2x4 device mesh."""
70- self._setup_mock_platform(mock_platform, world_size=8)47+ mock_platform.get_rank.return_value = 0
71- return init_device_mesh(device_type="npu", mesh_shape=(2, 4), mesh_dim_names=("dp", "mp"))48+ mock_platform.get_world_size.return_value = 8
49+ mock_platform.platform_type = MagicMock()
50+ return init_device_mesh(device_type="cpu", mesh_shape=(2, 4),
51+ mesh_dim_names=("dp", "mp"), init_backend=False)
72 52 
73 @patch("hyper_parallel.core.dtensor.device_mesh.platform")53 @patch("hyper_parallel.core.dtensor.device_mesh.platform")
74 def test_linear_layout_data_parallel(self, mock_platform):54 def test_linear_layout_data_parallel(self, mock_platform):
75- """55+ """Test Linear layout with Data Parallel."""
76- Feature: Linear data parallel
77- Description: Data parallel scenario
78- Expectation: Success
79- """
80 mesh = self._make_2x4_mesh(mock_platform)56 mesh = self._make_2x4_mesh(mock_platform)
81 x_layout = _build_layout(mesh, (Shard(0), Replicate()), 2)57 x_layout = _build_layout(mesh, (Shard(0), Replicate()), 2)
82 w_layout = _build_layout(mesh, (Replicate(), Replicate()), 2)58 w_layout = _build_layout(mesh, (Replicate(), Replicate()), 2)
83- output_layout = op.infer_layout((x_layout, w_layout, None), ())59+ cache_values = [x_layout, w_layout, None]
60+ output_layouts, extra_info = op.infer_layout(cache_values)
61+ output_layout = output_layouts[0]
84 expected_map = (1, -1)62 expected_map = (1, -1)
85 assert output_layout.tensor_map == expected_map, (63 assert output_layout.tensor_map == expected_map, (
86- f"Data Parallel with transpose_a test failed. Expected {expected_map},"64+ f"Data Parallel with transpose.a test failed. Expected {expected_map},"
87 f" got {output_layout.tensor_map}"65 f" got {output_layout.tensor_map}"
88 )66 )
89 67 
90- assert op.get_expand_impl(None, output_layout, (x_layout, w_layout, None), ())is None, (68+ assert extra_info is None, f"extra_info should be None, got {extra_info}"
91- f"get_expand_impl test failed. Expected None, "
92- f"got {self.op.get_expand_impl(None, output_layout, (x_layout,), None)}"
93- )
94 69 
95 @patch("hyper_parallel.core.dtensor.device_mesh.platform")70 @patch("hyper_parallel.core.dtensor.device_mesh.platform")
96 def test_linear_layout_hybrid_parallel(self, mock_platform):71 def test_linear_layout_hybrid_parallel(self, mock_platform):
97- """72+ """Test Linear layout with Hybrid Parallel."""
98- Feature: Linear hybrid parallel
99- Description: Hybrid parallel scenario
100- Expectation: Success
101- """
102 mesh = self._make_2x4_mesh(mock_platform)73 mesh = self._make_2x4_mesh(mock_platform)
103 x_layout = _build_layout(mesh, (Shard(0), Replicate()), 2)74 x_layout = _build_layout(mesh, (Shard(0), Replicate()), 2)
104 w_layout = _build_layout(mesh, (Replicate(), Shard(0)), 2)75 w_layout = _build_layout(mesh, (Replicate(), Shard(0)), 2)
105 bias_layout = _build_layout(mesh, (Replicate(), Shard(0)), 1)76 bias_layout = _build_layout(mesh, (Replicate(), Shard(0)), 1)
106- output_layout = op.infer_layout((x_layout, w_layout, bias_layout), ())77+ cache_values = [x_layout, w_layout, bias_layout]
78+ output_layouts, extra_info = op.infer_layout(cache_values)
79+ output_layout = output_layouts[0]
107 expected_map = (1, 0)80 expected_map = (1, 0)
108 assert output_layout.tensor_map == expected_map, (81 assert output_layout.tensor_map == expected_map, (
109 f"Hybrid Parallel test failed. Expected {expected_map}, "82 f"Hybrid Parallel test failed. Expected {expected_map}, "
110 f"got {output_layout.tensor_map}"83 f"got {output_layout.tensor_map}"
111 )84 )
112- assert op.get_expand_impl(None, output_layout, (x_layout, w_layout, bias_layout), ())is None, (85+ assert extra_info is None, f"extra_info should be None, got {extra_info}"
113- f"get_expand_impl test failed. Expected None, "
114- f"got {self.op.get_expand_impl(None, output_layout, (x_layout,), None)}"
115- )
116 86 
117 @patch("hyper_parallel.core.dtensor.device_mesh.platform")87 @patch("hyper_parallel.core.dtensor.device_mesh.platform")
118 def test_linear_layout_hybrid_tensor_parallel(self, mock_platform):88 def test_linear_layout_hybrid_tensor_parallel(self, mock_platform):
119- """89+ """Test Linear layout with Hybrid Tensor Parallel."""
120- Feature: Linear hybrid tensor parallel
121- Description: Hybrid tensor parallel scenario
122- Expectation: Success
123- """
124 mesh = self._make_2x4_mesh(mock_platform)90 mesh = self._make_2x4_mesh(mock_platform)
125 x_layout = _build_layout(mesh, (Shard(0), Shard(1)), 2)91 x_layout = _build_layout(mesh, (Shard(0), Shard(1)), 2)
126 w_layout = _build_layout(mesh, (Replicate(), Shard(1)), 2)92 w_layout = _build_layout(mesh, (Replicate(), Shard(1)), 2)
127 93 
128- output_layout = op.infer_layout((x_layout, w_layout, None), ())94+ cache_values = [x_layout, w_layout, None]
95+ output_layouts, extra_info = op.infer_layout(cache_values)
96+ output_layout = output_layouts[0]
129 expected_map = (1, -1)97 expected_map = (1, -1)
130 assert output_layout.tensor_map == expected_map, (98 assert output_layout.tensor_map == expected_map, (
131 f"Hybrid Tensor Parallel test failed. Expected {expected_map}, "99 f"Hybrid Tensor Parallel test failed. Expected {expected_map}, "
@@ -134,29 +102,24 @@ class TestParallelLinear(unittest.TestCase):
134 102 
135 @patch("hyper_parallel.core.dtensor.device_mesh.platform")103 @patch("hyper_parallel.core.dtensor.device_mesh.platform")
136 def test_linear_layout_hybrid_tensor_parallel_with_bias(self, mock_platform):104 def test_linear_layout_hybrid_tensor_parallel_with_bias(self, mock_platform):
137- """105+ """Test Linear layout with Hybrid Tensor Parallel and Bias."""
138- Feature: Linear hybrid tensor parallel
139- Description: Hybrid tensor parallel scenario
140- Expectation: raise error
141- """
142 mesh = self._make_2x4_mesh(mock_platform)106 mesh = self._make_2x4_mesh(mock_platform)
143 x_layout = _build_layout(mesh, (Shard(0), Shard(1)), 2)107 x_layout = _build_layout(mesh, (Shard(0), Shard(1)), 2)
144 w_layout = _build_layout(mesh, (Replicate(), Shard(1)), 2)108 w_layout = _build_layout(mesh, (Replicate(), Shard(1)), 2)
145 bias_layout = _build_layout(mesh, (Shard(0),), 1)109 bias_layout = _build_layout(mesh, (Shard(0),), 1)
110+ cache_values = [x_layout, w_layout, bias_layout]
146 with self.assertRaisesRegex(ValueError, "Output dimensions must have same sharding"):111 with self.assertRaisesRegex(ValueError, "Output dimensions must have same sharding"):
147- _ = op.infer_layout((x_layout, w_layout, bias_layout), ())112+ _ = op.infer_layout(cache_values)
148 113 
149 @patch("hyper_parallel.core.dtensor.device_mesh.platform")114 @patch("hyper_parallel.core.dtensor.device_mesh.platform")
150 def test_linear_layout_partial_with_sharded_contract_dim(self, mock_platform):115 def test_linear_layout_partial_with_sharded_contract_dim(self, mock_platform):
151- """116+ """Test Linear layout with Partial status with sharded contract dim."""
152- Feature: Linear partial status with sharded contract dimension
153- Description: Test that partial status is set when contract dimension is sharded
154- Expectation: Success
155- """
156 mesh = self._make_2x4_mesh(mock_platform)117 mesh = self._make_2x4_mesh(mock_platform)
157 x_layout = _build_layout(mesh, (Shard(0), Shard(1)), 2)118 x_layout = _build_layout(mesh, (Shard(0), Shard(1)), 2)
158 w_layout = _build_layout(mesh, (Replicate(), Shard(1)), 2)119 w_layout = _build_layout(mesh, (Replicate(), Shard(1)), 2)
159- output_layout = op.infer_layout((x_layout, w_layout, None), ())120+ cache_values = [x_layout, w_layout, None]
121+ output_layouts, extra_info = op.infer_layout(cache_values)
122+ output_layout = output_layouts[0]
160 123 
161 expected_partial = [None, 'sum']124 expected_partial = [None, 'sum']
162 assert output_layout.partial == expected_partial, (125 assert output_layout.partial == expected_partial, (
@@ -164,32 +127,25 @@ class TestParallelLinear(unittest.TestCase):
164 f"got {output_layout.partial}"127 f"got {output_layout.partial}"
165 )128 )
166 129 
167- assert op.get_expand_impl(None, output_layout, (x_layout, w_layout, None), ())is None, (130+ assert extra_info is None, f"extra_info should be None, got {extra_info}"
168- f"get_expand_impl test failed. Expected None, "
169- f"got {self.op.get_expand_impl(None, output_layout, (x_layout,), None)}"
170- )
171 131 
172 @patch("hyper_parallel.core.dtensor.device_mesh.platform")132 @patch("hyper_parallel.core.dtensor.device_mesh.platform")
173 def test_linear_layout_partial_without_sharded_contract_dim(self, mock_platform):133 def test_linear_layout_partial_without_sharded_contract_dim(self, mock_platform):
174- """134+ """Test Linear layout with Partial status without sharded contract dim."""
175- Feature: Linear partial status without sharded contract dimension
176- Description: Test that partial status is None when contract dimension is not sharded
177- Expectation: Success
178- """
179 mesh = self._make_2x4_mesh(mock_platform)135 mesh = self._make_2x4_mesh(mock_platform)
180 x_layout = _build_layout(mesh, (Shard(0), Replicate()), 2)136 x_layout = _build_layout(mesh, (Shard(0), Replicate()), 2)
181 w_layout = _build_layout(mesh, (Replicate(), Replicate()), 2)137 w_layout = _build_layout(mesh, (Replicate(), Replicate()), 2)
182- output_layout = op.infer_layout((x_layout, w_layout, None), ())138+ cache_values = [x_layout, w_layout, None]
139+ output_layouts, extra_info = op.infer_layout(cache_values)
140+ output_layout = output_layouts[0]
183 141 
184 expected_partial = [None, None]142 expected_partial = [None, None]
185 assert output_layout.partial == expected_partial, (143 assert output_layout.partial == expected_partial, (
186 f"Partial status test failed. Expected {expected_partial}, "144 f"Partial status test failed. Expected {expected_partial}, "
187 f"got {output_layout.partial}"145 f"got {output_layout.partial}"
188 )146 )
189- assert op.get_expand_impl(None, output_layout, (x_layout, w_layout, None), ())is None, (147+ assert extra_info is None, f"extra_info should be None, got {extra_info}"
190- f"get_expand_impl test failed. Expected None, "148+ 
191- f"got {self.op.get_expand_impl(None, output_layout, (x_layout,), None)}"
192- )
193 149 
194if __name__ == "__main__":150if __name__ == "__main__":
195 unittest.main()151 unittest.main()
@@ -21,7 +21,7 @@ os.environ["HYPER_PARALLEL_PLATFORM"] = "mindspore"
21 21 
22from hyper_parallel.core.dtensor.dtensor import _build_layout22from hyper_parallel.core.dtensor.dtensor import _build_layout
23from hyper_parallel.core.dtensor.placement_types import Shard, Replicate23from hyper_parallel.core.dtensor.placement_types import Shard, Replicate
24-from hyper_parallel.core.shard.ops.parallel_npu_flash_attention_score import FlashAttentionScoreDistributedOp24+from hyper_parallel.core.shard.ops.parallel_npu_flash_attention_score import NPUFlashAttentionScoreDistributedOp
25from hyper_parallel.platform import get_platform25from hyper_parallel.platform import get_platform
26from hyper_parallel.core.dtensor.device_mesh import (26from hyper_parallel.core.dtensor.device_mesh import (
27 init_device_mesh,27 init_device_mesh,
@@ -29,11 +29,11 @@ from hyper_parallel.core.dtensor.device_mesh import (
29)29)
30from hyper_parallel.platform.platform import EXISTING_COMM_GROUPS30from hyper_parallel.platform.platform import EXISTING_COMM_GROUPS
31 31 
32-op = FlashAttentionScoreDistributedOp("npu_fusion_attention")32+op = NPUFlashAttentionScoreDistributedOp("npu_fusion_attention")
33 33 
34 34 
35class TestParallelNpuFlashAttentionScore(unittest.TestCase):35class TestParallelNpuFlashAttentionScore(unittest.TestCase):
36- """Unit tests for FlashAttentionScoreDistributedOp."""36+ """Unit tests for NPUFlashAttentionScoreDistributedOp."""
37 def setUp(self):37 def setUp(self):
38 """Set up test fixtures before each test method."""38 """Set up test fixtures before each test method."""
39 EXISTING_COMM_GROUPS.clear()39 EXISTING_COMM_GROUPS.clear()
@@ -15,14 +15,14 @@
15"""parallel_sort test"""15"""parallel_sort test"""
16import os16import os
17import unittest17import unittest
18-from unittest.mock import patch18+from unittest.mock import MagicMock, patch
19+ 
19import numpy as np20import numpy as np
20-os.environ["HYPER_PARALLEL_PLATFORM"] = "mindspore"21+os.environ["HYPER_PARALLEL_PLATFORM"] = "torch"
21 22 
22from hyper_parallel.core.dtensor.dtensor import _build_layout23from hyper_parallel.core.dtensor.dtensor import _build_layout
23from hyper_parallel.core.dtensor.placement_types import Shard, Replicate24from hyper_parallel.core.dtensor.placement_types import Shard, Replicate
24from hyper_parallel.core.shard.ops.parallel_sort import SortDistributedOp25from hyper_parallel.core.shard.ops.parallel_sort import SortDistributedOp
25-from hyper_parallel.platform import get_platform
26from hyper_parallel.core.dtensor.device_mesh import (26from hyper_parallel.core.dtensor.device_mesh import (
27 init_device_mesh,27 init_device_mesh,
28 _DEVICE_MESH_MAP28 _DEVICE_MESH_MAP
@@ -33,69 +33,53 @@ op = SortDistributedOp("sort")
33 33 
34 34 
35class TestParallelSort(unittest.TestCase):35class TestParallelSort(unittest.TestCase):
36- """Unit tests for SortDistributedOp."""36+ """Test parallel_sort ops."""
37 def setUp(self):37 def setUp(self):
38- """Set up test fixtures before each test method.
39- 
40- Clears global caches to ensure test isolation and initializes
41- the platform for testing.
42- """
43 EXISTING_COMM_GROUPS.clear()38 EXISTING_COMM_GROUPS.clear()
44 _DEVICE_MESH_MAP.clear()39 _DEVICE_MESH_MAP.clear()
45- self.platform = get_platform()
46 40 
47 def tearDown(self):41 def tearDown(self):
48- """Clean up after each test method."""
49 EXISTING_COMM_GROUPS.clear()42 EXISTING_COMM_GROUPS.clear()
50 _DEVICE_MESH_MAP.clear()43 _DEVICE_MESH_MAP.clear()
51 44 
52- def _setup_mock_platform(self, mock_platform, platform_type=None, world_size=8):
53- """Configure common mock-platform attributes used across tests.
54- 
55- Args:
56- mock_platform: The MagicMock object injected by @patch.
57- platform_type: Optional PlatformType to set on the mock.
58- world_size: Value returned by mock_platform.get_world_size().
59- """
60- if platform_type is not None:
61- mock_platform.platform_type = platform_type
62- mock_platform.get_rank.return_value = 0
63- mock_platform.get_world_size.return_value = world_size
64- mock_platform.tensor_to_numpy.side_effect = (
65- lambda t: t.numpy() if hasattr(t, "numpy") else np.array(t)
66- )
67- 
68 def _make_2x4_mesh(self, mock_platform):45 def _make_2x4_mesh(self, mock_platform):
69- """Set up mock and return a standard 2x4 (dp, mp) mesh via init_device_mesh."""46+ """Mock a 2x4 device mesh."""
70- self._setup_mock_platform(mock_platform, world_size=8)47+ mock_platform.get_rank.return_value = 0
71- return init_device_mesh(device_type="npu", mesh_shape=(2, 4), mesh_dim_names=("dp", "mp"))48+ mock_platform.get_world_size.return_value = 8
49+ mock_platform.platform_type = MagicMock()
50+ return init_device_mesh(device_type="cpu", mesh_shape=(2, 4),
51+ mesh_dim_names=("dp", "mp"), init_backend=False)
72 52 
73 def _make_2x2_mesh(self, mock_platform):53 def _make_2x2_mesh(self, mock_platform):
74- """Set up mock and return a standard 2x2 (dp, tp) mesh via init_device_mesh."""54+ """Mock a 2x2 device mesh."""
75- self._setup_mock_platform(mock_platform, world_size=4)55+ mock_platform.get_rank.return_value = 0
76- return init_device_mesh(device_type="npu", mesh_shape=(2, 2), mesh_dim_names=("dp", "tp"))56+ mock_platform.get_world_size.return_value = 4
57+ mock_platform.platform_type = MagicMock()
58+ return init_device_mesh(device_type="cpu", mesh_shape=(2, 2),
59+ mesh_dim_names=("dp", "tp"), init_backend=False)
77 60 
78 def _make_2x2x2_mesh(self, mock_platform):61 def _make_2x2x2_mesh(self, mock_platform):
79- """Set up mock and return a standard 2x2x2 (dp, tp, mp) mesh via init_device_mesh."""62+ """Mock a 2x2x2 device mesh."""
80- self._setup_mock_platform(mock_platform, world_size=8)63+ mock_platform.get_rank.return_value = 0
81- return init_device_mesh(device_type="npu", mesh_shape=(2, 2, 2), mesh_dim_names=("dp", "tp", "mp"))64+ mock_platform.get_world_size.return_value = 8
65+ mock_platform.platform_type = MagicMock()
66+ return init_device_mesh(device_type="cpu", mesh_shape=(2, 2, 2),
67+ mesh_dim_names=("dp", "tp", "mp"), init_backend=False)
82 68 
83 @patch("hyper_parallel.core.dtensor.device_mesh.platform")69 @patch("hyper_parallel.core.dtensor.device_mesh.platform")
84 def test_sort_layout_inference_basic(self, mock_platform):70 def test_sort_layout_inference_basic(self, mock_platform):
85- """71+ """Test Sort layout inference with basic sharding."""
86- Feature: Sort along an unsharded dimension
87- Description: Input is sharded on dim0, sorting is performed on dim1 (unsharded).
88- Expectation: Returns a tuple of two layouts (values, indices), both identical to input layout.
89- """
90 mesh = self._make_2x4_mesh(mock_platform)72 mesh = self._make_2x4_mesh(mock_platform)
91 x_placements = (Shard(0), Replicate())73 x_placements = (Shard(0), Replicate())
92 x_layout = _build_layout(mesh, x_placements, 2)74 x_layout = _build_layout(mesh, x_placements, 2)
93 75 
94- output_layouts = op.infer_layout((x_layout,), extra_args=(1, False, False))76+ cache_values = [x_layout, 1]
77+ output_layouts, extra_info = op.infer_layout(cache_values)
95 78 
96 assert isinstance(output_layouts, tuple) and len(output_layouts) == 2, (79 assert isinstance(output_layouts, tuple) and len(output_layouts) == 2, (
97 "Sort must return a tuple of two layouts (values, indices)"80 "Sort must return a tuple of two layouts (values, indices)"
98 )81 )
82+ assert extra_info is None, f"Sort extra_info should be None, got {extra_info}"
99 83 
100 values_layout, indices_layout = output_layouts84 values_layout, indices_layout = output_layouts
101 85 
@@ -110,82 +94,66 @@ class TestParallelSort(unittest.TestCase):
110 f"got {indices_layout.tensor_map}"94 f"got {indices_layout.tensor_map}"
111 )95 )
112 96 
113- # Since `get_expand_impl` is not overridden, it returns None by default.
114- # The same applies to other test classes, so it is unnecessary to test its return value.
115- assert op.get_expand_impl(None, output_layouts, (x_layout,), (1, False, False)) is None, (
116- f"get_expand_impl test failed. Expected None, "
117- f"got {op.get_expand_impl(None, output_layouts, (x_layout,), (1, False, False))}"
118- )
119- 
120 @patch("hyper_parallel.core.dtensor.device_mesh.platform")97 @patch("hyper_parallel.core.dtensor.device_mesh.platform")
121 def test_sort_layout_inference_sharded_dim_error(self, mock_platform):98 def test_sort_layout_inference_sharded_dim_error(self, mock_platform):
122- """99+ """Test Sort layout inference with sharded dimension error."""
123- Feature: Sort along a sharded dimension
124- Description: Input is sharded on dim0, attempt to sort on dim0.
125- Expectation: Should raise ValueError because sorting requires global data along the sort axis.
126- """
127 mesh = self._make_2x4_mesh(mock_platform)100 mesh = self._make_2x4_mesh(mock_platform)
128 x_placements = (Shard(0), Replicate())101 x_placements = (Shard(0), Replicate())
129 x_layout = _build_layout(mesh, x_placements, 2)102 x_layout = _build_layout(mesh, x_placements, 2)
130 103 
131- with self.assertRaisesRegex(ValueError, "sorting along a sharded dimension .* is not supported"):104+ cache_values = [x_layout, 0]
132- op.infer_layout((x_layout,), extra_args=(0, True, False))105+ with self.assertRaisesRegex(ValueError, "sharded dimension"):
106+ op.infer_layout(cache_values)
133 107 
134 @patch("hyper_parallel.core.dtensor.device_mesh.platform")108 @patch("hyper_parallel.core.dtensor.device_mesh.platform")
135 def test_sort_layout_inference_negative_dim(self, mock_platform):109 def test_sort_layout_inference_negative_dim(self, mock_platform):
136- """110+ """Test Sort layout inference with negative dimension."""
137- Feature: Sort with negative dimension index
138- Description: Input (2D) sharded on dim0, sort on dim=-1 (last dim, which is unsharded).
139- Expectation: Successfully infers layout, converting -1 to correct dimension index.
140- """
141 mesh = self._make_2x2_mesh(mock_platform)111 mesh = self._make_2x2_mesh(mock_platform)
142 x_placements = (Shard(0), Replicate())112 x_placements = (Shard(0), Replicate())
143 x_layout = _build_layout(mesh, x_placements, 2)113 x_layout = _build_layout(mesh, x_placements, 2)
144 114 
145- output_layouts = op.infer_layout((x_layout,), extra_args=(-1, False, True))115+ cache_values = [x_layout, -1]
116+ output_layouts, extra_info = op.infer_layout(cache_values)
146 117 
147 values_layout, indices_layout = output_layouts118 values_layout, indices_layout = output_layouts
148 expected_map = (1, -1)119 expected_map = (1, -1)
149 120 
150 assert values_layout.tensor_map == expected_map121 assert values_layout.tensor_map == expected_map
151 assert indices_layout.tensor_map == expected_map122 assert indices_layout.tensor_map == expected_map
123+ assert extra_info is None
152 124 
153 @patch("hyper_parallel.core.dtensor.device_mesh.platform")125 @patch("hyper_parallel.core.dtensor.device_mesh.platform")
154 def test_sort_layout_inference_preserve_other_dims(self, mock_platform):126 def test_sort_layout_inference_preserve_other_dims(self, mock_platform):
155- """127+ """Test Sort layout inference with preserve other dims."""
156- Feature: Sort preserves sharding on other dimensions
157- Description: 3D input sharded on dim0 and dim2. Sort on dim1 (unsharded).
158- Expectation: Output layouts preserve sharding on dim0 and dim2.
159- """
160 mesh = self._make_2x2x2_mesh(mock_platform)128 mesh = self._make_2x2x2_mesh(mock_platform)
161 x_placements = (Shard(0), Replicate(), Shard(2))129 x_placements = (Shard(0), Replicate(), Shard(2))
162 x_layout = _build_layout(mesh, x_placements, 3)130 x_layout = _build_layout(mesh, x_placements, 3)
163 131 
164- output_layouts = op.infer_layout((x_layout,), extra_args=(1, False, False))132+ cache_values = [x_layout, 1]
133+ output_layouts, extra_info = op.infer_layout(cache_values)
165 134 
166 values_layout, indices_layout = output_layouts135 values_layout, indices_layout = output_layouts
167 expected_map = (2, -1, 0)136 expected_map = (2, -1, 0)
168 137 
169 assert values_layout.tensor_map == expected_map138 assert values_layout.tensor_map == expected_map
170 assert indices_layout.tensor_map == expected_map139 assert indices_layout.tensor_map == expected_map
140+ assert extra_info is None
171 141 
172 @patch("hyper_parallel.core.dtensor.device_mesh.platform")142 @patch("hyper_parallel.core.dtensor.device_mesh.platform")
173 def test_sort_layout_inference_all_replicate(self, mock_platform):143 def test_sort_layout_inference_all_replicate(self, mock_platform):
174- """144+ """Test Sort layout inference with all Replicate."""
175- Feature: Sort on fully replicated tensor
176- Description: Input is fully replicated. Sort on any dimension.
177- Expectation: Output is fully replicated.
178- """
179 mesh = self._make_2x2_mesh(mock_platform)145 mesh = self._make_2x2_mesh(mock_platform)
180 x_placements = (Replicate(), Replicate())146 x_placements = (Replicate(), Replicate())
181 x_layout = _build_layout(mesh, x_placements, 2)147 x_layout = _build_layout(mesh, x_placements, 2)
182 148 
183- output_layouts = op.infer_layout((x_layout,), extra_args=(0, False, False))149+ cache_values = [x_layout, 0]
150+ output_layouts, extra_info = op.infer_layout(cache_values)
184 151 
185 values_layout, _ = output_layouts152 values_layout, _ = output_layouts
186 expected_map = (-1, -1)153 expected_map = (-1, -1)
187 154 
188 assert values_layout.tensor_map == expected_map155 assert values_layout.tensor_map == expected_map
156+ assert extra_info is None
189 157 
190 158 
191if __name__ == "__main__":159if __name__ == "__main__":
@@ -18,13 +18,26 @@ Unit tests for OpDispatcher with custom distributed ops (e.g., StackExt).
18import importlib18import importlib
19import os19import os
20import sys20import sys
21+import unittest
21from pathlib import Path22from pathlib import Path
22from typing import Optional, Tuple23from typing import Optional, Tuple
24+from unittest.mock import MagicMock, patch
23 25 
26+import numpy as np
24import pytest27import pytest
25 28 
26from hyper_parallel.core.dtensor.layout import Layout29from hyper_parallel.core.dtensor.layout import Layout
27from hyper_parallel.core.dtensor.dtensor import DTensor30from hyper_parallel.core.dtensor.dtensor import DTensor
31+from hyper_parallel.core.dtensor.dtensor import _build_layout
32+from hyper_parallel.core.dtensor.placement_types import Shard, Replicate
33+from hyper_parallel.core.dtensor.device_mesh import (
34+ init_device_mesh,
35+ _DEVICE_MESH_MAP
36+)
37+ 
38+from hyper_parallel.platform import get_platform
39+from hyper_parallel.platform.platform import EXISTING_COMM_GROUPS
40+from hyper_parallel.core.shard._op_dispatch import LayoutCacheKey
28 41 
29_TEST_FILE_DIR = Path(__file__).resolve().parent42_TEST_FILE_DIR = Path(__file__).resolve().parent
30_TESTS_ROOT_DIR = _TEST_FILE_DIR.parent.parent.parent.parent43_TESTS_ROOT_DIR = _TEST_FILE_DIR.parent.parent.parent.parent
@@ -33,92 +46,20 @@ _CUSTOM_OPS_DIR = _TESTS_ROOT_DIR / "tests" / "custom_ops"
33HYPER_PARALLEL_OPS_YAML_DIR = str(_CUSTOM_OPS_DIR)46HYPER_PARALLEL_OPS_YAML_DIR = str(_CUSTOM_OPS_DIR)
34HYPER_PARALLEL_OPS_PYTHON_PATH = str(_CUSTOM_OPS_DIR)47HYPER_PARALLEL_OPS_PYTHON_PATH = str(_CUSTOM_OPS_DIR)
35 48 
36-_YAML_DIR_ENV_KEYS = ("HP_TEST_OPS_YAML_DIR", "HYPER_PARALLEL_OPS_YAML_DIR")
37-_PY_PATH_ENV_KEYS = ("HP_TEST_OPS_PYTHON_PATH", "HYPER_PARALLEL_OPS_PYTHON_PATH")
38 49 
39- 50+def _reload_op_dispatch_with_env_str(yaml_dir: str, python_path: str):
40-def _first_env(*keys: str) -> Optional[str]:
41 """51 """
42- Feature: Read first available environment variable52+ Reload OpDispatcher module with custom environment variables.
43- Description: Iterate over provided environment variable keys and return the first non-empty value.53+ 
44- Expectation: Returns a string value when found; otherwise returns None.54+ Args:
55+ yaml_dir (str): Path to the directory containing op dispatch YAML files.
56+ python_path (str): Path to the directory containing custom op implementations.
57+ 
58+ Returns:
59+ module: The reloaded OpDispatcher module.
45 """60 """
46- for k in keys:61+ os.environ["HYPER_PARALLEL_OPS_YAML_DIR"] = yaml_dir
47- v = os.environ.get(k)62+ os.environ["HYPER_PARALLEL_OPS_PYTHON_PATH"] = python_path
48- if v:
49- return v
50- return None
51- 
52- 
53-def _normalize_yaml_dir(yaml_dir_or_file: str) -> Path:
54- """
55- _op_dispatch.safe_load_yaml_from_dir() expects a DIRECTORY containing *.yaml.
56- If user passes a YAML file path, accept it and use its parent directory.
57- """
58- p = Path(yaml_dir_or_file).expanduser().resolve(strict=False)
59- if p.suffix.lower() in (".yml", ".yaml"):
60- return p.parent
61- return p
62- 
63- 
64-def _validate_paths(yaml_dir: Path, python_path: str) -> None:
65- """
66- Feature: Validate custom ops YAML directory and Python path
67- Description: Ensure yaml_dir exists and is a directory, and python_path contains at least one
68- existing directory (split by ':').
69- Expectation: Raises ValueError when validation fails; otherwise returns None.
70- """
71- if not yaml_dir.exists() or not yaml_dir.is_dir():
72- raise ValueError(
73- f"Invalid yaml directory path: {yaml_dir}\n"
74- f"Expected a DIRECTORY containing *.yaml files.\n"
75- )
76- 
77- py_dirs = [
78- Path(x).expanduser().resolve(strict=False) for x in python_path.split(":") if x
79- ]
80- if not py_dirs or not any(d.exists() and d.is_dir() for d in py_dirs):
81- raise ValueError(
82- f"Invalid python path: {python_path}\n"
83- f"Expected at least one existing DIRECTORY.\n"
84- )
85- 
86- 
87-def _get_custom_paths_from_env() -> Tuple[Path, str]:
88- """
89- Feature: Resolve custom ops search paths from environment variables
90- Description: Read YAML directory and Python import search path from the first available
91- environment variables in _YAML_DIR_ENV_KEYS / _PY_PATH_ENV_KEYS, normalize
92- YAML path (file->parent dir), and validate both paths exist.
93- Expectation: Returns (yaml_dir, python_path) when env vars are present and valid;
94- otherwise raises RuntimeError/ValueError.
95- """
96- yaml_raw = _first_env(*_YAML_DIR_ENV_KEYS)
97- py_raw = _first_env(*_PY_PATH_ENV_KEYS)
98- 
99- if not yaml_raw or not py_raw:
100- raise RuntimeError(
101- "Missing env vars for custom ops paths.\n"
102- "Set either:\n"
103- " HP_TEST_OPS_YAML_DIR and HP_TEST_OPS_PYTHON_PATH\n"
104- "or:\n"
105- " HYPER_PARALLEL_OPS_YAML_DIR and HYPER_PARALLEL_OPS_PYTHON_PATH\n"
106- )
107- 
108- yaml_dir = _normalize_yaml_dir(yaml_raw)
109- python_path = py_raw
110- _validate_paths(yaml_dir, python_path)
111- return yaml_dir, python_path
112- 
113- 
114-def _reload_op_dispatch_with_env(
115- monkeypatch: pytest.MonkeyPatch, yaml_dir: str, python_path: str
116-):
117- """
118- MUST set env BEFORE importing _op_dispatch because it creates _OP_DISPATCHER at import time.
119- """
120- monkeypatch.setenv("HYPER_PARALLEL_OPS_YAML_DIR", yaml_dir)
121- monkeypatch.setenv("HYPER_PARALLEL_OPS_PYTHON_PATH", python_path)
122 63 
123 target_mod = "hyper_parallel.core.shard._op_dispatch"64 target_mod = "hyper_parallel.core.shard._op_dispatch"
124 if target_mod in sys.modules:65 if target_mod in sys.modules:
@@ -129,159 +70,251 @@ def _reload_op_dispatch_with_env(
129 return mod70 return mod
130 71 
131 72 
132-def _require_mindspore():73+class TestStackExtDispatch(unittest.TestCase):
133 """74 """
134- Feature: Conditional dependency gate for MindSpore75+ Feature: StackExt Dispatch and Layout Cache
135- Description: Check MindSpore availability at runtime; skip tests when not installed.76+ Description: Test StackExt distributed operator dispatch and layout caching.
136- Expectation: Test is skipped (pytest.skip) if MindSpore cannot be imported.77+ Expectation: dispatch should return correct DTensor output with proper layout,
78+ and layout cache should work correctly.
137 """79 """
138- try:80+ 
139- importlib.import_module("mindspore")81+ def setUp(self):
140- except ImportError as e:82+ EXISTING_COMM_GROUPS.clear()
141- pytest.skip(f"mindspore not available: {e}")83+ _DEVICE_MESH_MAP.clear()
84+ self.platform = get_platform()
85+ 
86+ def tearDown(self):
87+ EXISTING_COMM_GROUPS.clear()
88+ _DEVICE_MESH_MAP.clear()
89+ 
90+ def _make_mesh(self, mock_platform, mesh_shape, mesh_dim_names):
91+ """Create a device mesh for testing."""
92+ EXISTING_COMM_GROUPS.clear()
93+ _DEVICE_MESH_MAP.clear()
94+ mock_platform.get_rank.return_value = 0
95+ mock_platform.get_world_size.return_value = np.prod(mesh_shape)
96+ return init_device_mesh(
97+ device_type="npu",
98+ mesh_shape=mesh_shape,
99+ mesh_dim_names=mesh_dim_names,
100+ init_backend=False,
101+ )
102+ 
103+ @patch("hyper_parallel.core.dtensor.device_mesh.platform")
104+ def test_stack_ext_dispatch_and_layout(self, mock_platform):
105+ """Test StackExt dispatch and layout cache with two input DTensors."""
106+ op_dispatch = _reload_op_dispatch_with_env_str(
107+ HYPER_PARALLEL_OPS_YAML_DIR, HYPER_PARALLEL_OPS_PYTHON_PATH
108+ )
109+ 
110+ mesh = self._make_mesh(mock_platform, (1, 1, 1), ("dp", "cp", "mp"))
111+ base_layout = _build_layout(mesh, (Replicate(), Replicate(), Replicate()), 2)
112+ 
113+ from hyper_parallel.core.shard._op_dispatch import LayoutCacheManager
114+ 
115+ dist_op = LayoutCacheManager.get_instance().distributed_op("StackExt")
116+ 
117+ np_obj = np
118+ local_tensor0 = np_obj.arange(6).reshape(2, 3).astype(np_obj.int32)
119+ local_tensor1 = np_obj.arange(6, 12).reshape(2, 3).astype(np_obj.int32)
120+ 
121+ d0 = MagicMock(spec=DTensor)
122+ d0._local_tensor = local_tensor0
123+ d0.layout = base_layout
124+ d0._layout = base_layout
125+ d0.to_local.return_value = local_tensor0
126+ 
127+ d1 = MagicMock(spec=DTensor)
128+ d1._local_tensor = local_tensor1
129+ d1.layout = base_layout
130+ d1._layout = base_layout
131+ d1.to_local.return_value = local_tensor1
132+ 
133+ output_layout = dist_op.infer_layout((d0.layout, d1.layout), (0,))
134+ 
135+ assert output_layout is not None
136+ assert tuple(output_layout.to_dict()["tensor_map"]) == (-1, -1, -1)
137+ 
138+ @patch("hyper_parallel.core.dtensor.device_mesh.platform")
139+ def test_stack_ext_layout_cache(self, mock_platform):
140+ """Test StackExt layout cache with multiple input layouts."""
141+ op_dispatch = _reload_op_dispatch_with_env_str(
142+ HYPER_PARALLEL_OPS_YAML_DIR, HYPER_PARALLEL_OPS_PYTHON_PATH
143+ )
144+ 
145+ mesh = self._make_mesh(mock_platform, (1, 1, 1), ("dp", "cp", "mp"))
146+ base_layout = _build_layout(mesh, (Replicate(), Replicate(), Replicate()), 2)
147+ 
148+ from hyper_parallel.core.shard._op_dispatch import LayoutCacheManager
149+ 
150+ dist_op = LayoutCacheManager.get_instance().distributed_op("StackExt")
151+ 
152+ np_obj = np
153+ local_tensor0 = np_obj.arange(6).reshape(2, 3).astype(np_obj.int32)
154+ local_tensor1 = np_obj.arange(6, 12).reshape(2, 3).astype(np_obj.int32)
155+ 
156+ d0 = MagicMock(spec=DTensor)
157+ d0._local_tensor = local_tensor0
158+ d0.layout = base_layout
159+ d0._layout = base_layout
160+ 
161+ d1 = MagicMock(spec=DTensor)
162+ d1._local_tensor = local_tensor1
163+ d1.layout = base_layout
164+ d1._layout = base_layout
165+ 
166+ output_layout = dist_op.infer_layout((d0.layout, d1.layout), (0,))
167+ 
168+ assert output_layout is not None
169+ assert tuple(output_layout.to_dict()["tensor_map"]) == (-1, -1, -1)
142 170 
143 171 
144-base_mesh_shape = (1, 1, 1)172+class TestNewDispatchFlow(unittest.TestCase):
145-base_alias_name = ("dp", "cp", "mp")
146-base_rank_list = [0]
147- 
148- 
149-def _make_layout_2d_replicated():
150 """173 """
151- 2D tensor layout, fully replicated: ("None", "None")174+ Feature: New Dispatch Flow with Preprocess and Infer Layout
152- Uses the same Layout(...) + layout(...) pattern as your elementwise UT.175+ Description: Test the new dispatch flow with preprocess and infer_layout methods for a distributed operator.
176+ Expectation: preprocess should return valid local_args, local_kwargs, and cache_values, and infer_layout should
177+ return the correct output layouts.
153 """178 """
154- layout = Layout(base_mesh_shape, base_alias_name, base_rank_list)
155- return layout("None", "None")
156 179 
180+ def setUp(self):
181+ EXISTING_COMM_GROUPS.clear()
182+ _DEVICE_MESH_MAP.clear()
183+ self.platform = get_platform()
157 184 
158-def _make_dtensors_ms():185+ def tearDown(self):
159- """186+ EXISTING_COMM_GROUPS.clear()
160- Feature: Create MindSpore DTensor inputs for StackExt dispatch tests187+ _DEVICE_MESH_MAP.clear()
161- Description: Create two 2x3 MindSpore tensors, wrap them into DTensor with a fully
162- replicated 2D layout.
163- Expectation: Returns (d0, d1) where both are DTensor and share the same replicated layout.
164- """
165- _require_mindspore()
166 188 
167- np = importlib.import_module("numpy")189+ def _make_mesh(self, mock_platform, mesh_shape, mesh_dim_names):
168- ms = importlib.import_module("mindspore")190+ """Create a device mesh for testing."""
191+ EXISTING_COMM_GROUPS.clear()
192+ _DEVICE_MESH_MAP.clear()
193+ mock_platform.get_rank.return_value = 0
194+ mock_platform.get_world_size.return_value = np.prod(mesh_shape)
195+ return init_device_mesh(
196+ device_type="npu",
197+ mesh_shape=mesh_shape,
198+ mesh_dim_names=mesh_dim_names,
199+ init_backend=False,
200+ )
169 201 
170- x0 = ms.Tensor(np.arange(6).reshape(2, 3), ms.int32)202+ @patch("hyper_parallel.core.dtensor.device_mesh.platform")
171- x1 = ms.Tensor(np.arange(6, 12).reshape(2, 3), ms.int32)203+ def test_from_cache_values_with_layout(self, mock_platform):
204+ """Test that LayoutCacheKey from_cache_values with layout returns correct key."""
205+ mesh = self._make_mesh(mock_platform, (2, 2), ("dp", "mp"))
206+ layout = _build_layout(mesh, (Replicate(), Shard(1)), 2)
207+ cache_values = [layout, 1, True]
208+ key = LayoutCacheKey.from_cache_values(cache_values)
209+ expected = [str(layout.compact_str), "1", "True"]
210+ assert list(key._tuple) == expected, (
211+ f"Expected {expected}, got {list(key._tuple)}"
212+ )
172 213 
173- x_layout = _make_layout_2d_replicated()214+ @patch("hyper_parallel.core.dtensor.device_mesh.platform")
174- d0 = DTensor.from_local(x0, x_layout.mesh, x_layout.placements)215+ def test_from_cache_values_consistency(self, mock_platform):
175- d1 = DTensor.from_local(x1, x_layout.mesh, x_layout.placements)216+ """Test that LayoutCacheKey from_cache_values is consistent with the same cache_values."""
176- return d0, d1217+ mesh = self._make_mesh(mock_platform, (2, 2), ("dp", "mp"))
218+ layout = _build_layout(mesh, (Replicate(), Shard(1)), 2)
219+ cache_values1 = [layout, 1, True]
220+ cache_values2 = [layout, 1, True]
221+ key1 = LayoutCacheKey.from_cache_values(cache_values1)
222+ key2 = LayoutCacheKey.from_cache_values(cache_values2)
223+ assert key1 == key2, f"Keys should be equal: {key1} != {key2}"
224+ assert hash(key1) == hash(key2), "Hashes should be equal"
177 225 
226+ @patch("hyper_parallel.core.dtensor.device_mesh.platform")
227+ def test_from_cache_values_different_values(self, mock_platform):
228+ """Test that LayoutCacheKey from_cache_values differs with different values."""
229+ mesh = self._make_mesh(mock_platform, (2, 2), ("dp", "mp"))
230+ layout = _build_layout(mesh, (Replicate(), Shard(1)), 2)
231+ key1 = LayoutCacheKey.from_cache_values([layout, 1, True])
232+ key2 = LayoutCacheKey.from_cache_values([layout, 0, True])
233+ assert key1 != key2, "Keys should differ"
178 234 
179-def _stack_ext_local_ms(x0, x1, axis: int):235+ @patch("hyper_parallel.core.dtensor.device_mesh.platform")
180- """236+ def test_equality_with_legacy_key(self, mock_platform):
181- Local op implementation simulating StackExt(x0, x1, axis).237+ """Test that LayoutCacheKey is equal to legacy key."""
182- """238+ mesh = self._make_mesh(mock_platform, (2, 2), ("dp", "mp"))
183- _require_mindspore()239+ layout = _build_layout(mesh, (Replicate(), Shard(1)), 2)
184- ops = importlib.import_module("mindspore.ops")240+ key_new = LayoutCacheKey.from_cache_values([layout, 1, True])
241+ key_legacy = LayoutCacheKey([str(layout.compact_str), "1", "True"])
242+ assert key_new == key_legacy, "Keys should be equal"
185 243 
186- if hasattr(ops, "stack"):244+ @patch("hyper_parallel.core.dtensor.device_mesh.platform")
187- return ops.stack([x0, x1], axis)245+ def test_dispatch_new_flow_with_preprocess(self, mock_platform):
188- return ops.Stack(axis)([x0, x1])246+ """Test that dispatch preprocess returns valid local_args, local_kwargs, and cache_values."""
247+ from hyper_parallel.core.shard.ops.parallel_sort import SortDistributedOp
189 248 
249+ op = SortDistributedOp("sort")
250+ mesh = self._make_mesh(mock_platform, (2,), ("dp",))
251+ layout = _build_layout(mesh, (Replicate(),), 2)
190 252 
191-def _to_numpy_ms(x):253+ mock_tensor = MagicMock()
192- """254+ mock_tensor._layout = layout
193- Feature: Convert MindSpore tensor-like to numpy-like255+ mock_tensor.layout = layout
194- Description: Use asnumpy() when available; otherwise return the input.256+ mock_tensor.to_local.return_value = np.random.randn(4, 4)
195- Expectation: Returns a numpy array for MindSpore tensors, or the original object.
196- """
197- return x.asnumpy() if hasattr(x, "asnumpy") else x
198 257 
258+ result = op.preprocess((mock_tensor, -1), {})
259+ assert result is not None, "preprocess should return tuple for DTensor input"
260+ local_args, local_kwargs, cache_values = result
199 261 
200-def _patch_op_name(monkeypatch, op_dispatch_mod):262+ assert local_kwargs.get("dim") == -1, "Expected dim=-1 in kwargs"
201- """263+ assert len(cache_values) == 2, "Expected 2 cache values"
202- Ensure platform.get_op_name(_stack_ext_local_ms) == "StackExt".264+ assert cache_values[0] is layout, "Expected layout in cache_values"
203- """
204- orig_get_op_name = op_dispatch_mod.platform.get_op_name
205 265 
206- def patched_get_op_name(func):266+ @patch("hyper_parallel.core.dtensor.device_mesh.platform")
207- if func is _stack_ext_local_ms:267+ def test_dispatch_new_flow_infer_layout_with_cache_values(self, mock_platform):
208- return "StackExt"268+ """Test that dispatch infer_layout with cache_values returns correct output layouts."""
209- return orig_get_op_name(func)269+ from hyper_parallel.core.shard.ops.parallel_sort import SortDistributedOp
210 270 
211- monkeypatch.setattr(271+ op = SortDistributedOp("sort")
212- op_dispatch_mod.platform, "get_op_name", patched_get_op_name, raising=True272+ mesh = self._make_mesh(mock_platform, (2,), ("dp",))
213- )273+ layout = _build_layout(mesh, (Replicate(),), 2)
274+ cache_values = [layout, -1]
214 275 
276+ infer_result = op.infer_layout(cache_values)
277+ assert isinstance(infer_result, tuple), "Expected tuple"
278+ output_layouts, extra_info = infer_result
279+ assert isinstance(output_layouts, tuple), "Expected tuple of output layouts"
280+ assert len(output_layouts) == 2, "Expected 2 output layouts"
281+ assert extra_info is None, "Expected extra_info=None"
215 282 
216-def test_stack_ext_dispatch_and_layout(monkeypatch):283+ @patch("hyper_parallel.core.dtensor.device_mesh.platform")
217- """284+ def test_dispatch_new_flow_normalizes_args(self, mock_platform):
218- Feature: OpDispatcher dispatch and StackExt layout inference integration285+ """Test that dispatch normalizes args and kwargs."""
219- Description: Reload _op_dispatch with custom YAML/Python paths from environment, create two replicated286+ from hyper_parallel.core.shard.ops.parallel_sort import SortDistributedOp
220- 2D DTensor inputs, patch platform.get_op_name so the local function maps to "StackExt",
221- and dispatch via _OP_DISPATCHER.dispatch using axis=0.
222- Expectation: dispatch returns a DTensor whose output layout tensor_map inserts a replicated dimension
223- at axis=0 (i.e., from (-1, -1) to (-1, -1, -1)), and the numerical result matches the
224- local MindSpore stack reference output.
225- """
226- op_dispatch = _reload_op_dispatch_with_env(
227- monkeypatch, HYPER_PARALLEL_OPS_YAML_DIR, HYPER_PARALLEL_OPS_PYTHON_PATH
228- )
229 287 
230- d0, d1 = _make_dtensors_ms()288+ op = SortDistributedOp("sort")
231- axis = 0289+ mesh = self._make_mesh(mock_platform, (2,), ("dp",))
290+ layout = _build_layout(mesh, (Replicate(),), 2)
232 291 
233- _patch_op_name(monkeypatch, op_dispatch)292+ mock_tensor = MagicMock()
293+ mock_tensor._layout = layout
294+ mock_tensor.layout = layout
295+ mock_tensor.to_local.return_value = np.random.randn(4, 4)
234 296 
235- out = op_dispatch._OP_DISPATCHER.dispatch( # pylint: disable=protected-access297+ result1 = op.preprocess((mock_tensor, 1, True, False), {})
236- _stack_ext_local_ms, (d0, d1, axis), {}298+ result2 = op.preprocess((mock_tensor,), {"dim": 1, "descending": True, "stable": False})
237- )
238 299 
239- assert isinstance(out, DTensor)300+ assert result1 is not None, "preprocess should return tuple"
301+ assert result2 is not None, "preprocess should return tuple"
240 302 
241- # The input 2D replication tensor map should be (-1, -1)303+ _, kwargs1, cv1 = result1
242- assert tuple(d0.layout.to_dict()["tensor_map"]) == (-1, -1)304+ _, kwargs2, cv2 = result2
305+ assert kwargs1 == kwargs2, "Normalized kwargs should match"
306+ assert cv1[1] == cv2[1], "Normalized cache_values dim should match"
243 307 
244- # axis=0: outputs rank=3, with a replicated inserted at the axis position => (-1, -1, -1)308+ def test_dispatch_falls_back_to_legacy_when_preprocess_returns_none(self):
245- assert tuple(out.layout.to_dict()["tensor_map"]) == (-1, -1, -1)309+ """Test that dispatch falls back to legacy preprocess when new preprocess returns None."""
310+ from hyper_parallel.core.shard.ops.parallel_ops import DistributedOp
246 311 
247- # The number is correct312+ class DummyOp(DistributedOp):
248- ref = _stack_ext_local_ms(d0.to_local(), d1.to_local(), axis)313+ def infer_layout(self, layouts, extra_args=None):
249- assert (_to_numpy_ms(out.to_local()) == _to_numpy_ms(ref)).all()314+ return layouts[0]
250 315 
316+ op = DummyOp("dummy")
317+ assert op.preprocess((1, 2), {"a": 3}) is None, "Default preprocess should return None for non-DTensor inputs"
251 318 
252-def test_stack_ext_layout_cache(monkeypatch):319+if __name__ == "__main__":
253- """320+ unittest.main()
254- Feature: LayoutCacheManager effectiveness for StackExt infer_layout
255- Description: Reload _op_dispatch with custom YAML/Python paths, dispatch the same StackExt call twice
256- with identical inputs and axis=1, and wrap the distributed op infer_layout to count calls.
257- Expectation: infer_layout is invoked exactly once due to layout cache hit on the second dispatch
258- (call_count["n"] == 1).
259- """
260- op_dispatch = _reload_op_dispatch_with_env(
261- monkeypatch, HYPER_PARALLEL_OPS_YAML_DIR, HYPER_PARALLEL_OPS_PYTHON_PATH
262- )
263- 
264- d0, d1 = _make_dtensors_ms()
265- axis = 1
266- 
267- _patch_op_name(monkeypatch, op_dispatch)
268- 
269- dist_op = op_dispatch.LayoutCacheManager.get_instance().distributed_op("StackExt")
270- 
271- call_count = {"n": 0}
272- orig_infer = dist_op.infer_layout
273- 
274- def wrapped_infer(layouts, extra_args):
275- call_count["n"] += 1
276- return orig_infer(layouts, extra_args)
277- 
278- monkeypatch.setattr(dist_op, "infer_layout", wrapped_infer, raising=True)
279- 
280- _ = op_dispatch._OP_DISPATCHER.dispatch( # pylint: disable=protected-access
281- _stack_ext_local_ms, (d0, d1, axis), {}
282- )
283- _ = op_dispatch._OP_DISPATCHER.dispatch( # pylint: disable=protected-access
284- _stack_ext_local_ms, (d0, d1, axis), {}
285- )
286- 
287- assert call_count["n"] == 1