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

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# ----------------------------------------------------------------------------------------------------------- 

12 

13"""Python pass registry and decorators.""" 

14 

15import inspect 

16import threading 

17from collections.abc import Iterable as IterableABC 

18from dataclasses import dataclass, field 

19from typing import Dict, Iterable, List, Optional, Type 

20 

21from .base import DecomposePass, FusionBasePass, PassStage, PatternFusionPass 

22 

23PASS_KIND_FUSION_BASE = "fusion_base" 

24PASS_KIND_PATTERN_FUSION = "pattern_fusion" 

25PASS_KIND_DECOMPOSE = "decompose" 

26 

27 

28@dataclass(frozen=True) 

29class PassDescriptor: 

30 """Normalized Python pass descriptor.""" 

31 

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) 

40 

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 } 

51 

52 

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] = {} 

58 

59 def clear(self) -> None: 

60 with self._lock: 

61 self._descriptor_key_to_desc.clear() 

62 self._pass_name_to_desc.clear() 

63 

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 

73 

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) 

77 

78 def get_all(self) -> List[PassDescriptor]: 

79 with self._lock: 

80 return list(self._descriptor_key_to_desc.values()) 

81 

82 

83_PASS_REGISTRY = _PassRegistry() 

84 

85 

86def _build_descriptor_key(module_name: str, class_name: str, pass_name: str) -> str: 

87 return f"{module_name}:{class_name}:{pass_name}" 

88 

89 

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") 

93 

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") 

97 

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 

102 

103 

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 

127 

128 

129def register_fusion_pass(*, name: str, stage: PassStage, kind: Optional[str] = None) -> callable: 

130 """Decorator for FusionBasePass and PatternFusionPass.""" 

131 

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) 

139 

140 return decorator 

141 

142 

143def register_decompose_pass(*, name: str, stage: PassStage, op_types: Iterable[str]) -> callable: 

144 """Decorator for DecomposePass.""" 

145 

146 normalized_op_types = _normalize_decompose_op_types(op_types) 

147 

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 ) 

159 

160 return decorator 

161 

162 

163def clear_registered_passes() -> None: 

164 _PASS_REGISTRY.clear() 

165 

166 

167def get_registered_passes() -> List[PassDescriptor]: 

168 return _PASS_REGISTRY.get_all() 

169 

170 

171def get_registered_pass_dicts() -> List[dict]: 

172 return [item.to_bridge_dict() for item in get_registered_passes()] 

173 

174 

175def get_registered_pass_by_descriptor_key( 

176 descriptor_key: str, 

177) -> Optional[PassDescriptor]: 

178 return _PASS_REGISTRY.get_by_descriptor_key(descriptor_key)