( # noqa: C901
edge_compile_config: Optional[EdgeCompileConfig] = None,
class_only: bool = False,
core_aten_ops_exception_list: Optional[List[torch._ops.OpOverload]] = None,
preserve_ops: Optional[List[torch._ops.OpOverload]] = None,
)
| 240 | |
| 241 | |
| 242 | def EXIREdgeDialectVerifier( # noqa: C901 |
| 243 | edge_compile_config: Optional[EdgeCompileConfig] = None, |
| 244 | class_only: bool = False, |
| 245 | core_aten_ops_exception_list: Optional[List[torch._ops.OpOverload]] = None, |
| 246 | preserve_ops: Optional[List[torch._ops.OpOverload]] = None, |
| 247 | ): |
| 248 | _core_aten_ops_exception_list = core_aten_ops_exception_list or [] |
| 249 | _preserve_ops = preserve_ops or [] |
| 250 | # merge the exception list from edge_compile_config and exception_list |
| 251 | if edge_compile_config: |
| 252 | if edge_compile_config._core_aten_ops_exception_list: |
| 253 | _core_aten_ops_exception_list.extend( |
| 254 | edge_compile_config._core_aten_ops_exception_list |
| 255 | ) |
| 256 | if edge_compile_config.preserve_ops: |
| 257 | _preserve_ops.extend(edge_compile_config.preserve_ops) |
| 258 | |
| 259 | class _EXIREdgeDialectVerifier(Verifier): |
| 260 | dialect = "EDGE" |
| 261 | |
| 262 | def __init__(self) -> None: |
| 263 | _edge_compile_config = edge_compile_config or EdgeCompileConfig() |
| 264 | |
| 265 | self.enable = _edge_compile_config._check_ir_validity |
| 266 | self.check_edge_ops = _edge_compile_config._use_edge_ops |
| 267 | self.use_dim_order = not _edge_compile_config._skip_dim_order |
| 268 | |
| 269 | self._core_aten_ops_exception_list = _core_aten_ops_exception_list |
| 270 | self._preserve_ops = _preserve_ops |
| 271 | |
| 272 | self.aten_op_verifier = EXIRATenDialectVerifier( |
| 273 | core_aten_ops_exception_list=_core_aten_ops_exception_list, |
| 274 | preserve_ops=_preserve_ops, |
| 275 | ) |
| 276 | self.check_valid_aten_op = self.aten_op_verifier.check_valid_op |
| 277 | |
| 278 | if self.check_edge_ops: |
| 279 | self.check_valid_op = self.check_valid_edge_op |
| 280 | else: |
| 281 | self.check_valid_op = self.check_valid_aten_op |
| 282 | |
| 283 | def allowed_getattr_types(self) -> Tuple[Type[Any], ...]: |
| 284 | return ( |
| 285 | torch.fx.GraphModule, |
| 286 | LoweredBackendModule, |
| 287 | torch.Tensor, |
| 288 | torch.ScriptObject, |
| 289 | ) |
| 290 | |
| 291 | def allowed_op_types(self): |
| 292 | return super().allowed_op_types() + (EdgeOpOverload, types.FunctionType) |
| 293 | |
| 294 | def check_valid_edge_op(self, op): |
| 295 | if not self.enable: |
| 296 | return |
| 297 | if ( |
| 298 | op |
| 299 | in [operator.getitem] |
no outgoing calls