diff --git a/graph_pattern.py b/graph_pattern.py new file mode 100644 index 0000000..5e09fe4 --- /dev/null +++ b/graph_pattern.py @@ -0,0 +1,249 @@ +# Copyright 2020 Huawei Technologies Co., Ltd +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================ +"""Patterns for describing graphs""" +#描述图形的模式 +from mindspore.ops import Primitive +from mindspore.common.tensor import Tensor +from mindspore._c_expression import Pattern, OneOf_, Prim_, Call_, NoneOf_, Any, NewTensor_, NewParameter_, Imm + +__all__ = [ + "OneOf", + "Prim", + "Call", + "NoneOf", + "Any", + "NewTensor", + "NewParameter", + "Imm" +] + + +class OneOf(OneOf_): + r""" + Express a pattern which allows a list of patterns. + 表达一个允许模式列表的模式 + """ + def __init__(self, patterns=None): + r""" + Args: + patterns(Union[:class:`mindspore.graph_utils.graph_pattern`, + tuple[:class:`mindspore.graph_utils.graph_pattern`], + list[:class:`mindspore.graph_utils.graph_pattern`]]): list of allowed patterns, + each element should be one of the exposed Pattern instance. + + Raises: + TypeError: raise type error for invalid inputs. + """ + self.patterns = patterns + if isinstance(patterns, Pattern): + OneOf_.__init__(self, [patterns]) + #如果patterns是mindspore.graph_utils.graph_pattern类的实例,则使用包含patterns的单一元素列表初始化基类OneOf_。 + elif isinstance(patterns, (tuple, list)) and all(isinstance(pattern, Pattern) for pattern in patterns): + OneOf_.__init__(self, patterns) + #如果patterns是元组或列表,并且其中的所有元素都是mindspore.graph_utils.graph_pattern类的实例, + #则使用包含patterns中模式的列表初始化基类OneOf_ + else: + raise TypeError(f"Expect patterns to be a list of Patterns/Pattern, got : {patterns}") + #如果patterns是其他类型的对象,或者其中包含不是mindspore.graph_utils.graph_pattern类的实例, + #则抛出TypeError并附带相应的错误消息。 + + +class Prim(Prim_): + r""" + Express a pattern of certain primitive type(s). + 表示某种基本类型的模式 + + NOTE: + This pattern will match and only match the primitive value node. If matching primitive CNode is needed, + please refer to CallWith pattern. + """ + def __init__(self, types, name=None): + r""" + Args: + 支持三种不同的types参数输入方式: + types (Union[str, :class:`mindspore.ops.Primitive`, list[:class:`mindspore.ops.Primitive`], + tuple[:class:`mindspore.ops.Primitive`]): + Specify allowed types. + If it is a string, the form could be + 1) a single primitive type, e.g. 'Conv2D'若types是一个字符串,则可以是单个基本类型,例如 'Conv2D' + 2) a set of primitive types separated by '|', e.g. 'MatMul|Conv2D' + 多个基本类型,用|符号分隔,例如 'MatMul|Conv2D' + It can also be a Primitive or a list/tuple of Primitives, e.g. [ops.Conv2D(1, 6)] + 如果types是一个Primitive对象,则表示仅匹配该具体的基本类型。 + 如果types是一个列表或元组,并且列表中的元素都是Primitive对象,则表示匹配多个基本类型。 + name (str): name of the pattern, optional. Default: None. + + Raises: + TypeError: raise type error for invalid argument. + """ + # 如果提供了 name 参数,但不是字符串类型,则抛出 TypeError + if name is not None and not isinstance(name, str): + raise TypeError(f"Expect string, got : {name}") + # 将参数 name 赋值给对象的 name 属性 + self.name = name + # 根据 types 参数的类型进行不同的处理 + if isinstance(types, str): + if self.name is None: + self.name = types + self.types = types.split('|') + # 如果 types 是一个字符串,则将其按照 '|' 符号拆分,并存储在 self.types 中 + elif isinstance(types, Primitive): + # 如果 types 是一个 Primitive 对象,则表示只允许匹配这个具体的基本类型 + if self.name is None: + self.name = types.name + self.types = [types] + elif isinstance(types, (tuple, list)) and all(isinstance(tp, Primitive) for tp in types): + # 如果 types 是一个包含 Primitive 对象的列表或元组,则表示允许匹配多个基本类型 + if self.name is None: + self.name = "" + for prim in types: + self.name += prim.name + self.types = types + else: + # 如果 types 参数不是允许的类型,则抛出 TypeError + raise TypeError(f"Expecting a primitive type string or a list of Primitives, got : {types}") + # 调用父类 Prim_ 的初始化方法,将处理好的 types 和 name 作为参数传递给父类 + Prim_.__init__(self, self.types, self.name) + + +class Call(Call_): + #初始化对象的属性 + r""" + Express a primitive CNode. + """ + def __init__(self, prim_pattern, inputs=None): + r""" + Args: + prim_pattern (Union[str, :class:`mindspore.graph_utils.graph_pattern.IsPrimTypeOf`, + :class:`mindspore.ops.Primitive`]): Primitive ValueNode in the Primitive CNode. + inputs (Union[list[:class:`mindspore.graph_utils.graph_pattern`], + tuple[:class:`mindspore.graph_utils.graph_pattern`]]): + Specify inputs pattern for the primitive(s), optional. If None, accepts any inputs; if specified, input + patterns should be of right order and each element should be one of the exposed Pattern instance. + 模式的顺序应该正确,每个元素应该是公开的Pattern实例之一 + Raises: + TypeError: raise type error for invalid argument. + """ + #检查prim_pattern的类型是否是Pattern、Primitive或字符串类型。 + #如果不是这些类型之一,则抛出TypeError,表示期望prim_pattern是Pattern、Primitive或字符串类型。 + if not isinstance(prim_pattern, (Pattern, str, Primitive)): + raise TypeError(f"Expect prim_pattern to be Pattern, Primitive or string, got : {prim_pattern}") + #将传入的prim_pattern赋值给实例变量self.prim_pattern,以保存原语模式(Primitive Pattern)或原语名称(Primitive name) + self.prim_pattern = prim_pattern + #将self.inputs初始化为空列表 + self.inputs = [] + #None:什么都不做 + if inputs is None: + pass + #检查inputs是否为Pattern的列表或元组,且其中的所有元素都是Pattern类型。 + elif isinstance(inputs, (tuple, list)) and all(isinstance(input, Pattern) for input in inputs): + self.inputs = inputs + #如果inputs不满足上述条件,抛出TypeError,表示期望inputs是Pattern的列表。 + else: + raise TypeError(f"Expect inputs to be a list of Patterns, got : {inputs}") + Call_.__init__(self, self.prim_pattern, self.inputs) + #调用基类Call_的构造函数来完成初始化,传入self.prim_pattern作为原语模式或名称, + #以及self.inputs作为输入模式列表。 + +class NoneOf(NoneOf_): + r""" + Express a pattern which forbids a list of patterns. + 表达一个禁止模式列表的模式。 + NOTE: + NoneOf pattern should not be the root pattern. + #NoneOf模式不是根模式 + """ + def __init__(self, patterns=None): + r""" + Args: + patterns(Union[list[:class:`mindspore.graph_utils.graph_pattern`]]: list of forbidden patterns, each + element should be one of the exposed Pattern instance. + 禁用模式列表,每个元素应该是公开的Pattern实例之一。 + Raises: + TypeError: raise type error for invalid argument. + """ + self.patterns = patterns + if patterns is None: + NoneOf_.__init__(self, ()) + #如果 patterns 参数为 None,即未提供 patterns,则将 None 传递给 NoneOf_ 类的初始化方法。 + elif isinstance(patterns, Pattern): + NoneOf_.__init__(self, [patterns]) + #如果 patterns 参数是一个单独的 Pattern 对象,则将包含这个 Pattern 对象的列表传递给 NoneOf_ 类的初始化方法。 + elif isinstance(patterns, (tuple, list)) and all(isinstance(pattern, Pattern) for pattern in patterns): + NoneOf_.__init__(self, patterns) + #如果 patterns 参数是一个列表或元组,并且列表中的元素都是 Pattern 对象, + #则将整个列表传递给 NoneOf_ 类的初始化方法。 + else: + raise TypeError(f"Expect list of Patterns/Pattern, got : {patterns}") + + +class NewTensor(NewTensor_): + r""" + New Tensor to be used in the target. + 要在目标中使用的新张量 + """ + def __init__(self, input_tensor): + r""" + Args: + input_tensor(:class:`mindspore.common.tensor.Tensor`): new tensor to be used in the target. + 要在目标中使用的新张量 + Raises: + TypeError: raise type error for invalid argument. + """ + self.input_tensor = input_tensor + # 将传入的 input_tensor 参数赋值给对象的 input_tensor 属性 + if isinstance(input_tensor, Tensor): + # 检查 input_tensor 参数的类型 + NewTensor_.__init__(self, input_tensor) + # 如果 input_tensor 是一个 Tensor 对象,则调用 NewTensor_ 类的初始化方法,并传入 input_tensor 作为参数 + else: + raise TypeError(f"Expect input_tensor to be a Tensor, got : {input_tensor}") + # 如果 input_tensor 参数不是 Tensor 对象,则抛出 TypeError,指示输入参数类型错误 + + +class NewParameter(NewParameter_): + r""" + New Parameter to be used in the target. + """ + def __init__(self, para_name, default_tensor, requires_grad=False, layerwise_parallel=False): + r""" + Args: + para_name(str): name for the new Parameter. + default_tensor(:class:`mindspore.common.tensor.Tensor`): default value for the new Parameter. + requires_grad(bool): True if the parameter requires gradient. Default: True. + 如果参数需要梯度,则为True + layerwise_parallel(bool): switch for layerwise parallel mode. Default: False. + Layerwise_parallel (bool):分层并行模式开关。 + + Raises: + TypeError: raise type error for invalid argument. + """ + # 将传入的 para_name、default_tensor、requires_grad 和 layerwise_parallel 参数赋值给对象的相应属性 + self.para_name = para_name + self.default_tensor = default_tensor + self.requires_grad = requires_grad + self.layerwise_parallel = layerwise_parallel + # 检查传入的参数类型是否正确,并根据结果选择初始化父类 NewParameter_ + if isinstance(para_name, str) and isinstance(default_tensor, Tensor) and isinstance(requires_grad, bool) and\ + isinstance(layerwise_parallel, bool): + # 如果 para_name 是一个字符串,default_tensor 是一个 Tensor 对象,requires_grad 和 layerwise_parallel 都是布尔值, + # 则调用 NewParameter_ 类的初始化方法,并传入 para_name、default_tensor、requires_grad 和 layerwise_parallel 作为参数。 + NewParameter_.__init__(self, self.para_name, self.default_tensor, self.requires_grad, + self.layerwise_parallel) + else: + # 如果有任何一个参数类型不正确,则抛出 TypeError,指示输入参数类型错误。 + raise TypeError(f"Expect para_name(str), default_tensor(Tensor), requires_grad(bool), \ + layerwise_parallel(bool), got : {para_name}, {default_tensor}, \ + {requires_grad}, {layerwise_parallel}")