已合并
[master]refactor: move shape handling into dynamo utils (#4241) #45007
黄桂军创建于 8月20日
[master]refactor: move shape handling into dynamo utils (#4241) #45007
已合并
共 4 个文件变更+549-511
| @@ -23,6 +23,7 @@ device = "npu" | |||
| 23 | def model_fn(A, B): | 23 | def model_fn(A, B): |
| 24 | return A + B | 24 | return A + B |
| 25 | 25 | ||
| 26 | + | ||
| 26 | shape_options = { | 27 | shape_options = { |
| 27 | "enable_shape_handling": True, | 28 | "enable_shape_handling": True, |
| 28 | "shape_handling_configs": [ | 29 | "shape_handling_configs": [ |
| @@ -436,12 +436,34 @@ class TorchCompileTriggerTests(unittest.TestCase): | |||
| 436 | """ | 436 | """ |
| 437 | ) | 437 | ) |
| 438 | 438 | ||
| 439 | - # Verify shape handling is installed only after the selected backend scope. | 439 | + # Verify scope entry failures restore the process environment. |
| 440 | - def test_shape_handling_initializes_after_backend_selection(self): | 440 | + def test_npu_backend_scope_restores_env_after_entry_failure(self): |
| 441 | + self.run_in_subprocess( | ||
| 442 | + """ | ||
| 443 | + import os | ||
| 444 | + | ||
| 445 | + import torch_npu | ||
| 446 | + from torch_npu.utils import _dynamo | ||
| 447 | + | ||
| 448 | + env_name = "TORCHINDUCTOR_NPU_BACKEND" | ||
| 449 | + original_env = os.environ.get(env_name) | ||
| 450 | + try: | ||
| 451 | + with _dynamo._NpuBackendScope(1): | ||
| 452 | + raise AssertionError("scope entry must fail") | ||
| 453 | + except TypeError as error: | ||
| 454 | + assert "str" in str(error) | ||
| 455 | + else: | ||
| 456 | + raise AssertionError("a non-string backend must fail") | ||
| 457 | + | ||
| 458 | + assert os.environ.get(env_name) == original_env | ||
| 459 | + """ | ||
| 460 | + ) | ||
| 461 | + | ||
| 462 | + # Verify shape handling is installed before selecting the requested backend. | ||
| 463 | + def test_shape_handling_initializes_before_backend_selection(self): | ||
| 441 | self.run_in_subprocess( | 464 | self.run_in_subprocess( |
| 442 | """ | 465 | """ |
| 443 | import os | 466 | import os |
| 444 | - import types | ||
| 445 | from unittest import mock | 467 | from unittest import mock |
| 446 | 468 | ||
| 447 | import torch | 469 | import torch |
| @@ -455,29 +477,52 @@ class TorchCompileTriggerTests(unittest.TestCase): | |||
| 455 | events = [] | 477 | events = [] |
| 456 | 478 | ||
| 457 | def scope_register(): | 479 | def scope_register(): |
| 458 | - events.append( | 480 | + events.append(("scope_register", os.environ["TORCHINDUCTOR_NPU_BACKEND"])) |
| 459 | - ("scope_register", os.environ.get("TORCHINDUCTOR_NPU_BACKEND")) | ||
| 460 | - ) | ||
| 461 | 481 | ||
| 462 | - fake_inductor = types.SimpleNamespace( | ||
| 463 | - patch_shape_handling=lambda: events.append( | ||
| 464 | - ("shape_handling", None) | ||
| 465 | - ) | ||
| 466 | - ) | ||
| 467 | with mock.patch.object( | 482 | with mock.patch.object( |
| 468 | - _dynamo, "_lazy_dynamo_setup" | 483 | + _dynamo, "_lazy_dynamo_setup", lambda: None |
| 469 | ), mock.patch.object( | 484 | ), mock.patch.object( |
| 470 | - _dynamo, "_lazy_inductor_setup" | 485 | + _dynamo, "_patch_shape_handling", |
| 486 | + lambda: events.append(("shape_handling", None)), | ||
| 487 | + ), mock.patch.object( | ||
| 488 | + _dynamo, "_lazy_inductor_setup", lambda: None | ||
| 471 | ), mock.patch.object( | 489 | ), mock.patch.object( |
| 472 | _dynamo, "register_inductor_npu", scope_register | 490 | _dynamo, "register_inductor_npu", scope_register |
| 473 | - ), mock.patch.object( | ||
| 474 | - torch_npu, "_inductor", fake_inductor, create=True | ||
| 475 | ): | 491 | ): |
| 476 | wrapper = torch._TorchCompileInductorWrapper(None, options, None) | 492 | wrapper = torch._TorchCompileInductorWrapper(None, options, None) |
| 477 | 493 | ||
| 494 | + normalized_options = { | ||
| 495 | + "npu_backend": "mlir", | ||
| 496 | + "enable_shape_handling": True, | ||
| 497 | + } | ||
| 478 | assert wrapper.config["npu_backend"] == "mlir" | 498 | assert wrapper.config["npu_backend"] == "mlir" |
| 479 | assert wrapper.config["enable_shape_handling"] is True | 499 | assert wrapper.config["enable_shape_handling"] is True |
| 480 | - assert events == [("scope_register", "mlir")], events | 500 | + assert events == [ |
| 501 | + ("shape_handling", None), | ||
| 502 | + ("scope_register", "mlir"), | ||
| 503 | + ], events | ||
| 504 | + """ | ||
| 505 | + ) | ||
| 506 | + | ||
| 507 | + # Shape Handling must not defeat deferred NPU Inductor loading. | ||
| 508 | + def test_shape_handling_backend_load_is_deferred_until_first_call(self): | ||
| 509 | + self.run_in_subprocess( | ||
| 510 | + """ | ||
| 511 | + import sys | ||
| 512 | + | ||
| 513 | + import torch | ||
| 514 | + import torch_npu | ||
| 515 | + from torch_npu.utils import _dynamo | ||
| 516 | + | ||
| 517 | + torch.compile( | ||
| 518 | + lambda x: x + 1, | ||
| 519 | + backend="inductor", | ||
| 520 | + options={"enable_shape_handling": True}, | ||
| 521 | + ) | ||
| 522 | + | ||
| 523 | + assert _dynamo._lazy_dynamo_setup.has_run | ||
| 524 | + assert not _dynamo._lazy_inductor_setup.has_run | ||
| 525 | + assert "torch_npu._inductor" not in sys.modules | ||
| 481 | """ | 526 | """ |
| 482 | ) | 527 | ) |
| 483 | 528 | ||
| @@ -665,27 +710,6 @@ class TorchCompileTriggerTests(unittest.TestCase): | |||
| 665 | """ | 710 | """ |
| 666 | ) | 711 | ) |
| 667 | 712 | ||
| 668 | - # Verify creating an Inductor wrapper does not load the backend prematurely. | ||
| 669 | - def test_inductor_backend_load_is_deferred_until_first_call(self): | ||
| 670 | - self.run_in_subprocess( | ||
| 671 | - """ | ||
| 672 | - import sys | ||
| 673 | - import torch | ||
| 674 | - import torch_npu | ||
| 675 | - from torch_npu.utils import _dynamo | ||
| 676 | - | ||
| 677 | - torch.compile( | ||
| 678 | - lambda x: x + 1, | ||
| 679 | - backend="inductor", | ||
| 680 | - options={"enable_shape_handling": True}, | ||
| 681 | - ) | ||
| 682 | - | ||
| 683 | - assert _dynamo._lazy_dynamo_setup.has_run | ||
| 684 | - assert not _dynamo._lazy_inductor_setup.has_run | ||
| 685 | - assert "torch_npu._inductor" not in sys.modules | ||
| 686 | - """ | ||
| 687 | - ) | ||
| 688 | - | ||
| 689 | # Verify lazy setup completes before compile backend lookup. | 713 | # Verify lazy setup completes before compile backend lookup. |
| 690 | def test_compile_triggers_setup_before_backend_lookup(self): | 714 | def test_compile_triggers_setup_before_backend_lookup(self): |
| 691 | self.run_in_subprocess( | 715 | self.run_in_subprocess( |
| @@ -1,477 +1,14 @@ | |||
| 1 | +from torch_npu.utils import _dynamo as _impl | ||
| 2 | +from torch_npu.utils._dynamo import NPUShapeHandling | ||
| 3 | + | ||
| 1 | __all__ = ["NPUShapeHandling"] | 4 | __all__ = ["NPUShapeHandling"] |
| 2 | 5 | ||
| 3 | -from typing import Any, Callable, Dict, List, Optional, Tuple | 6 | +unified_copy = _impl.unified_copy |
| 4 | -import copy | ||
| 5 | -import logging | ||
| 6 | -import warnings | ||
| 7 | -from torch.utils._pytree import tree_flatten, tree_unflatten, TreeSpec | ||
| 8 | -import torch | ||
| 9 | -import torch_npu._C | ||
| 10 | - | ||
| 11 | - | ||
| 12 | -class NPUShapeHandling(torch_npu._C._NPUShapeHandling): | ||
| 13 | - r"""Wrapper around a NPU shape handling configuration. | ||
| 14 | - Args: | ||
| 15 | - configs: List of configuration dictionaries that define shape handling rules. | ||
| 16 | - transform_pre_fn: Pre-processing function to convert inputs to tensor lists for transformation (optional). | ||
| 17 | - transform_post_fn: Post-processing function to convert tensor lists to structured outputs for transformation (optional). | ||
| 18 | - recover_pre_fn: Pre-processing function to convert inputs to tensor lists for recovery (optional). | ||
| 19 | - recover_post_fn: Post-processing function tp convert tensor lists to structured outputs for recovery (optional). | ||
| 20 | - | ||
| 21 | - Each config dictionary in configs supports the following keys: | ||
| 22 | - - type (str): | ||
| 23 | - Logical dimension type. Supported values: | ||
| 24 | - "BATCHSIZE" | "SEQLEN" | ||
| 25 | - - dimensions (int or List[int]): | ||
| 26 | - For BATCHSIZE: all affected tensors must share the same batch dimensions, if a list is provided, only the | ||
| 27 | - first element is used. | ||
| 28 | - For SEQLEN: if an int or a single-element list is provided, the value is automatically applied to all affected tensors; | ||
| 29 | - if a list is provided, it specifies the sequence dimension position for each affected tensor respectively, | ||
| 30 | - allowing different tensors to have the sequence dimension at different positions. | ||
| 31 | - - indices (List[int]): | ||
| 32 | - Indices of tensors that this rule applies to. Empty list means "apply to all tensors". | ||
| 33 | - - value (float): | ||
| 34 | - Padding value used when increasing size to reach the next gear. | ||
| 35 | - - gears (List[int]): | ||
| 36 | - Explicit list of allowed sizes(gears). If non-empty, overrides min_size/max_size/policy. | ||
| 37 | - - min_size (int): | ||
| 38 | - Minimum allowed size for this dimension (inclusive). Default: 1. | ||
| 39 | - - max_size (int): | ||
| 40 | - Maximum allowed size for this dimension (inclusive). Default: 1024. | ||
| 41 | - - policy (str): | ||
| 42 | - Gear generation strategy. Supported values: | ||
| 43 | - "TIMES" | "CUSTOM" | ||
| 44 | - | ||
| 45 | - If no configs are provided at construction, a default configuration handling batch size on dimension 0 is created. | ||
| 46 | - """ | ||
| 47 | - def __init__( | ||
| 48 | - self, | ||
| 49 | - configs: List[Dict[str, Any]] = None, | ||
| 50 | - transform_pre_fn: Optional[Callable[..., List[torch.Tensor]]] = None, | ||
| 51 | - transform_post_fn: Optional[Callable[[List[List[torch.Tensor]]], Tuple[List[Tuple], List[Dict]]]] = None, | ||
| 52 | - recover_pre_fn: Optional[Callable[[List[Any]], List[List[torch.Tensor]]]] = None, | ||
| 53 | - recover_post_fn: Optional[Callable[[List[torch.Tensor]], torch.Tensor]] = None, | ||
| 54 | - ) -> None: | ||
| 55 | - super().__init__() | ||
| 56 | - self.delay_init = False | ||
| 57 | - self.shape_type_map = { | ||
| 58 | - "BATCHSIZE": torch_npu._C.ShapeType.BATCHSIZE, | ||
| 59 | - "SEQLEN": torch_npu._C.ShapeType.SEQLEN | ||
| 60 | - } | ||
| 61 | - | ||
| 62 | - self.policy_map = { | ||
| 63 | - "TIMES": torch_npu._C.ShapePolicy.TIMES, | ||
| 64 | - "CUSTOM": torch_npu._C.ShapePolicy.CUSTOM | ||
| 65 | - } | ||
| 66 | - | ||
| 67 | - # Register processing functions | ||
| 68 | - self.transform_pre_fn = transform_pre_fn | ||
| 69 | - self.transform_post_fn = transform_post_fn | ||
| 70 | - self.recover_pre_fn = recover_pre_fn | ||
| 71 | - self.recover_post_fn = recover_post_fn | ||
| 72 | - if configs and len(configs) > 0: | ||
| 73 | - self._validate_configs(configs) | ||
| 74 | - self.configs = configs | ||
| 75 | - self._initialize_from_configs(configs) | ||
| 76 | - else: | ||
| 77 | - self.delay_init = True | ||
| 78 | - self.configs = [{ | ||
| 79 | - "type": "BATCHSIZE", | ||
| 80 | - "dimensions": [0], | ||
| 81 | - "indices": [], | ||
| 82 | - "value": 0.0, | ||
| 83 | - "gears": [], | ||
| 84 | - "min_size": 1, | ||
| 85 | - "max_size": 1024, | ||
| 86 | - "policy": "TIMES" | ||
| 87 | - }] | ||
| 88 | - | ||
| 89 | - def _validate_configs(self, configs: List[Dict[str, Any]]) -> None: | ||
| 90 | - if not configs or len(configs) == 0: | ||
| 91 | - return | ||
| 92 | - if len(configs) > 2: | ||
| 93 | - raise ValueError("NPUShapeHandling currently supports only two dimensions.") | ||
| 94 | - | ||
| 95 | - required_fields = ["type"] | ||
| 96 | - int_list_fields = ["dimensions", "indices", "gears"] | ||
| 97 | - int_fields = ["min_size", "max_size"] | ||
| 98 | - for i, config in enumerate(configs): | ||
| 99 | - for field in required_fields: | ||
| 100 | - if field not in config: | ||
| 101 | - raise ValueError(f"Config {i} missing required field: {field}.") | ||
| 102 | - | ||
| 103 | - if not isinstance(config["type"], str): | ||
| 104 | - raise ValueError(f"Config {i} {field} must be a str, got {type(config['type'])}.") | ||
| 105 | - if config["type"] not in self.shape_type_map: | ||
| 106 | - raise ValueError( | ||
| 107 | - f"Invalid 'type' in config[{i}]: {config['type']}. " | ||
| 108 | - f"Must be one of: {', '.join(repr(k) for k in self.shape_type_map.keys())}." | ||
| 109 | - ) | ||
| 110 | - | ||
| 111 | - for field in int_list_fields: | ||
| 112 | - if field not in config: | ||
| 113 | - continue | ||
| 114 | - | ||
| 115 | - if field == "dimensions": | ||
| 116 | - if isinstance(config[field], int): | ||
| 117 | - config[field] = [config[field]] | ||
| 118 | - if config["type"] == "BATCHSIZE" and len(config[field]) > 1: | ||
| 119 | - warnings.warn("For BATCHSIZE, only the first element of 'dimensions' is used") | ||
| 120 | - config[field] = config[field][0] | ||
| 121 | - | ||
| 122 | - if not isinstance(config[field], (list, tuple)): | ||
| 123 | - raise ValueError(f"Config {i} {field} must be a list, got {type(config[field])}.") | ||
| 124 | - | ||
| 125 | - for item in config[field]: | ||
| 126 | - if not isinstance(item, int): | ||
| 127 | - raise ValueError(f"Config {i} {field} must contain integers, got {type(item)}.") | ||
| 128 | - | ||
| 129 | - for field in int_fields: | ||
| 130 | - if field not in config: | ||
| 131 | - continue | ||
| 132 | - if not isinstance(config[field], int): | ||
| 133 | - raise ValueError(f"Config {i} {field} must be an integer, got {type(config[field])}.") | ||
| 134 | - | ||
| 135 | - if "value" in config and not isinstance(config["value"], (int, float)): | ||
| 136 | - raise ValueError(f"Config {i} 'value' must be a number, got {type(config['value'])}.") | ||
| 137 | - | ||
| 138 | - if "policy" in config: | ||
| 139 | - if not isinstance(config["policy"], str): | ||
| 140 | - raise ValueError(f"Config {i} 'policy' must be a str, got {type(config['policy'])}.") | ||
| 141 | - if config["policy"] not in self.policy_map: | ||
| 142 | - raise ValueError( | ||
| 143 | - f"Invalid 'policy' in config[{i}]: {config['policy']}. " | ||
| 144 | - f"Must be one of: {', '.join(repr(k) for k in self.policy_map.keys())}." | ||
| 145 | - ) | ||
| 146 | - | ||
| 147 | - if len(configs) == 2 and configs[0]["type"] == configs[1]["type"]: | ||
| 148 | - raise ValueError("Cannot initialize the same type repeatedly.") | ||
| 149 | - | ||
| 150 | - def _initialize_from_configs(self, configs: List[Dict[str, Any]]) -> None: | ||
| 151 | - for config in configs: | ||
| 152 | - shape_type = self.shape_type_map.get(config.get("type")) | ||
| 153 | - indices = config.get("indices", []) | ||
| 154 | - value = config.get("value", 0.0) | ||
| 155 | - gears = config.get("gears", []) | ||
| 156 | - | ||
| 157 | - dimensions = config.get("dimensions", []) | ||
| 158 | - if shape_type == torch_npu._C.ShapeType.BATCHSIZE: | ||
| 159 | - if not dimensions: | ||
| 160 | - # Empty list | ||
| 161 | - dimensions = [0] | ||
| 162 | - elif shape_type == torch_npu._C.ShapeType.SEQLEN and len(indices) != 0: | ||
| 163 | - if len(dimensions) == 1: | ||
| 164 | - dimensions = [dimensions[0] for _ in range(len(indices))] | ||
| 165 | - if not dimensions: | ||
| 166 | - dimensions = [1 for _ in range(len(indices))] | ||
| 167 | - | ||
| 168 | - | ||
| 169 | - if len(dimensions) == 0 or len(indices) == 0: | ||
| 170 | - self.delay_init = True | ||
| 171 | - continue | ||
| 172 | - | ||
| 173 | - if len(gears) > 0: | ||
| 174 | - self.initialize(shape_type, gears, dimensions, indices, value) | ||
| 175 | - else: | ||
| 176 | - min_size = config.get("min_size", 1) | ||
| 177 | - max_size = config.get("max_size", 1024) | ||
| 178 | - policy = self.policy_map.get(config.get("policy", "TIMES")) | ||
| 179 | - self.initialize(shape_type, min_size, max_size, policy, dimensions, indices, value) | ||
| 180 | - | ||
| 181 | - def _construct_indices(self, tensors: List[torch.Tensor], dimensions, dimension_type): | ||
| 182 | - if dimension_type == "BATCHSIZE": | ||
| 183 | - if not dimensions: | ||
| 184 | - dimensions = [0] | ||
| 185 | - dimensions = [dimensions[0] for _ in range(len(tensors))] | ||
| 186 | - | ||
| 187 | - if dimension_type == "SEQLEN": | ||
| 188 | - if not dimensions: | ||
| 189 | - dimensions = [1] | ||
| 190 | - if len(dimensions) == 1: | ||
| 191 | - dimensions = [dimensions[0] for _ in range(len(tensors))] | ||
| 192 | - | ||
| 193 | - index = 0 | ||
| 194 | - indices = [] | ||
| 195 | - for dimension, tensor in zip(dimensions, tensors): | ||
| 196 | - if tensor.ndim > dimension: | ||
| 197 | - indices.append(index) | ||
| 198 | - index += 1 | ||
| 199 | - | ||
| 200 | - return indices | ||
| 201 | - | ||
| 202 | - def delay_initialize(self, tensors: List[torch.Tensor]): | ||
| 203 | - delay_init_configs = [] | ||
| 204 | - for config in self.configs: | ||
| 205 | - init_flag = False | ||
| 206 | - if "indices" not in config or len(config["indices"]) == 0: | ||
| 207 | - init_flag = True | ||
| 208 | - config["indices"] = self._construct_indices(tensors, config.get("dimensions", []), config["type"]) | ||
| 209 | - | ||
| 210 | - if init_flag: | ||
| 211 | - delay_init_configs.append(config) | ||
| 212 | - if len(delay_init_configs) > 0: | ||
| 213 | - self._initialize_from_configs(delay_init_configs) | ||
| 214 | - self.delay_init = False | ||
| 215 | - | ||
| 216 | - def transform(self, tensors: List[torch.Tensor]) -> List[List[torch.Tensor]]: | ||
| 217 | - if self.delay_init: | ||
| 218 | - self.delay_initialize(tensors) | ||
| 219 | - return super().transform(tensors) | ||
| 220 | - | ||
| 221 | - def recover(self, tensor_groups: List[List[torch.Tensor]]) -> List[torch.Tensor]: | ||
| 222 | - return super().recover(tensor_groups) | ||
| 223 | - | ||
| 224 | - def get_shape_safe(self, item): | ||
| 225 | - """递归获取 shape 的辅助函数""" | ||
| 226 | - if isinstance(item, torch.Tensor): | ||
| 227 | - return list(item.shape) | ||
| 228 | - elif isinstance(item, (list, tuple)): | ||
| 229 | - # 如果是列表,递归处理内部元素,并标注这是个容器 | ||
| 230 | - return [self.get_shape_safe(i) for i in item] | ||
| 231 | - else: | ||
| 232 | - return type(item) | ||
| 233 | - | ||
| 234 | - | ||
| 235 | - def transform_hook( | ||
| 236 | - self, | ||
| 237 | - *args: Any, | ||
| 238 | - **kwargs: Any | ||
| 239 | - ) -> Tuple[List[Tuple], List[Dict]]: | ||
| 240 | - # 获取 logger | ||
| 241 | - logger = logging.getLogger(__name__) | ||
| 242 | - # 预处理阶段优化:统一使用预定义函数或默认逻辑 | ||
| 243 | - if self.transform_pre_fn: | ||
| 244 | - inputs = self.transform_pre_fn(*args, **kwargs) | ||
| 245 | - else: | ||
| 246 | - inputs, indices, leaves, spec = self._process_inputs(args, kwargs) | ||
| 247 | - | ||
| 248 | - # 提取转换前的形状 (inputs 通常是 Tensor 列表) | ||
| 249 | - if logger.isEnabledFor(logging.INFO): | ||
| 250 | - pre_shapes = [self.get_shape_safe(t) for t in inputs] | ||
| 251 | - logger.info("[Transform] Starting. Input tensors: %s, Shapes: %s", len(inputs), pre_shapes) | ||
| 252 | - | ||
| 253 | - # 执行核心转换操作 | ||
| 254 | - trans_outputs = self.transform(tensors=inputs) | ||
| 255 | - | ||
| 256 | - # 提取转换后的形状 | ||
| 257 | - if logger.isEnabledFor(logging.INFO): | ||
| 258 | - post_shapes = [self.get_shape_safe(t) for t in trans_outputs] | ||
| 259 | - logger.info("> Post-transform content: %s", post_shapes) | ||
| 260 | - | ||
| 261 | - # 后处理阶段优化:避免嵌套循环 | ||
| 262 | - if self.transform_post_fn: | ||
| 263 | - outputs = self.transform_post_fn(trans_outputs) | ||
| 264 | - else: | ||
| 265 | - outputs = self._recover_inputs(trans_outputs, indices, leaves, spec) | ||
| 266 | - | ||
| 267 | - if not outputs: | ||
| 268 | - logger.error("CRITICAL: _recover_inputs returned NULL") | ||
| 269 | - | ||
| 270 | - return outputs | ||
| 271 | - | ||
| 272 | - def flatten_to_tensors(self, structure: Any) -> Tuple[List[torch.Tensor], List[int], List[Any], TreeSpec]: | ||
| 273 | - leaves, spec = tree_flatten(structure) | ||
| 274 | - indexed_tensors = [(i, leaf) for i, leaf in enumerate(leaves) if isinstance(leaf, torch.Tensor)] | ||
| 275 | - indices = [] | ||
| 276 | - tensors = [] | ||
| 277 | - if indexed_tensors is not None and len(indexed_tensors) > 0: | ||
| 278 | - indices, tensors = zip(*indexed_tensors) | ||
| 279 | - return tensors, indices, leaves, spec | ||
| 280 | - | ||
| 281 | - def unflatten_from_tensors( | ||
| 282 | - self, | ||
| 283 | - tensors: List[torch.Tensor], | ||
| 284 | - indices: List[int], | ||
| 285 | - leaves: List[Any], | ||
| 286 | - spec: TreeSpec | ||
| 287 | - ) -> Any: | ||
| 288 | - for idx, tensor in zip(indices, tensors): | ||
| 289 | - leaves[idx] = tensor | ||
| 290 | - return tree_unflatten(leaves, spec) | ||
| 291 | - | ||
| 292 | - def _process_inputs(self, args: Tuple, kwargs: dict) -> List[torch.Tensor]: | ||
| 293 | - return self.flatten_to_tensors((args, kwargs)) | ||
| 294 | - | ||
| 295 | - def _recover_inputs( | ||
| 296 | - self, | ||
| 297 | - transform_res: List[List[torch.Tensor]], | ||
| 298 | - indices: List[int], | ||
| 299 | - leaves: List[Any], | ||
| 300 | - spec: TreeSpec | ||
| 301 | - ) -> Tuple[List[Tuple], List[Dict]]: | ||
| 302 | - res = [] | ||
| 303 | - for processd_tensors in transform_res: | ||
| 304 | - res.append(self.unflatten_from_tensors(processd_tensors, indices, list(leaves), spec)) | ||
| 305 | - return zip(*res) | ||
| 306 | - | ||
| 307 | - def _process_outputs( | ||
| 308 | - self, | ||
| 309 | - outputs_list: List[Any] | ||
| 310 | - ) -> Tuple[List[List[torch.Tensor]], List[int], List[Any], TreeSpec]: | ||
| 311 | - tensors_list = [] | ||
| 312 | - leaves = [] | ||
| 313 | - indices = [] | ||
| 314 | - spec = None | ||
| 315 | - for output in outputs_list: | ||
| 316 | - tensors, indices, leaves, spec = self.flatten_to_tensors(output) | ||
| 317 | - tensors_list.append(tensors) | ||
| 318 | - return tensors_list, indices, leaves, spec | ||
| 319 | - | ||
| 320 | - def _recover_outputs( | ||
| 321 | - self, | ||
| 322 | - recover_res: List[torch.Tensor], | ||
| 323 | - indices: List[int], | ||
| 324 | - leaves: List[Any], | ||
| 325 | - spec: TreeSpec | ||
| 326 | - ) -> Any: | ||
| 327 | - return self.unflatten_from_tensors(recover_res, indices, leaves, spec) | ||
| 328 | - | ||
| 329 | - def recover_hook( | ||
| 330 | - self, | ||
| 331 | - groups: List[Any] | ||
| 332 | - ) -> Any: | ||
| 333 | - """ | ||
| 334 | - Process input groups through recovery pipeline. | ||
| 335 | - | ||
| 336 | - Args: | ||
| 337 | - groups: List of input data to be processed. | ||
| 338 | - | ||
| 339 | - Returns: | ||
| 340 | - Processed outputs after recovery and postprocessing. | ||
| 341 | - """ | ||
| 342 | - # 预处理:使用自定义函数或默认方法 | ||
| 343 | - if self.recover_pre_fn: | ||
| 344 | - inputs = self.recover_pre_fn(groups) | ||
| 345 | - else: | ||
| 346 | - inputs, indices, leaves, spec = self._process_outputs(groups) | ||
| 347 | - | ||
| 348 | - # 执行恢复操作 | ||
| 349 | - re_outputs = self.recover(tensor_groups=inputs) | ||
| 350 | - | ||
| 351 | - # 后处理:使用自定义函数或默认方法 | ||
| 352 | - if self.recover_post_fn: | ||
| 353 | - outputs = self.recover_post_fn(re_outputs) | ||
| 354 | - else: | ||
| 355 | - outputs = self._recover_outputs(re_outputs, indices, leaves, spec) | ||
| 356 | - | ||
| 357 | - return outputs | ||
| 358 | - | ||
| 359 | - | ||
| 360 | -def unified_copy(data: Any) -> Any: | ||
| 361 | - """ | ||
| 362 | - 对输入数据进行安全且统一的深拷贝。 | ||
| 363 | - 支持PyTorch Tensor、字典、列表等常见数据类型。 | ||
| 364 | - | ||
| 365 | - Args: | ||
| 366 | - data: 输入数据,可以是Tensor、dict、list等 | ||
| 367 | - | ||
| 368 | - Returns: | ||
| 369 | - 数据的独立副本 | ||
| 370 | - """ | ||
| 371 | - if data is None: | ||
| 372 | - return None | ||
| 373 | - | ||
| 374 | - # 处理PyTorch Tensor | ||
| 375 | - if isinstance(data, torch.Tensor): | ||
| 376 | - return data.clone().detach() | ||
| 377 | - | ||
| 378 | - # 处理字典类型 | ||
| 379 | - elif isinstance(data, dict): | ||
| 380 | - return {key: unified_copy(value) for key, value in data.items()} | ||
| 381 | - | ||
| 382 | - # 处理列表类型 | ||
| 383 | - elif isinstance(data, list): | ||
| 384 | - return [unified_copy(item) for item in data] | ||
| 385 | - | ||
| 386 | - # 处理元组类型 | ||
| 387 | - elif isinstance(data, tuple): | ||
| 388 | - return tuple(unified_copy(item) for item in data) | ||
| 389 | - | ||
| 390 | - else: | ||
| 391 | - try: | ||
| 392 | - return copy.deepcopy(data) | ||
| 393 | - except (TypeError, ValueError): | ||
| 394 | - return data | ||
| 395 | 7 | ||
| 396 | 8 | ||
| 397 | def patch_dynamo_context(): | 9 | def patch_dynamo_context(): |
| 398 | - import contextlib | 10 | + if not getattr(_impl._patch_shape_handling, "_is_patched", False): |
| 399 | - import inspect | 11 | + _impl._patch_dynamo_context() |
| 400 | - from torch._dynamo.eval_frame import _TorchDynamoContext | ||
| 401 | - from torch._dynamo.types import DynamoCallback | ||
| 402 | - from torch._dynamo.convert_frame import CatchErrorsWrapper, ConvertFrame | ||
| 403 | - from torch._dynamo.repro.after_dynamo import WrapBackendDebug | ||
| 404 | - src_call = _TorchDynamoContext.__call__ | ||
| 405 | - src_init = _TorchDynamoContext.__init__ | ||
| 406 | - null_context = contextlib.nullcontext | ||
| 407 | - | ||
| 408 | - def is_enable_shape_handling(callback: DynamoCallback, compiler_config=None): | ||
| 409 | - """ | ||
| 410 | - The shape handling feature is only available when enable_shape_handling is True and the backend is inductor | ||
| 411 | - """ | ||
| 412 | - if compiler_config is None or not compiler_config.get("enable_shape_handling", False): | ||
| 413 | - return False | ||
| 414 | - | ||
| 415 | - if callback is None or not isinstance(callback, CatchErrorsWrapper): | ||
| 416 | - return False | ||
| 417 | - | ||
| 418 | - convert_frame = getattr(callback, "_torchdynamo_orig_backend", None) | ||
| 419 | - if not isinstance(convert_frame, ConvertFrame): | ||
| 420 | - return False | ||
| 421 | - | ||
| 422 | - backend_debug = getattr(convert_frame, "_torchdynamo_orig_backend", None) | ||
| 423 | - if not isinstance(backend_debug, WrapBackendDebug): | ||
| 424 | - return False | ||
| 425 | - | ||
| 426 | - return getattr(backend_debug, "_compiler_name", None) == "inductor" | ||
| 427 | - | ||
| 428 | - def nothing(): | ||
| 429 | - pass | ||
| 430 | - | ||
| 431 | - def new_init(self, callback: DynamoCallback, *args, **kwargs) -> None: | ||
| 432 | - src_init(self, callback, *args, **kwargs) | ||
| 433 | - compiler_config = kwargs.get("compiler_config") | ||
| 434 | - if (is_enable_shape_handling(callback, compiler_config=compiler_config)): | ||
| 435 | - trans_pre_fn = None | ||
| 436 | - trans_post_fn = None | ||
| 437 | - re_pre_fn = None | ||
| 438 | - re_post_fn = None | ||
| 439 | - function_dict = compiler_config.get("shape_handling_dict") | ||
| 440 | - if function_dict is not None: | ||
| 441 | - trans_pre_fn = function_dict.get("trans_pre_fn", None) | ||
| 442 | - trans_post_fn = function_dict.get("trans_post_fn", None) | ||
| 443 | - re_pre_fn = function_dict.get("re_pre_fn", None) | ||
| 444 | - re_post_fn = function_dict.get("re_post_fn", None) | ||
| 445 | - | ||
| 446 | - self.shape_handling = NPUShapeHandling( | ||
| 447 | - configs=compiler_config.get("shape_handling_configs"), | ||
| 448 | - transform_pre_fn=trans_pre_fn, | ||
| 449 | - transform_post_fn=trans_post_fn, | ||
| 450 | - recover_pre_fn=re_pre_fn, | ||
| 451 | - recover_post_fn=re_post_fn, | ||
| 452 | - ) | ||
| 453 | - | ||
| 454 | - def new_call(self, fn): | ||
| 455 | - src_fn = src_call(self, fn) | ||
| 456 | - if isinstance(fn, torch.nn.Module) or inspect.isclass(fn): | ||
| 457 | - return src_fn | ||
| 458 | - | ||
| 459 | - def new_fn(*args, **kwargs): | ||
| 460 | - if (is_enable_shape_handling(self.callback, compiler_config=self.compiler_config)): | ||
| 461 | - new_args, new_kwargs = self.shape_handling.transform_hook(*args, **kwargs) | ||
| 462 | - args_is_split = len(args) != 0 and len(new_args) > 1 | ||
| 463 | - kwargs_is_split = len(kwargs) != 0 and len(new_kwargs) > 1 | ||
| 464 | - zipped_params = zip(new_args, new_kwargs) | ||
| 465 | - res = [ | ||
| 466 | - unified_copy(src_fn(*arg, **kwargs)) if args_is_split or kwargs_is_split | ||
| 467 | - else src_fn(*arg, **kwargs) | ||
| 468 | - for arg, kwargs in zipped_params | ||
| 469 | - ] | ||
| 470 | - return self.shape_handling.recover_hook(res) | ||
| 471 | - return src_fn(*args, **kwargs) | ||
| 472 | - return new_fn | ||
| 473 | - _TorchDynamoContext.__call__ = new_call | ||
| 474 | - _TorchDynamoContext.__init__ = new_init | ||
| 475 | 12 | ||
| 476 | 13 | ||
| 477 | def patch_shape_handling(): | 14 | def patch_shape_handling(): |
| @@ -479,3 +16,4 @@ def patch_shape_handling(): | |||
| 479 | return | 16 | return |
| 480 | patch_dynamo_context() | 17 | patch_dynamo_context() |
| 481 | patch_shape_handling._is_patched = True | 18 | patch_shape_handling._is_patched = True |
| 19 | + _impl._patch_shape_handling._is_patched = True | ||
| @@ -1,15 +1,19 @@ | |||
| 1 | import importlib | 1 | import importlib |
| 2 | import importlib.abc | 2 | import importlib.abc |
| 3 | +import copy | ||
| 3 | import functools | 4 | import functools |
| 4 | import inspect | 5 | import inspect |
| 5 | import logging | 6 | import logging |
| 6 | import os | 7 | import os |
| 7 | import sys | 8 | import sys |
| 8 | import threading | 9 | import threading |
| 9 | -from typing import Any, Optional, TYPE_CHECKING | 10 | +import warnings |
| 11 | +from typing import Any, Callable, Dict, List, Optional, Tuple, TYPE_CHECKING | ||
| 10 | 12 | ||
| 11 | import torch | 13 | import torch |
| 12 | import torch_npu | 14 | import torch_npu |
| 15 | +import torch_npu._C | ||
| 16 | +from torch.utils._pytree import TreeSpec, tree_flatten, tree_unflatten | ||
| 13 | from torch import _TorchCompileWrapper | 17 | from torch import _TorchCompileWrapper |
| 14 | 18 | ||
| 15 | 19 | ||
| @@ -210,8 +214,8 @@ def patch_inductor_wrapper(): | |||
| 210 | if shape_handling_requested: | 214 | if shape_handling_requested: |
| 211 | if getattr(self, "_npu_defer_shape_handling", False): | 215 | if getattr(self, "_npu_defer_shape_handling", False): |
| 212 | self._npu_shape_handling_requested = True | 216 | self._npu_shape_handling_requested = True |
| 213 | - # Shape handling is installed in new_call, after the selected | 217 | + else: |
| 214 | - # backend scope has loaded the matching NPU Inductor backend. | 218 | + _patch_shape_handling() |
| 215 | 219 | ||
| 216 | def new_get_config_copy(self) -> dict[str, Any]: | 220 | def new_get_config_copy(self) -> dict[str, Any]: |
| 217 | ori_dict = src_get_config_copy(self) | 221 | ori_dict = src_get_config_copy(self) |
| @@ -255,10 +259,13 @@ def patch_inductor_wrapper(): | |||
| 255 | src_init(self, mode, options, dynamic, name) | 259 | src_init(self, mode, options, dynamic, name) |
| 256 | else: | 260 | else: |
| 257 | src_init(self, mode, options, dynamic) | 261 | src_init(self, mode, options, dynamic) |
| 262 | + shape_handling_requested = self._npu_shape_handling_requested | ||
| 258 | finally: | 263 | finally: |
| 259 | del self._npu_defer_shape_handling | 264 | del self._npu_defer_shape_handling |
| 260 | del self._npu_shape_handling_requested | 265 | del self._npu_shape_handling_requested |
| 261 | _lazy_dynamo_setup() | 266 | _lazy_dynamo_setup() |
| 267 | + if shape_handling_requested: | ||
| 268 | + _patch_shape_handling() | ||
| 262 | backend = _resolve_npu_backend_from_wrapper(self) | 269 | backend = _resolve_npu_backend_from_wrapper(self) |
| 263 | if backend == "mlir": | 270 | if backend == "mlir": |
| 264 | with _NpuBackendScope(backend): | 271 | with _NpuBackendScope(backend): |
| @@ -272,8 +279,6 @@ def patch_inductor_wrapper(): | |||
| 272 | def new_call(self, model_, inputs_): | 279 | def new_call(self, model_, inputs_): |
| 273 | backend = _resolve_npu_backend_from_wrapper(self) | 280 | backend = _resolve_npu_backend_from_wrapper(self) |
| 274 | with _NpuBackendScope(backend): | 281 | with _NpuBackendScope(backend): |
| 275 | - if self.config.get("enable_shape_handling", False): | ||
| 276 | - torch_npu._inductor.patch_shape_handling() | ||
| 277 | if backend == "ascendc": | 282 | if backend == "ascendc": |
| 278 | from torch_npu.dynamo._deterministic_guard import ( | 283 | from torch_npu.dynamo._deterministic_guard import ( |
| 279 | install_npu_deterministic_level_guard, | 284 | install_npu_deterministic_level_guard, |
| @@ -768,3 +773,473 @@ def add_dynamo_methods(): | |||
| 768 | sys.modules["npugraph_ex"] = _LazyNpuGraphEx("npugraph_ex") | 773 | sys.modules["npugraph_ex"] = _LazyNpuGraphEx("npugraph_ex") |
| 769 | patch_inductor_wrapper() | 774 | patch_inductor_wrapper() |
| 770 | install_npugraph_mark_step_trigger() | 775 | install_npugraph_mark_step_trigger() |
| 776 | + | ||
| 777 | + | ||
| 778 | +class NPUShapeHandling(torch_npu._C._NPUShapeHandling): | ||
| 779 | + r"""Wrapper around a NPU shape handling configuration. | ||
| 780 | + Args: | ||
| 781 | + configs: List of configuration dictionaries that define shape handling rules. | ||
| 782 | + transform_pre_fn: Pre-processing function to convert inputs to tensor lists for transformation (optional). | ||
| 783 | + transform_post_fn: Post-processing function to convert tensor lists to structured outputs for transformation (optional). | ||
| 784 | + recover_pre_fn: Pre-processing function to convert inputs to tensor lists for recovery (optional). | ||
| 785 | + recover_post_fn: Post-processing function tp convert tensor lists to structured outputs for recovery (optional). | ||
| 786 | + | ||
| 787 | + Each config dictionary in configs supports the following keys: | ||
| 788 | + - type (str): | ||
| 789 | + Logical dimension type. Supported values: | ||
| 790 | + "BATCHSIZE" | "SEQLEN" | ||
| 791 | + - dimensions (int or List[int]): | ||
| 792 | + For BATCHSIZE: all affected tensors must share the same batch dimensions, if a list is provided, only the | ||
| 793 | + first element is used. | ||
| 794 | + For SEQLEN: if an int or a single-element list is provided, the value is automatically applied to all affected tensors; | ||
| 795 | + if a list is provided, it specifies the sequence dimension position for each affected tensor respectively, | ||
| 796 | + allowing different tensors to have the sequence dimension at different positions. | ||
| 797 | + - indices (List[int]): | ||
| 798 | + Indices of tensors that this rule applies to. Empty list means "apply to all tensors". | ||
| 799 | + - value (float): | ||
| 800 | + Padding value used when increasing size to reach the next gear. | ||
| 801 | + - gears (List[int]): | ||
| 802 | + Explicit list of allowed sizes(gears). If non-empty, overrides min_size/max_size/policy. | ||
| 803 | + - min_size (int): | ||
| 804 | + Minimum allowed size for this dimension (inclusive). Default: 1. | ||
| 805 | + - max_size (int): | ||
| 806 | + Maximum allowed size for this dimension (inclusive). Default: 1024. | ||
| 807 | + - policy (str): | ||
| 808 | + Gear generation strategy. Supported values: | ||
| 809 | + "TIMES" | "CUSTOM" | ||
| 810 | + | ||
| 811 | + If no configs are provided at construction, a default configuration handling batch size on dimension 0 is created. | ||
| 812 | + """ | ||
| 813 | + def __init__( | ||
| 814 | + self, | ||
| 815 | + configs: List[Dict[str, Any]] = None, | ||
| 816 | + transform_pre_fn: Optional[Callable[..., List[torch.Tensor]]] = None, | ||
| 817 | + transform_post_fn: Optional[Callable[[List[List[torch.Tensor]]], Tuple[List[Tuple], List[Dict]]]] = None, | ||
| 818 | + recover_pre_fn: Optional[Callable[[List[Any]], List[List[torch.Tensor]]]] = None, | ||
| 819 | + recover_post_fn: Optional[Callable[[List[torch.Tensor]], torch.Tensor]] = None, | ||
| 820 | + ) -> None: | ||
| 821 | + super().__init__() | ||
| 822 | + self.delay_init = False | ||
| 823 | + self.shape_type_map = { | ||
| 824 | + "BATCHSIZE": torch_npu._C.ShapeType.BATCHSIZE, | ||
| 825 | + "SEQLEN": torch_npu._C.ShapeType.SEQLEN | ||
| 826 | + } | ||
| 827 | + | ||
| 828 | + self.policy_map = { | ||
| 829 | + "TIMES": torch_npu._C.ShapePolicy.TIMES, | ||
| 830 | + "CUSTOM": torch_npu._C.ShapePolicy.CUSTOM | ||
| 831 | + } | ||
| 832 | + | ||
| 833 | + # Register processing functions | ||
| 834 | + self.transform_pre_fn = transform_pre_fn | ||
| 835 | + self.transform_post_fn = transform_post_fn | ||
| 836 | + self.recover_pre_fn = recover_pre_fn | ||
| 837 | + self.recover_post_fn = recover_post_fn | ||
| 838 | + if configs and len(configs) > 0: | ||
| 839 | + self._validate_configs(configs) | ||
| 840 | + self.configs = configs | ||
| 841 | + self._initialize_from_configs(configs) | ||
| 842 | + else: | ||
| 843 | + self.delay_init = True | ||
| 844 | + self.configs = [{ | ||
| 845 | + "type": "BATCHSIZE", | ||
| 846 | + "dimensions": [0], | ||
| 847 | + "indices": [], | ||
| 848 | + "value": 0.0, | ||
| 849 | + "gears": [], | ||
| 850 | + "min_size": 1, | ||
| 851 | + "max_size": 1024, | ||
| 852 | + "policy": "TIMES" | ||
| 853 | + }] | ||
| 854 | + | ||
| 855 | + def _validate_configs(self, configs: List[Dict[str, Any]]) -> None: | ||
| 856 | + if not configs or len(configs) == 0: | ||
| 857 | + return | ||
| 858 | + if len(configs) > 2: | ||
| 859 | + raise ValueError("NPUShapeHandling currently supports only two dimensions.") | ||
| 860 | + | ||
| 861 | + required_fields = ["type"] | ||
| 862 | + int_list_fields = ["dimensions", "indices", "gears"] | ||
| 863 | + int_fields = ["min_size", "max_size"] | ||
| 864 | + for i, config in enumerate(configs): | ||
| 865 | + for field in required_fields: | ||
| 866 | + if field not in config: | ||
| 867 | + raise ValueError(f"Config {i} missing required field: {field}.") | ||
| 868 | + | ||
| 869 | + if not isinstance(config["type"], str): | ||
| 870 | + raise ValueError(f"Config {i} {field} must be a str, got {type(config['type'])}.") | ||
| 871 | + if config["type"] not in self.shape_type_map: | ||
| 872 | + raise ValueError( | ||
| 873 | + f"Invalid 'type' in config[{i}]: {config['type']}. " | ||
| 874 | + f"Must be one of: {', '.join(repr(k) for k in self.shape_type_map.keys())}." | ||
| 875 | + ) | ||
| 876 | + | ||
| 877 | + for field in int_list_fields: | ||
| 878 | + if field not in config: | ||
| 879 | + continue | ||
| 880 | + | ||
| 881 | + if field == "dimensions": | ||
| 882 | + if isinstance(config[field], int): | ||
| 883 | + config[field] = [config[field]] | ||
| 884 | + if config["type"] == "BATCHSIZE" and len(config[field]) > 1: | ||
| 885 | + warnings.warn("For BATCHSIZE, only the first element of 'dimensions' is used") | ||
| 886 | + config[field] = config[field][0] | ||
| 887 | + | ||
| 888 | + if not isinstance(config[field], (list, tuple)): | ||
| 889 | + raise ValueError(f"Config {i} {field} must be a list, got {type(config[field])}.") | ||
| 890 | + | ||
| 891 | + for item in config[field]: | ||
| 892 | + if not isinstance(item, int): | ||
| 893 | + raise ValueError(f"Config {i} {field} must contain integers, got {type(item)}.") | ||
| 894 | + | ||
| 895 | + for field in int_fields: | ||
| 896 | + if field not in config: | ||
| 897 | + continue | ||
| 898 | + if not isinstance(config[field], int): | ||
| 899 | + raise ValueError(f"Config {i} {field} must be an integer, got {type(config[field])}.") | ||
| 900 | + | ||
| 901 | + if "value" in config and not isinstance(config["value"], (int, float)): | ||
| 902 | + raise ValueError(f"Config {i} 'value' must be a number, got {type(config['value'])}.") | ||
| 903 | + | ||
| 904 | + if "policy" in config: | ||
| 905 | + if not isinstance(config["policy"], str): | ||
| 906 | + raise ValueError(f"Config {i} 'policy' must be a str, got {type(config['policy'])}.") | ||
| 907 | + if config["policy"] not in self.policy_map: | ||
| 908 | + raise ValueError( | ||
| 909 | + f"Invalid 'policy' in config[{i}]: {config['policy']}. " | ||
| 910 | + f"Must be one of: {', '.join(repr(k) for k in self.policy_map.keys())}." | ||
| 911 | + ) | ||
| 912 | + | ||
| 913 | + if len(configs) == 2 and configs[0]["type"] == configs[1]["type"]: | ||
| 914 | + raise ValueError("Cannot initialize the same type repeatedly.") | ||
| 915 | + | ||
| 916 | + def _initialize_from_configs(self, configs: List[Dict[str, Any]]) -> None: | ||
| 917 | + for config in configs: | ||
| 918 | + shape_type = self.shape_type_map.get(config.get("type")) | ||
| 919 | + indices = config.get("indices", []) | ||
| 920 | + value = config.get("value", 0.0) | ||
| 921 | + gears = config.get("gears", []) | ||
| 922 | + | ||
| 923 | + dimensions = config.get("dimensions", []) | ||
| 924 | + if shape_type == torch_npu._C.ShapeType.BATCHSIZE: | ||
| 925 | + if not dimensions: | ||
| 926 | + # Empty list | ||
| 927 | + dimensions = [0] | ||
| 928 | + elif shape_type == torch_npu._C.ShapeType.SEQLEN and len(indices) != 0: | ||
| 929 | + if len(dimensions) == 1: | ||
| 930 | + dimensions = [dimensions[0] for _ in range(len(indices))] | ||
| 931 | + if not dimensions: | ||
| 932 | + dimensions = [1 for _ in range(len(indices))] | ||
| 933 | + | ||
| 934 | + | ||
| 935 | + if len(dimensions) == 0 or len(indices) == 0: | ||
| 936 | + self.delay_init = True | ||
| 937 | + continue | ||
| 938 | + | ||
| 939 | + if len(gears) > 0: | ||
| 940 | + self.initialize(shape_type, gears, dimensions, indices, value) | ||
| 941 | + else: | ||
| 942 | + min_size = config.get("min_size", 1) | ||
| 943 | + max_size = config.get("max_size", 1024) | ||
| 944 | + policy = self.policy_map.get(config.get("policy", "TIMES")) | ||
| 945 | + self.initialize(shape_type, min_size, max_size, policy, dimensions, indices, value) | ||
| 946 | + | ||
| 947 | + def _construct_indices(self, tensors: List[torch.Tensor], dimensions, dimension_type): | ||
| 948 | + if dimension_type == "BATCHSIZE": | ||
| 949 | + if not dimensions: | ||
| 950 | + dimensions = [0] | ||
| 951 | + dimensions = [dimensions[0] for _ in range(len(tensors))] | ||
| 952 | + | ||
| 953 | + if dimension_type == "SEQLEN": | ||
| 954 | + if not dimensions: | ||
| 955 | + dimensions = [1] | ||
| 956 | + if len(dimensions) == 1: | ||
| 957 | + dimensions = [dimensions[0] for _ in range(len(tensors))] | ||
| 958 | + | ||
| 959 | + index = 0 | ||
| 960 | + indices = [] | ||
| 961 | + for dimension, tensor in zip(dimensions, tensors): | ||
| 962 | + if tensor.ndim > dimension: | ||
| 963 | + indices.append(index) | ||
| 964 | + index += 1 | ||
| 965 | + | ||
| 966 | + return indices | ||
| 967 | + | ||
| 968 | + def delay_initialize(self, tensors: List[torch.Tensor]): | ||
| 969 | + delay_init_configs = [] | ||
| 970 | + for config in self.configs: | ||
| 971 | + init_flag = False | ||
| 972 | + if "indices" not in config or len(config["indices"]) == 0: | ||
| 973 | + init_flag = True | ||
| 974 | + config["indices"] = self._construct_indices(tensors, config.get("dimensions", []), config["type"]) | ||
| 975 | + | ||
| 976 | + if init_flag: | ||
| 977 | + delay_init_configs.append(config) | ||
| 978 | + if len(delay_init_configs) > 0: | ||
| 979 | + self._initialize_from_configs(delay_init_configs) | ||
| 980 | + self.delay_init = False | ||
| 981 | + | ||
| 982 | + def transform(self, tensors: List[torch.Tensor]) -> List[List[torch.Tensor]]: | ||
| 983 | + if self.delay_init: | ||
| 984 | + self.delay_initialize(tensors) | ||
| 985 | + return super().transform(tensors) | ||
| 986 | + | ||
| 987 | + def recover(self, tensor_groups: List[List[torch.Tensor]]) -> List[torch.Tensor]: | ||
| 988 | + return super().recover(tensor_groups) | ||
| 989 | + | ||
| 990 | + def get_shape_safe(self, item): | ||
| 991 | + """递归获取 shape 的辅助函数""" | ||
| 992 | + if isinstance(item, torch.Tensor): | ||
| 993 | + return list(item.shape) | ||
| 994 | + elif isinstance(item, (list, tuple)): | ||
| 995 | + # 如果是列表,递归处理内部元素,并标注这是个容器 | ||
| 996 | + return [self.get_shape_safe(i) for i in item] | ||
| 997 | + else: | ||
| 998 | + return type(item) | ||
| 999 | + | ||
| 1000 | + | ||
| 1001 | + def transform_hook( | ||
| 1002 | + self, | ||
| 1003 | + *args: Any, | ||
| 1004 | + **kwargs: Any | ||
| 1005 | + ) -> Tuple[List[Tuple], List[Dict]]: | ||
| 1006 | + # 获取 logger | ||
| 1007 | + logger = logging.getLogger(__name__) | ||
| 1008 | + # 预处理阶段优化:统一使用预定义函数或默认逻辑 | ||
| 1009 | + if self.transform_pre_fn: | ||
| 1010 | + inputs = self.transform_pre_fn(*args, **kwargs) | ||
| 1011 | + else: | ||
| 1012 | + inputs, indices, leaves, spec = self._process_inputs(args, kwargs) | ||
| 1013 | + | ||
| 1014 | + # 提取转换前的形状 (inputs 通常是 Tensor 列表) | ||
| 1015 | + if logger.isEnabledFor(logging.INFO): | ||
| 1016 | + pre_shapes = [self.get_shape_safe(t) for t in inputs] | ||
| 1017 | + logger.info("[Transform] Starting. Input tensors: %s, Shapes: %s", len(inputs), pre_shapes) | ||
| 1018 | + | ||
| 1019 | + # 执行核心转换操作 | ||
| 1020 | + trans_outputs = self.transform(tensors=inputs) | ||
| 1021 | + | ||
| 1022 | + # 提取转换后的形状 | ||
| 1023 | + if logger.isEnabledFor(logging.INFO): | ||
| 1024 | + post_shapes = [self.get_shape_safe(t) for t in trans_outputs] | ||
| 1025 | + logger.info("> Post-transform content: %s", post_shapes) | ||
| 1026 | + | ||
| 1027 | + # 后处理阶段优化:避免嵌套循环 | ||
| 1028 | + if self.transform_post_fn: | ||
| 1029 | + outputs = self.transform_post_fn(trans_outputs) | ||
| 1030 | + else: | ||
| 1031 | + outputs = self._recover_inputs(trans_outputs, indices, leaves, spec) | ||
| 1032 | + | ||
| 1033 | + if not outputs: | ||
| 1034 | + logger.error("CRITICAL: _recover_inputs returned NULL") | ||
| 1035 | + | ||
| 1036 | + return outputs | ||
| 1037 | + | ||
| 1038 | + def flatten_to_tensors(self, structure: Any) -> Tuple[List[torch.Tensor], List[int], List[Any], TreeSpec]: | ||
| 1039 | + leaves, spec = tree_flatten(structure) | ||
| 1040 | + indexed_tensors = [(i, leaf) for i, leaf in enumerate(leaves) if isinstance(leaf, torch.Tensor)] | ||
| 1041 | + indices = [] | ||
| 1042 | + tensors = [] | ||
| 1043 | + if indexed_tensors is not None and len(indexed_tensors) > 0: | ||
| 1044 | + indices, tensors = zip(*indexed_tensors) | ||
| 1045 | + return tensors, indices, leaves, spec | ||
| 1046 | + | ||
| 1047 | + def unflatten_from_tensors( | ||
| 1048 | + self, | ||
| 1049 | + tensors: List[torch.Tensor], | ||
| 1050 | + indices: List[int], | ||
| 1051 | + leaves: List[Any], | ||
| 1052 | + spec: TreeSpec | ||
| 1053 | + ) -> Any: | ||
| 1054 | + for idx, tensor in zip(indices, tensors): | ||
| 1055 | + leaves[idx] = tensor | ||
| 1056 | + return tree_unflatten(leaves, spec) | ||
| 1057 | + | ||
| 1058 | + def _process_inputs(self, args: Tuple, kwargs: dict) -> List[torch.Tensor]: | ||
| 1059 | + return self.flatten_to_tensors((args, kwargs)) | ||
| 1060 | + | ||
| 1061 | + def _recover_inputs( | ||
| 1062 | + self, | ||
| 1063 | + transform_res: List[List[torch.Tensor]], | ||
| 1064 | + indices: List[int], | ||
| 1065 | + leaves: List[Any], | ||
| 1066 | + spec: TreeSpec | ||
| 1067 | + ) -> Tuple[List[Tuple], List[Dict]]: | ||
| 1068 | + res = [] | ||
| 1069 | + for processd_tensors in transform_res: | ||
| 1070 | + res.append(self.unflatten_from_tensors(processd_tensors, indices, list(leaves), spec)) | ||
| 1071 | + return zip(*res) | ||
| 1072 | + | ||
| 1073 | + def _process_outputs( | ||
| 1074 | + self, | ||
| 1075 | + outputs_list: List[Any] | ||
| 1076 | + ) -> Tuple[List[List[torch.Tensor]], List[int], List[Any], TreeSpec]: | ||
| 1077 | + tensors_list = [] | ||
| 1078 | + leaves = [] | ||
| 1079 | + indices = [] | ||
| 1080 | + spec = None | ||
| 1081 | + for output in outputs_list: | ||
| 1082 | + tensors, indices, leaves, spec = self.flatten_to_tensors(output) | ||
| 1083 | + tensors_list.append(tensors) | ||
| 1084 | + return tensors_list, indices, leaves, spec | ||
| 1085 | + | ||
| 1086 | + def _recover_outputs( | ||
| 1087 | + self, | ||
| 1088 | + recover_res: List[torch.Tensor], | ||
| 1089 | + indices: List[int], | ||
| 1090 | + leaves: List[Any], | ||
| 1091 | + spec: TreeSpec | ||
| 1092 | + ) -> Any: | ||
| 1093 | + return self.unflatten_from_tensors(recover_res, indices, leaves, spec) | ||
| 1094 | + | ||
| 1095 | + def recover_hook( | ||
| 1096 | + self, | ||
| 1097 | + groups: List[Any] | ||
| 1098 | + ) -> Any: | ||
| 1099 | + """ | ||
| 1100 | + Process input groups through recovery pipeline. | ||
| 1101 | + | ||
| 1102 | + Args: | ||
| 1103 | + groups: List of input data to be processed. | ||
| 1104 | + | ||
| 1105 | + Returns: | ||
| 1106 | + Processed outputs after recovery and postprocessing. | ||
| 1107 | + """ | ||
| 1108 | + # 预处理:使用自定义函数或默认方法 | ||
| 1109 | + if self.recover_pre_fn: | ||
| 1110 | + inputs = self.recover_pre_fn(groups) | ||
| 1111 | + else: | ||
| 1112 | + inputs, indices, leaves, spec = self._process_outputs(groups) | ||
| 1113 | + | ||
| 1114 | + # 执行恢复操作 | ||
| 1115 | + re_outputs = self.recover(tensor_groups=inputs) | ||
| 1116 | + | ||
| 1117 | + # 后处理:使用自定义函数或默认方法 | ||
| 1118 | + if self.recover_post_fn: | ||
| 1119 | + outputs = self.recover_post_fn(re_outputs) | ||
| 1120 | + else: | ||
| 1121 | + outputs = self._recover_outputs(re_outputs, indices, leaves, spec) | ||
| 1122 | + | ||
| 1123 | + return outputs | ||
| 1124 | + | ||
| 1125 | + | ||
| 1126 | +def unified_copy(data: Any) -> Any: | ||
| 1127 | + """ | ||
| 1128 | + 对输入数据进行安全且统一的深拷贝。 | ||
| 1129 | + 支持PyTorch Tensor、字典、列表等常见数据类型。 | ||
| 1130 | + | ||
| 1131 | + Args: | ||
| 1132 | + data: 输入数据,可以是Tensor、dict、list等 | ||
| 1133 | + | ||
| 1134 | + Returns: | ||
| 1135 | + 数据的独立副本 | ||
| 1136 | + """ | ||
| 1137 | + if data is None: | ||
| 1138 | + return None | ||
| 1139 | + | ||
| 1140 | + # 处理PyTorch Tensor | ||
| 1141 | + if isinstance(data, torch.Tensor): | ||
| 1142 | + return data.clone().detach() | ||
| 1143 | + | ||
| 1144 | + # 处理字典类型 | ||
| 1145 | + elif isinstance(data, dict): | ||
| 1146 | + return {key: unified_copy(value) for key, value in data.items()} | ||
| 1147 | + | ||
| 1148 | + # 处理列表类型 | ||
| 1149 | + elif isinstance(data, list): | ||
| 1150 | + return [unified_copy(item) for item in data] | ||
| 1151 | + | ||
| 1152 | + # 处理元组类型 | ||
| 1153 | + elif isinstance(data, tuple): | ||
| 1154 | + return tuple(unified_copy(item) for item in data) | ||
| 1155 | + | ||
| 1156 | + else: | ||
| 1157 | + try: | ||
| 1158 | + return copy.deepcopy(data) | ||
| 1159 | + except (TypeError, ValueError): | ||
| 1160 | + return data | ||
| 1161 | + | ||
| 1162 | + | ||
| 1163 | +def _patch_dynamo_context(): | ||
| 1164 | + import inspect | ||
| 1165 | + from torch._dynamo.eval_frame import _TorchDynamoContext | ||
| 1166 | + from torch._dynamo.types import DynamoCallback | ||
| 1167 | + from torch._dynamo.convert_frame import CatchErrorsWrapper, ConvertFrame | ||
| 1168 | + from torch._dynamo.repro.after_dynamo import WrapBackendDebug | ||
| 1169 | + src_call = _TorchDynamoContext.__call__ | ||
| 1170 | + src_init = _TorchDynamoContext.__init__ | ||
| 1171 | + | ||
| 1172 | + def is_enable_shape_handling(callback: DynamoCallback, compiler_config=None): | ||
| 1173 | + """ | ||
| 1174 | + The shape handling feature is only available when enable_shape_handling is True and the backend is inductor | ||
| 1175 | + """ | ||
| 1176 | + if compiler_config is None or not compiler_config.get("enable_shape_handling", False): | ||
| 1177 | + return False | ||
| 1178 | + | ||
| 1179 | + if callback is None or not isinstance(callback, CatchErrorsWrapper): | ||
| 1180 | + return False | ||
| 1181 | + | ||
| 1182 | + convert_frame = getattr(callback, "_torchdynamo_orig_backend", None) | ||
| 1183 | + if not isinstance(convert_frame, ConvertFrame): | ||
| 1184 | + return False | ||
| 1185 | + | ||
| 1186 | + backend_debug = getattr(convert_frame, "_torchdynamo_orig_backend", None) | ||
| 1187 | + if not isinstance(backend_debug, WrapBackendDebug): | ||
| 1188 | + return False | ||
| 1189 | + | ||
| 1190 | + return getattr(backend_debug, "_compiler_name", None) == "inductor" | ||
| 1191 | + | ||
| 1192 | + def nothing(): | ||
| 1193 | + pass | ||
| 1194 | + | ||
| 1195 | + def new_init(self, callback: DynamoCallback, *args, **kwargs) -> None: | ||
| 1196 | + src_init(self, callback, *args, **kwargs) | ||
| 1197 | + compiler_config = kwargs.get("compiler_config") | ||
| 1198 | + if (is_enable_shape_handling(callback, compiler_config=compiler_config)): | ||
| 1199 | + trans_pre_fn = None | ||
| 1200 | + trans_post_fn = None | ||
| 1201 | + re_pre_fn = None | ||
| 1202 | + re_post_fn = None | ||
| 1203 | + function_dict = compiler_config.get("shape_handling_dict") | ||
| 1204 | + if function_dict is not None: | ||
| 1205 | + trans_pre_fn = function_dict.get("trans_pre_fn", None) | ||
| 1206 | + trans_post_fn = function_dict.get("trans_post_fn", None) | ||
| 1207 | + re_pre_fn = function_dict.get("re_pre_fn", None) | ||
| 1208 | + re_post_fn = function_dict.get("re_post_fn", None) | ||
| 1209 | + | ||
| 1210 | + self.shape_handling = NPUShapeHandling( | ||
| 1211 | + configs=compiler_config.get("shape_handling_configs"), | ||
| 1212 | + transform_pre_fn=trans_pre_fn, | ||
| 1213 | + transform_post_fn=trans_post_fn, | ||
| 1214 | + recover_pre_fn=re_pre_fn, | ||
| 1215 | + recover_post_fn=re_post_fn, | ||
| 1216 | + ) | ||
| 1217 | + | ||
| 1218 | + def new_call(self, fn): | ||
| 1219 | + src_fn = src_call(self, fn) | ||
| 1220 | + if isinstance(fn, torch.nn.Module) or inspect.isclass(fn): | ||
| 1221 | + return src_fn | ||
| 1222 | + | ||
| 1223 | + def new_fn(*args, **kwargs): | ||
| 1224 | + if (is_enable_shape_handling(self.callback, compiler_config=self.compiler_config)): | ||
| 1225 | + new_args, new_kwargs = self.shape_handling.transform_hook(*args, **kwargs) | ||
| 1226 | + args_is_split = len(args) != 0 and len(new_args) > 1 | ||
| 1227 | + kwargs_is_split = len(kwargs) != 0 and len(new_kwargs) > 1 | ||
| 1228 | + zipped_params = zip(new_args, new_kwargs) | ||
| 1229 | + res = [ | ||
| 1230 | + unified_copy(src_fn(*arg, **kwargs)) if args_is_split or kwargs_is_split | ||
| 1231 | + else src_fn(*arg, **kwargs) | ||
| 1232 | + for arg, kwargs in zipped_params | ||
| 1233 | + ] | ||
| 1234 | + return self.shape_handling.recover_hook(res) | ||
| 1235 | + return src_fn(*args, **kwargs) | ||
| 1236 | + return new_fn | ||
| 1237 | + _TorchDynamoContext.__call__ = new_call | ||
| 1238 | + _TorchDynamoContext.__init__ = new_init | ||
| 1239 | + | ||
| 1240 | + | ||
| 1241 | +def _patch_shape_handling(): | ||
| 1242 | + if getattr(_patch_shape_handling, "_is_patched", False): | ||
| 1243 | + return | ||
| 1244 | + _patch_dynamo_context() | ||
| 1245 | + _patch_shape_handling._is_patched = True | ||