已合并
[master]refactor: move shape handling into dynamo utils (#4241) #45007
黄桂军创建于 8月20日
[master]refactor: move shape handling into dynamo utils (#4241) #45007
已合并
黄桂军创建于 8月20日
共 4 个文件变更+549-511
@@ -23,6 +23,7 @@ device = "npu"
23def model_fn(A, B):23def model_fn(A, B):
24 return A + B24 return A + B
25 25 
26+ 
26shape_options = {27shape_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 os466 import os
444- import types
445 from unittest import mock467 from unittest import mock
446 468 
447 import torch469 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_register490 _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 True499 assert wrapper.config["enable_shape_handling"] is True
480- assert events == [("scope_register", "mlir")], events500+ 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, Tuple6+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 
397def patch_dynamo_context():9def patch_dynamo_context():
398- import contextlib10+ if not getattr(_impl._patch_shape_handling, "_is_patched", False):
399- import inspect11+ _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 
477def patch_shape_handling():14def patch_shape_handling():
@@ -479,3 +16,4 @@ def patch_shape_handling():
479 return16 return
480 patch_dynamo_context()17 patch_dynamo_context()
481 patch_shape_handling._is_patched = True18 patch_shape_handling._is_patched = True
19+ _impl._patch_shape_handling._is_patched = True
@@ -1,15 +1,19 @@
1import importlib1import importlib
2import importlib.abc2import importlib.abc
3+import copy
3import functools4import functools
4import inspect5import inspect
5import logging6import logging
6import os7import os
7import sys8import sys
8import threading9import threading
9-from typing import Any, Optional, TYPE_CHECKING10+import warnings
11+from typing import Any, Callable, Dict, List, Optional, Tuple, TYPE_CHECKING
10 12 
11import torch13import torch
12import torch_npu14import torch_npu
15+import torch_npu._C
16+from torch.utils._pytree import TreeSpec, tree_flatten, tree_unflatten
13from torch import _TorchCompileWrapper17from 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 = True216 self._npu_shape_handling_requested = True
213- # Shape handling is installed in new_call, after the selected217+ 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_handling264 del self._npu_defer_shape_handling
260 del self._npu_shape_handling_requested265 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