Coverage for /opt/cloud/slavespace/usr1/096471637100f3de0fcfc01072822a80/dttest/api/python/ge/ge/passes/registry.py: 94%
90 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-07-27 10:02 +0800
« prev ^ index » next coverage.py v7.15.2, created at 2026-07-27 10:02 +0800
1#!/usr/bin/env python3
2# -*- coding: utf-8 -*-
3# -----------------------------------------------------------------------------------------------------------
4# Copyright (c) 2026 Huawei Technologies Co., Ltd.
5# This program is free software, you can redistribute it and/or modify it under the terms and conditions of
6# CANN Open Software License Agreement Version 2.0 (the "License").
7# Please refer to the License for details. You may not use this file except in compliance with the License.
8# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
9# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
10# See LICENSE in the root of the software repository for the full text of the License.
11# -----------------------------------------------------------------------------------------------------------
13"""Python pass registry and decorators."""
15import inspect
16import threading
17from collections.abc import Iterable as IterableABC
18from dataclasses import dataclass, field
19from typing import Dict, Iterable, List, Optional, Type
21from .base import DecomposePass, FusionBasePass, PassStage, PatternFusionPass
23PASS_KIND_FUSION_BASE = "fusion_base"
24PASS_KIND_PATTERN_FUSION = "pattern_fusion"
25PASS_KIND_DECOMPOSE = "decompose"
28@dataclass(frozen=True)
29class PassDescriptor:
30 """Normalized Python pass descriptor."""
32 descriptor_key: str
33 pass_name: str
34 module_name: str
35 class_name: str
36 stage: PassStage
37 kind: str
38 cls: Type[FusionBasePass]
39 op_types: List[str] = field(default_factory=list)
41 def to_bridge_dict(self) -> dict:
42 return {
43 "descriptor_key": self.descriptor_key,
44 "pass_name": self.pass_name,
45 "module_name": self.module_name,
46 "class_name": self.class_name,
47 "stage": self.stage.value,
48 "kind": self.kind,
49 "op_types": list(self.op_types),
50 }
53class _PassRegistry:
54 def __init__(self) -> None:
55 self._lock = threading.RLock()
56 self._descriptor_key_to_desc: Dict[str, PassDescriptor] = {}
57 self._pass_name_to_desc: Dict[str, PassDescriptor] = {}
59 def clear(self) -> None:
60 with self._lock:
61 self._descriptor_key_to_desc.clear()
62 self._pass_name_to_desc.clear()
64 def register(self, descriptor: PassDescriptor) -> PassDescriptor:
65 with self._lock:
66 if descriptor.descriptor_key in self._descriptor_key_to_desc:
67 raise ValueError(f"python pass descriptor_key already exists: {descriptor.descriptor_key}")
68 if descriptor.pass_name in self._pass_name_to_desc:
69 raise ValueError(f"python pass name already exists: {descriptor.pass_name}")
70 self._descriptor_key_to_desc[descriptor.descriptor_key] = descriptor
71 self._pass_name_to_desc[descriptor.pass_name] = descriptor
72 return descriptor
74 def get_by_descriptor_key(self, descriptor_key: str) -> Optional[PassDescriptor]:
75 with self._lock:
76 return self._descriptor_key_to_desc.get(descriptor_key)
78 def get_all(self) -> List[PassDescriptor]:
79 with self._lock:
80 return list(self._descriptor_key_to_desc.values())
83_PASS_REGISTRY = _PassRegistry()
86def _build_descriptor_key(module_name: str, class_name: str, pass_name: str) -> str:
87 return f"{module_name}:{class_name}:{pass_name}"
90def _normalize_decompose_op_types(op_types: Iterable[str]) -> List[str]:
91 if isinstance(op_types, (str, bytes)) or not isinstance(op_types, IterableABC):
92 raise TypeError("register_decompose_pass op_types must be an iterable of strings")
94 normalized_op_types = list(op_types)
95 if not normalized_op_types:
96 raise ValueError("register_decompose_pass requires at least one op type")
98 for op_type in normalized_op_types:
99 if not isinstance(op_type, str) or not op_type:
100 raise TypeError("register_decompose_pass op_types must contain non-empty strings")
101 return normalized_op_types
104def _register_pass_class(
105 cls: Type[FusionBasePass],
106 *,
107 kind: str,
108 name: str,
109 stage: PassStage,
110 op_types: Optional[Iterable[str]] = None,
111) -> Type[FusionBasePass]:
112 module_name = cls.__module__
113 class_name = cls.__name__
114 descriptor = PassDescriptor(
115 descriptor_key=_build_descriptor_key(module_name, class_name, name),
116 pass_name=name,
117 module_name=module_name,
118 class_name=class_name,
119 stage=stage,
120 kind=kind,
121 cls=cls,
122 op_types=list(op_types or []),
123 )
124 _PASS_REGISTRY.register(descriptor)
125 setattr(cls, "__ge_pass_descriptor__", descriptor)
126 return cls
129def register_fusion_pass(*, name: str, stage: PassStage, kind: Optional[str] = None) -> callable:
130 """Decorator for FusionBasePass and PatternFusionPass."""
132 def decorator(cls: Type[FusionBasePass]) -> Type[FusionBasePass]:
133 if not inspect.isclass(cls) or not issubclass(cls, FusionBasePass):
134 raise TypeError("register_fusion_pass expects a FusionBasePass subclass")
135 pass_kind = kind
136 if pass_kind is None:
137 pass_kind = PASS_KIND_PATTERN_FUSION if issubclass(cls, PatternFusionPass) else PASS_KIND_FUSION_BASE
138 return _register_pass_class(cls, kind=pass_kind, name=name, stage=stage)
140 return decorator
143def register_decompose_pass(*, name: str, stage: PassStage, op_types: Iterable[str]) -> callable:
144 """Decorator for DecomposePass."""
146 normalized_op_types = _normalize_decompose_op_types(op_types)
148 def decorator(cls: Type[DecomposePass]) -> Type[DecomposePass]:
149 if not inspect.isclass(cls) or not issubclass(cls, DecomposePass):
150 raise TypeError("register_decompose_pass expects a DecomposePass subclass")
151 setattr(cls, "op_types", list(normalized_op_types))
152 return _register_pass_class(
153 cls,
154 kind=PASS_KIND_DECOMPOSE,
155 name=name,
156 stage=stage,
157 op_types=normalized_op_types,
158 )
160 return decorator
163def clear_registered_passes() -> None:
164 _PASS_REGISTRY.clear()
167def get_registered_passes() -> List[PassDescriptor]:
168 return _PASS_REGISTRY.get_all()
171def get_registered_pass_dicts() -> List[dict]:
172 return [item.to_bridge_dict() for item in get_registered_passes()]
175def get_registered_pass_by_descriptor_key(
176 descriptor_key: str,
177) -> Optional[PassDescriptor]:
178 return _PASS_REGISTRY.get_by_descriptor_key(descriptor_key)