已合并
refactor: OpDispatcher 重构 #508
hedongdong创建于 4月8日
refactor: OpDispatcher 重构 #508
已合并
共 37 个文件变更+637-614
| @@ -131,23 +131,33 @@ def get_dtensor_dispatch(): | |||
| 131 | 131 | ||
| 132 | 132 | ||
| 133 | class LayoutCacheKey: | 133 | class LayoutCacheKey: |
| 134 | - """ | 134 | + """Immutable layout cache key.""" |
| 135 | - Layout cache key | 135 | + __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_ids | 138 | + self._tuple = tuple(layout_ids) |
| 139 | + self._hash = hash(self._tuple) | ||
| 140 | + | ||
| 141 | + | ||
| 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 False | 153 | return False |
| 143 | - return self.layout_ids == other.layout_ids | 154 | + return self._tuple == other._tuple |
| 144 | 155 | ||
| 145 | def __hash__(self): | 156 | def __hash__(self): |
| 146 | - seed = 0 | 157 | + 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 | ||
| 152 | class LayoutCacheManager: | 162 | class LayoutCacheManager: |
| 153 | """ | 163 | """ |
| @@ -339,17 +349,15 @@ class OpDispatcher: | |||
| 339 | 349 | ||
| 340 | 350 | ||
| 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.layout | 376 | layout = arg.layout |
| 369 | layout_id = layout.compact_str | 377 | 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.layout | 397 | layout = val.layout |
| 390 | layout_id = layout.compact_str | 398 | 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_kwargs | 404 | + 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 = None | 409 | 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 and | 410 | 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_key | 417 | + 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 | - return | 464 | + 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.layout | 472 | + 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 tensors | 511 | # 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 args | 514 | # 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 args | 517 | + 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.layout | 549 | 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.shape | 554 | + 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 | 584 | ||
| 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.layout | 618 | layout = arg.layout |
| 615 | layout_id = layout.compact_str | 619 | 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.shape | 630 | 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.layout | 646 | layout = val.layout |
| 643 | layout_id = layout.compact_str | 647 | 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.shape | 656 | 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_kwargs | 660 | + 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 layout | 726 | # input layout |
| 723 | - cache_key = LayoutCacheKey([]) | ||
| 724 | input_layouts = [] | 727 | input_layouts = [] |
| 725 | 728 | ||
| 726 | layout = input_tensor.layout | 729 | layout = input_tensor.layout |
| 727 | global_shape = input_tensor.shape | 730 | global_shape = input_tensor.shape |
| 728 | input_layouts.append(layout) | 731 | input_layouts.append(layout) |
| 729 | layout_id = layout.compact_str | 732 | 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 | |||
| 22 | class ArgMaxWithValueDistributedOp(DistributedOp): | 22 | class 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 | |||
| 22 | class ConcatDistributedOp(DistributedOp): | 22 | class 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_layout | 133 | 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 partial | 621 | 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 operator | 626 | 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 = 1 | 634 | 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_int | 122 | 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 | |||
| 22 | class ExpandDimsDistributedOp(DistributedOp): | 22 | class 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 | |||
| 24 | class IndexSelectDistributedOp(DistributedOp): | 24 | class 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_layout | 100 | 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 tensor | 178 | - 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_layout | 267 | 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): | |||
| 337 | class GatherNdDistributedOp(DistributedOp): | 336 | class 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 | ||
| 23 | class MatMulExtDistributedOp(DistributedOp): | 23 | class 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 | ||
| 83 | class MatMulDistributedOp(DistributedOp): | 83 | class 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): | |||
| 273 | class BatchMatMulDistributedOp(BaseBatchMatMulDistributedOp): | 273 | class 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 | + | ||
| 344 | class LinearDistributedOp(DistributedOp): | 348 | class 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 x | 383 | x_layout (Layout): Layout of input x |
| 359 | w_layout (Layout): Layout of input w | 384 | 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 tensor | 388 | 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_shape | 397 | 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_layout | 432 | + 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 operator | 436 | 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_map | 442 | x_map = x_layout.alias_tensor_map |
| 415 | x_contract_dim = len(x_map) - 1 | 443 | x_contract_dim = len(x_map) - 1 |
| 416 | scaling_factor = 1 | 444 | scaling_factor = 1 |
| @@ -19,13 +19,13 @@ | |||
| 19 | import copy | 19 | import copy |
| 20 | import warnings | 20 | import warnings |
| 21 | 21 | ||
| 22 | -from typing import List, Tuple, Optional, Any | 22 | +from typing import List, Tuple, Optional |
| 23 | from hyper_parallel.core.shard.ops.parallel_npu_flash_attention_score import ( # pylint: disable=C0415 | 23 | from hyper_parallel.core.shard.ops.parallel_npu_flash_attention_score import ( # pylint: disable=C0415 |
| 24 | _get_lb_override, | 24 | _get_lb_override, |
| 25 | ) | 25 | ) |
| 26 | from hyper_parallel.core.dtensor.layout import Layout | 26 | from hyper_parallel.core.dtensor.layout import Layout |
| 27 | from hyper_parallel.core.dtensor.placement_types import Shard, Replicate | 27 | from hyper_parallel.core.dtensor.placement_types import Shard, Replicate |
| 28 | -from hyper_parallel.core.shard.ops.parallel_ops_register import register_distributed_op | 28 | +from hyper_parallel.core.shard.ops.parallel_ops import DistributedOp |
| 29 | from hyper_parallel.platform import get_platform | 29 | from hyper_parallel.platform import get_platform |
| 30 | 30 | ||
| 31 | platform = get_platform() | 31 | platform = 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_name | 92 | + 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 the | 800 | 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 None | 805 | 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 and | 811 | if (key_layout_check is not None and value_layout_check is not None and |
| 820 | hasattr(key_layout_check, 'tensor_map') and | 812 | 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_layout | 839 | query_layout, key_layout, input_layout |
| 848 | ) | 840 | ) |
| @@ -23,7 +23,7 @@ from .parallel_ops import DistributedOp | |||
| 23 | class MultinomialDistributedOp(DistributedOp): | 23 | class 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 | |||
| 24 | class NewOnesDistributedOp(DistributedOp): | 24 | class 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 | |||
| 20 | import threading | 20 | import threading |
| 21 | import warnings | 21 | import warnings |
| 22 | 22 | ||
| 23 | -from typing import List, Tuple, Optional, Any | 23 | +from typing import List, Tuple, Optional |
| 24 | from hyper_parallel.core.dtensor.layout import Layout | 24 | from hyper_parallel.core.dtensor.layout import Layout |
| 25 | from hyper_parallel.core.dtensor.placement_types import Shard, Replicate | 25 | from hyper_parallel.core.dtensor.placement_types import Shard, Replicate |
| 26 | -from hyper_parallel.core.shard.ops.parallel_ops_register import register_distributed_op | 26 | +from hyper_parallel.core.shard.ops.parallel_ops import DistributedOp |
| 27 | from hyper_parallel.platform import get_platform | 27 | from hyper_parallel.platform import get_platform |
| 28 | 28 | ||
| 29 | platform = get_platform() | 29 | platform = 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_name | 77 | + 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_map | 469 | 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_map | 506 | 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_layout | 544 | 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_layout | 564 | 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 | return | 643 | 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 | return | 665 | 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 | return | 689 | 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 None | 886 | 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 and | 892 | 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=False | 920 | 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 None | 925 | 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, sync | 944 | 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, sync | 971 | 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_impl | 976 | 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 1 | 1133 | 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 1 | 1138 | return 1 |
| @@ -1151,7 +1143,7 @@ class FlashAttentionScoreDistributedOp: | |||
| 1151 | if isinstance(dim_map, tuple): | 1143 | if isinstance(dim_map, tuple): |
| 1152 | total = 1 | 1144 | 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 total | 1149 | 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 0 | 1162 | 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 0 | 1167 | 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 0 | 1181 | return 0 |
| @@ -28,7 +28,7 @@ platform = get_platform() | |||
| 28 | class OneHotExtDistributedOp(DistributedOp): | 28 | class 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_layout | 73 | 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 ms | 77 | import mindspore as ms |
| 78 | from mindspore import ops, Tensor | 78 | from mindspore import ops, Tensor |
| 79 | 79 | ||
| 80 | - del output_layout | 80 | + del infer_result |
| 81 | 81 | ||
| 82 | indices_layout = layouts[0] if layouts else None | 82 | 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=W0613 | 58 | # 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 None | 96 | return None |
| 80 | 97 | ||
| 81 | - @staticmethod | 98 | + # 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 operator | 101 | Get expand implementation for the operator |
| 85 | """ | 102 | """ |
| @@ -22,7 +22,7 @@ from .parallel_ops import DistributedOp | |||
| 22 | class PadDistributedOp(DistributedOp): | 22 | class 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_type | 44 | 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 | |||
| 23 | class RepeatInterleaveDistributedOp(DistributedOp): | 23 | class 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 @@ | |||
| 19 | import copy | 19 | import copy |
| 20 | import warnings | 20 | import warnings |
| 21 | 21 | ||
| 22 | -from typing import List, Tuple, Optional, Any | 22 | +from typing import List, Tuple, Optional |
| 23 | from hyper_parallel.core.shard.ops.parallel_npu_flash_attention_score import ( # pylint: disable=C0415 | 23 | from hyper_parallel.core.shard.ops.parallel_npu_flash_attention_score import ( # pylint: disable=C0415 |
| 24 | _get_lb_override, | 24 | _get_lb_override, |
| 25 | ) | 25 | ) |
| 26 | from hyper_parallel.core.dtensor.layout import Layout | 26 | from hyper_parallel.core.dtensor.layout import Layout |
| 27 | from hyper_parallel.core.dtensor.placement_types import Shard, Replicate | 27 | from hyper_parallel.core.dtensor.placement_types import Shard, Replicate |
| 28 | -from hyper_parallel.core.shard.ops.parallel_ops_register import register_distributed_op | 28 | +from hyper_parallel.core.shard.ops.parallel_ops import DistributedOp |
| 29 | from hyper_parallel.platform import get_platform | 29 | from hyper_parallel.platform import get_platform |
| 30 | 30 | ||
| 31 | platform = get_platform() | 31 | platform = get_platform() |
| 32 | Tensor = platform.Tensor | 32 | Tensor = 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+MP | 45 | - 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 | 48 | ||
| 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, value | 277 | 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_layout | 294 | 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 None | 300 | 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 and | 306 | 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 None | 327 | + 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 | |||
| 22 | class ScatterUpdateDistributedOp(DistributedOp): | 22 | class 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_dim | 55 | 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 | |||
| 22 | class SliceExtDistributedOp(DistributedOp): | 22 | class 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. | |||
| 19 | from .parallel_ops import DistributedOp | 19 | from .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 | + | ||
| 22 | class SortDistributedOp(DistributedOp): | 26 | class 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 device | 40 | + def infer_layout(self, cache_values): |
| 30 | - (i.e., not sharded). If the dimension is sharded, a global sort cannot be performed | 41 | + 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 layout | 76 | + return ((layout, layout), None) |
| 82 | - return (layout, layout) | ||
| @@ -23,7 +23,7 @@ from .parallel_ops import DistributedOp | |||
| 23 | class SplitWithSizeDistributedOp(DistributedOp): | 23 | class 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): | |||
| 54 | class SplitWithSizeViewDistributedOp(DistributedOp): | 54 | class 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): | |||
| 85 | class SplitDistributedOp(DistributedOp): | 85 | class 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): | |||
| 143 | class SplitTensorDistributedOp(DistributedOp): | 143 | class 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): | |||
| 178 | class SplitTensorViewDistributedOp(DistributedOp): | 178 | class 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): | |||
| 213 | class TensorSplitDistributedOp(DistributedOp): | 213 | class 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 | |||
| 22 | class SqueezeDistributedOp(DistributedOp): | 22 | class 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 | |||
| 23 | class TransposeDistributedOp(DistributedOp): | 23 | class 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 | |||
| 23 | class UnbindDistributedOp(DistributedOp): | 23 | class 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 @@ | |||
| 1 | npu_fusion_attention: | 1 | npu_fusion_attention: |
| 2 | dist_op_name: _torch_npu_fusion_attention_dist_op | 2 | dist_op_name: _torch_npu_fusion_attention_dist_op |
| 3 | - distributed_op_class: FlashAttentionScoreDistributedOp | 3 | + distributed_op_class: NPUFlashAttentionScoreDistributedOp |
| 4 | distributed_op_file: parallel_npu_flash_attention_score | 4 | distributed_op_file: parallel_npu_flash_attention_score |
| @@ -35,6 +35,7 @@ from dataclasses import dataclass, field | |||
| 35 | from contextlib import ExitStack | 35 | from contextlib import ExitStack |
| 36 | import torch | 36 | import torch |
| 37 | import torch.distributed as dist | 37 | import torch.distributed as dist |
| 38 | +from torch.distributed import Work | ||
| 38 | from hyper_parallel.core.fully_shard.utils import ( | 39 | from 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.Tensor | 140 | all_gather_output: torch.Tensor |
| 140 | metadata: AllGatherMetadata | 141 | metadata: AllGatherMetadata |
| 141 | - handle: dist.distributed_c10d.Work | None | 142 | + handle: Optional[Work] |
| 142 | 143 | ||
| 143 | 144 | ||
| 144 | 145 | ||
| @@ -159,8 +160,8 @@ class CommContext: | |||
| 159 | Layer N reduce_scatter ↔ Layer N-1 backward compute | 160 | Layer N reduce_scatter ↔ Layer N-1 backward compute |
| 160 | Layer N all_reduce ↔ Layer N-1 reduce_scatter | 161 | Layer N all_reduce ↔ Layer N-1 reduce_scatter |
| 161 | """ | 162 | """ |
| 162 | - comm_handle: dist.distributed_c10d.Work | None = None | 163 | + comm_handle: Optional[Work] = None |
| 163 | - all_reduce_handle: dist.distributed_c10d.Work | None = None | 164 | + all_reduce_handle: Optional[Work] = None |
| 164 | pre_param_group = None | 165 | pre_param_group = None |
| 165 | # Param group whose all_reduce has been issued but grad not yet applied | 166 | # Param group whose all_reduce has been issued but grad not yet applied |
| 166 | all_reduce_param_group = None | 167 | 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(): | |||
| 829 | def test_colossal_bsnd_fa_causal_leftup(): | 829 | def 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 converts | 832 | + Core test: NPUFlashAttentionScoreDistributedOp._compute_sparse_params converts |
| 833 | sparse_mode=2 to BAND (mode=4) and adjusts pre/next_tockens per rank so that | 833 | 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""" |
| 16 | import os | 16 | import os |
| 17 | import unittest | 17 | import unittest |
| 18 | -from unittest.mock import patch | 18 | +from unittest.mock import MagicMock, patch |
| 19 | + | ||
| 19 | import numpy as np | 20 | import numpy as np |
| 20 | -os.environ["HYPER_PARALLEL_PLATFORM"] = "mindspore" | 21 | +os.environ["HYPER_PARALLEL_PLATFORM"] = "torch" |
| 21 | 22 | ||
| 22 | from hyper_parallel.core.dtensor.dtensor import _build_layout | 23 | from hyper_parallel.core.dtensor.dtensor import _build_layout |
| 23 | from hyper_parallel.core.dtensor.placement_types import Shard, Replicate | 24 | from hyper_parallel.core.dtensor.placement_types import Shard, Replicate |
| 24 | from hyper_parallel.core.shard.ops.parallel_matmul import LinearDistributedOp | 25 | from hyper_parallel.core.shard.ops.parallel_matmul import LinearDistributedOp |
| 25 | -from hyper_parallel.platform import get_platform | ||
| 26 | from hyper_parallel.core.dtensor.device_mesh import ( | 26 | from hyper_parallel.core.dtensor.device_mesh import ( |
| 27 | init_device_mesh, | 27 | init_device_mesh, |
| 28 | _DEVICE_MESH_MAP | 28 | _DEVICE_MESH_MAP |
| @@ -33,99 +33,67 @@ op = LinearDistributedOp("Linear") | |||
| 33 | 33 | ||
| 34 | 34 | ||
| 35 | class TestParallelLinear(unittest.TestCase): | 35 | class 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 | 53 | ||
| 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 | 70 | ||
| 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 | 87 | ||
| 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 | 103 | ||
| 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 | 114 | ||
| 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 | 132 | ||
| 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 | ||
| 194 | if __name__ == "__main__": | 150 | if __name__ == "__main__": |
| 195 | unittest.main() | 151 | unittest.main() |
| @@ -21,7 +21,7 @@ os.environ["HYPER_PARALLEL_PLATFORM"] = "mindspore" | |||
| 21 | 21 | ||
| 22 | from hyper_parallel.core.dtensor.dtensor import _build_layout | 22 | from hyper_parallel.core.dtensor.dtensor import _build_layout |
| 23 | from hyper_parallel.core.dtensor.placement_types import Shard, Replicate | 23 | from hyper_parallel.core.dtensor.placement_types import Shard, Replicate |
| 24 | -from hyper_parallel.core.shard.ops.parallel_npu_flash_attention_score import FlashAttentionScoreDistributedOp | 24 | +from hyper_parallel.core.shard.ops.parallel_npu_flash_attention_score import NPUFlashAttentionScoreDistributedOp |
| 25 | from hyper_parallel.platform import get_platform | 25 | from hyper_parallel.platform import get_platform |
| 26 | from hyper_parallel.core.dtensor.device_mesh import ( | 26 | from 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 | ) |
| 30 | from hyper_parallel.platform.platform import EXISTING_COMM_GROUPS | 30 | from 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 | ||
| 35 | class TestParallelNpuFlashAttentionScore(unittest.TestCase): | 35 | class 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""" |
| 16 | import os | 16 | import os |
| 17 | import unittest | 17 | import unittest |
| 18 | -from unittest.mock import patch | 18 | +from unittest.mock import MagicMock, patch |
| 19 | + | ||
| 19 | import numpy as np | 20 | import numpy as np |
| 20 | -os.environ["HYPER_PARALLEL_PLATFORM"] = "mindspore" | 21 | +os.environ["HYPER_PARALLEL_PLATFORM"] = "torch" |
| 21 | 22 | ||
| 22 | from hyper_parallel.core.dtensor.dtensor import _build_layout | 23 | from hyper_parallel.core.dtensor.dtensor import _build_layout |
| 23 | from hyper_parallel.core.dtensor.placement_types import Shard, Replicate | 24 | from hyper_parallel.core.dtensor.placement_types import Shard, Replicate |
| 24 | from hyper_parallel.core.shard.ops.parallel_sort import SortDistributedOp | 25 | from hyper_parallel.core.shard.ops.parallel_sort import SortDistributedOp |
| 25 | -from hyper_parallel.platform import get_platform | ||
| 26 | from hyper_parallel.core.dtensor.device_mesh import ( | 26 | from hyper_parallel.core.dtensor.device_mesh import ( |
| 27 | init_device_mesh, | 27 | init_device_mesh, |
| 28 | _DEVICE_MESH_MAP | 28 | _DEVICE_MESH_MAP |
| @@ -33,69 +33,53 @@ op = SortDistributedOp("sort") | |||
| 33 | 33 | ||
| 34 | 34 | ||
| 35 | class TestParallelSort(unittest.TestCase): | 35 | class 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 | 69 | ||
| 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_layouts | 84 | 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 | 97 | ||
| 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 | 108 | ||
| 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_layouts | 118 | 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_map | 121 | assert values_layout.tensor_map == expected_map |
| 151 | assert indices_layout.tensor_map == expected_map | 122 | assert indices_layout.tensor_map == expected_map |
| 123 | + assert extra_info is None | ||
| 152 | 124 | ||
| 153 | 125 | ||
| 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_layouts | 135 | 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_map | 138 | assert values_layout.tensor_map == expected_map |
| 170 | assert indices_layout.tensor_map == expected_map | 139 | assert indices_layout.tensor_map == expected_map |
| 140 | + assert extra_info is None | ||
| 171 | 141 | ||
| 172 | 142 | ||
| 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_layouts | 152 | values_layout, _ = output_layouts |
| 186 | expected_map = (-1, -1) | 153 | expected_map = (-1, -1) |
| 187 | 154 | ||
| 188 | assert values_layout.tensor_map == expected_map | 155 | assert values_layout.tensor_map == expected_map |
| 156 | + assert extra_info is None | ||
| 189 | 157 | ||
| 190 | 158 | ||
| 191 | if __name__ == "__main__": | 159 | if __name__ == "__main__": |
| @@ -18,13 +18,26 @@ Unit tests for OpDispatcher with custom distributed ops (e.g., StackExt). | |||
| 18 | import importlib | 18 | import importlib |
| 19 | import os | 19 | import os |
| 20 | import sys | 20 | import sys |
| 21 | +import unittest | ||
| 21 | from pathlib import Path | 22 | from pathlib import Path |
| 22 | from typing import Optional, Tuple | 23 | from typing import Optional, Tuple |
| 24 | +from unittest.mock import MagicMock, patch | ||
| 23 | 25 | ||
| 26 | +import numpy as np | ||
| 24 | import pytest | 27 | import pytest |
| 25 | 28 | ||
| 26 | from hyper_parallel.core.dtensor.layout import Layout | 29 | from hyper_parallel.core.dtensor.layout import Layout |
| 27 | from hyper_parallel.core.dtensor.dtensor import DTensor | 30 | from 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().parent | 42 | _TEST_FILE_DIR = Path(__file__).resolve().parent |
| 30 | _TESTS_ROOT_DIR = _TEST_FILE_DIR.parent.parent.parent.parent | 43 | _TESTS_ROOT_DIR = _TEST_FILE_DIR.parent.parent.parent.parent |
| @@ -33,92 +46,20 @@ _CUSTOM_OPS_DIR = _TESTS_ROOT_DIR / "tests" / "custom_ops" | |||
| 33 | HYPER_PARALLEL_OPS_YAML_DIR = str(_CUSTOM_OPS_DIR) | 46 | HYPER_PARALLEL_OPS_YAML_DIR = str(_CUSTOM_OPS_DIR) |
| 34 | HYPER_PARALLEL_OPS_PYTHON_PATH = str(_CUSTOM_OPS_DIR) | 47 | HYPER_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 variable | 52 | + 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 mod | 70 | return mod |
| 130 | 71 | ||
| 131 | 72 | ||
| 132 | -def _require_mindspore(): | 73 | +class TestStackExtDispatch(unittest.TestCase): |
| 133 | """ | 74 | """ |
| 134 | - Feature: Conditional dependency gate for MindSpore | 75 | + 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 | + | ||
| 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 | + | ||
| 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 tests | 187 | + _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, d1 | 217 | + 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 | + | ||
| 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-like | 255 | + 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=True | 272 | + 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 integration | 285 | + """Test that dispatch normalizes args and kwargs.""" |
| 219 | - Description: Reload _op_dispatch with custom YAML/Python paths from environment, create two replicated | 286 | + 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 = 0 | 289 | + 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-access | 297 | + 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 correct | 312 | + 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 | ||