力争自助餐队代码评注 #11
|
|
@ -0,0 +1,265 @@
|
|||
# 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.
|
||||
# ============================================================================
|
||||
"""Python pass register"""
|
||||
from inspect import isfunction # 检查对象是否为函数
|
||||
from mindspore.graph_utils.graph_pattern import Pattern, NewParameter # 用于图模式匹配的类
|
||||
from mindspore._c_expression import PyPassManager_ # 用于管理优化传递的类
|
||||
|
||||
# __all__ 列表用于指定在使用 "from <module> import *" 时导出的符号
|
||||
# 它包含本脚本中定义的函数和类的名称,以便从外部访问
|
||||
|
||||
__all__ = [
|
||||
"register_pass", # 注册新的优化传递到传递管理器中的函数
|
||||
"unregister_pass", # 从传递管理器中注销优化传递的函数
|
||||
"gen_new_parameter", # 生成用于图转换的新参数的函数
|
||||
"cancel_new_parameter", # 取消生成新参数的函数
|
||||
"set_renorm", # 设置图转换的重新归一化标志的函数
|
||||
"set_reopt" # 设置图转换的重新优化标志的函数
|
||||
]
|
||||
|
||||
# PyPassManager类继承自PyPassManager_,用于注册和注销Python优化传递,以便在编译期间对图进行修改。
|
||||
class PyPassManager(PyPassManager_):
|
||||
r"""
|
||||
Used to register and unregister python passes which can be used to alter graphs.
|
||||
|
||||
Args:
|
||||
requires_grad(bool): Do automatic-differentiation after modified graph if true. Default: True
|
||||
run_only_once (bool): Specify whether or not to run pass only once. Default: False.
|
||||
|
||||
Raises:
|
||||
TypeError: If argument has invalid type.
|
||||
"""
|
||||
#创建一个PyPassManager对象,并用指定的参数requires_grad和run_only_once对其进行初始化。
|
||||
#PyPassManager是用于管理优化传递的类,它允许注册和注销Python优化传递函数,以便在编译期间对计算图进行修改。
|
||||
#通过创建PyPassManager对象,可以将Python优化传递函数注册到特定的编译阶段,并指定是否进行自动微分和是否只运行一次。
|
||||
def __init__(self, requires_grad=True, run_only_once=False):
|
||||
# 检查输入的参数run_only_once是否为bool类型,否则抛出TypeError异常
|
||||
if not isinstance(requires_grad, bool):
|
||||
raise TypeError(f"Expect bool, got : ({type(requires_grad)}){requires_grad}")
|
||||
# 检查输入的参数requires_grad是否为bool类型,否则抛出TypeError异常
|
||||
if not isinstance(run_only_once, bool):
|
||||
raise TypeError(f"Expect bool, got : ({type(run_only_once)}){run_only_once}")
|
||||
# 检查输入的参数run_only_once是否为bool类型,否则抛出TypeError异常
|
||||
self.requires_grad = requires_grad
|
||||
self.run_only_once_ = run_only_once
|
||||
# 将输入的参数requires_grad和run_only_once保存到实例变量中
|
||||
PyPassManager_.__init__(self)
|
||||
# 调用父类PyPassManager_的构造函数初始化PyPassManager_
|
||||
|
||||
# register函数的作用是将Python优化传递函数注册到优化传递管理器中,以便在编译期间对计算图进行修改和优化。
|
||||
|
||||
# 创建一个PyPassManager对象后,可以使用该对象调用register函数,将编写的Python优化传递函数注册到传递管理器中。
|
||||
# 注册后,该优化传递函数会在编译期间的特定阶段被自动调用,对计算图进行修改。
|
||||
def register(self, py_pass):
|
||||
if not isfunction(py_pass):
|
||||
raise TypeError(f"Expect function pass, got : ({type(py_pass)}){py_pass}")
|
||||
# 检查输入的py_pass是否为函数类型,如果不是,则抛出TypeError异常
|
||||
pattern, target = py_pass()
|
||||
# 调用py_pass函数,获取其返回的图模式匹配的模式和目标
|
||||
pass_name = py_pass.__name__
|
||||
# 获取py_pass函数的名称作为优化传递的名称
|
||||
if not isinstance(pattern, Pattern):
|
||||
raise TypeError(f"Expect pattern of Pattern type, got : ({type(pattern)}){pattern}")
|
||||
# 检查pattern是否为Pattern类型,如果不是,则抛出TypeError异常
|
||||
if not isinstance(target, Pattern):
|
||||
raise TypeError(f"Expect target of Pattern type, got : ({type(target)}){target}")
|
||||
# 检查target是否为Pattern类型,如果不是,则抛出TypeError异常
|
||||
super().register(pass_name, pattern, target, self.requires_grad, self.run_only_once_)
|
||||
# 调用父类PyPassManager_的register方法,将优化传递函数注册到传递管理器中
|
||||
# 参数包括优化传递名称、图模式匹配的模式、目标、是否进行自动微分和是否只运行一次的标志
|
||||
|
||||
#该unregister方法用于从传递管理器中注销已注册的优化传递函数。优化传递函数的注册是通过register方法完成的。
|
||||
def unregister(self, py_pass):
|
||||
#如果输入的py_pass参数是字符串类型,说明要注销一个已注册的优化传递函数
|
||||
if isinstance(py_pass, str):
|
||||
super().unregister(py_pass)
|
||||
return
|
||||
# 调用父类PyPassManager_的unregister方法,传递优化传递名称,从传递管理器中注销优化传递函数
|
||||
|
||||
# 如果输入的py_pass参数是函数类型,说明要注销一个已注册的优化传递函数
|
||||
if isfunction(py_pass):
|
||||
super().unregister(py_pass.__name__)
|
||||
return
|
||||
# 调用父类PyPassManager_的unregister方法,传递优化传递函数的名称,从传递管理器中注销优化传递函数
|
||||
|
||||
raise TypeError(f"Expect py_pass to be string or function, got ({type(py_pass)}){py_pass}")
|
||||
# 如果输入的py_pass参数既不是字符串也不是函数,则抛出TypeError异常
|
||||
|
||||
#__call__方法在Python类中被称为"调用"方法,它使得该类的实例可以像函数一样被调用。
|
||||
#在PyPassManager类中,__call__方法用于将输入的Python优化传递函数(py_pass)注册到传递管理器中,并在成功注册后返回该优化传递函数本身。
|
||||
def __call__(self, py_pass):
|
||||
self.register(py_pass)
|
||||
# 调用register方法,将输入的py_pass函数注册到传递管理器中
|
||||
return py_pass
|
||||
# 返回输入的py_pass函数本身
|
||||
|
||||
|
||||
#这个方法用于生成用于图转换的新参数。
|
||||
#首先,它检查传入的 pattern 参数是否是 NewParameter 类型的实例,如果不是,就会抛出异常。
|
||||
#然后,它调用父类的 gen_new_parameter 方法,完成新参数的生成过程。
|
||||
def gen_new_parameter(self, pattern):
|
||||
if not isinstance(pattern, NewParameter):
|
||||
raise TypeError(f"Expect pattern to be a NewParameter Pattern, got {pattern}")
|
||||
# 检查输入的pattern参数是否为NewParameter类型,如果不是,则抛出TypeError异常
|
||||
super().gen_new_parameter(pattern)
|
||||
# 调用父类PyPassManager_的gen_new_parameter方法,生成用于图转换的新参数
|
||||
|
||||
#这个方法用于设置一个参数 should_renorm,用于指示是否进行重新规范化。
|
||||
#首先,它检查传入的 should_renorm 参数是否是布尔类型,如果不是,就会抛出异常。
|
||||
#然后,它调用父类的 set_renorm 方法,将 should_renorm 参数传递给它,完成重新规范化设置的操作。
|
||||
def set_renorm(self, should_renorm):
|
||||
if not isinstance(should_renorm, bool):
|
||||
raise TypeError(f"Expect should_renorm to be a bool, got {should_renorm}")
|
||||
# 检查输入的should_renorm参数是否为bool类型,如果不是,则抛出TypeError异常
|
||||
super().set_renorm(should_renorm)
|
||||
# 调用父类的set_renorm方法,并将should_renorm参数传递给它
|
||||
|
||||
#这个方法用于设置一个参数 do_reopt,用于指示是否进行重新优化。
|
||||
#首先,它检查传入的 do_reopt 参数是否是布尔类型,如果不是,就会抛出异常。
|
||||
#然后,它调用父类的 set_reopt 方法,将 do_reopt 参数传递给它,完成重新优化设置的操作。
|
||||
|
||||
def set_reopt(self, do_reopt):
|
||||
if not isinstance(do_reopt, bool):
|
||||
raise TypeError(f"Expect do_reopt to be a bool, got {do_reopt}")
|
||||
# 检查输入的do_reopt参数是否为bool类型,如果不是,则抛出TypeError异常
|
||||
super().set_reopt(do_reopt)
|
||||
# 调用父类的set_reopt方法,并将do_reopt参数传递给它
|
||||
|
||||
def register_pass(requires_grad=True, run_only_once=False):
|
||||
"""
|
||||
Register python pass to specified pipeline phase which would be used in compilation.
|
||||
将python传递注册到指定的管道阶段,该阶段将在编译中使用。
|
||||
|
||||
Args: 参数:
|
||||
requires_grad(bool): Do automatic-differentiation after modified graph if true. Default: True.
|
||||
如果修改后的图为true,则执行自动微分。默认值:True
|
||||
run_only_once(bool): Run this pass only once if set true. Otherwise run the pass until converge. Default:
|
||||
False.
|
||||
如果设置为true,则只运行一次。否则,运行通道,直到收敛。默认值:False
|
||||
|
||||
Returns:
|
||||
This function should be used as a decorator, return the decoratorated pass function.
|
||||
这个函数应该被用作装饰器,返回经过装饰的传递函数。
|
||||
|
||||
Examples:
|
||||
>>> from mindspore.graph_utils.graph_pattern import Call, Any
|
||||
>>> from mindspore.ops import operations as P
|
||||
>>> @register_pass()
|
||||
>>> def toy_pass():
|
||||
>>> x = Any()
|
||||
>>> pattern = Call(P.Softmax(), [x])
|
||||
>>> target = Call(P.ReLU(), [x])
|
||||
>>> return pattern, target
|
||||
"""
|
||||
return PyPassManager(requires_grad, run_only_once)
|
||||
|
||||
|
||||
def unregister_pass(py_pass):
|
||||
"""
|
||||
Unregister python pass.
|
||||
注销python路径
|
||||
|
||||
Args:
|
||||
py_pass(Union(str, function)): target python pass to unregister.
|
||||
注销指定的python路径
|
||||
"""
|
||||
ppm = PyPassManager()
|
||||
ppm.unregister(py_pass)
|
||||
|
||||
|
||||
def gen_new_parameter(pattern):
|
||||
"""
|
||||
Generate specified parameter every time a network gets compiled.
|
||||
每次编译网络时生成指定的参数。
|
||||
|
||||
NOTE:
|
||||
In this way, every pass uses this pattern would be using the same Parameter. If use NewParameter without
|
||||
gen_new_parameter, every pass match would build a new Parameter.
|
||||
This would register a pass to add new parameter in the compilation pipeline, so later compilation would
|
||||
ALSO add this parameter unless the pass is unregistered. To unregister this pass, call
|
||||
cancel_new_parameter(pattern)
|
||||
这样,每次使用此模式的pass都将使用相同的Parameter。如使用gen_new_parameter之外的NewParameter ,
|
||||
每次通过匹配都会构建一个新的Parameter。
|
||||
这将注册一个在编译管道中添加新参数的pass,因此以后的编译将:除非pass未注册,否则也添加此参数。
|
||||
要注销此pass,请调用cancel_new_parameter(模式)
|
||||
|
||||
Args:
|
||||
pattern (NewParameter): NewParameter type, could be used to build nested patterns across multiple passes
|
||||
after gen_new_parameter.
|
||||
NewParameter类型,可用于在gen_new_parameter之后跨多个pass构建嵌套模式。
|
||||
|
||||
Raises:
|
||||
TypeError: If argument has invalid type.
|
||||
参数的类型无效
|
||||
|
||||
Examples:
|
||||
>>> from mindspore.graph_utils.graph_pattern import NewParameter
|
||||
>>> abc = NewParameter("abc")
|
||||
>>> gen_new_parameter(abc)
|
||||
"""
|
||||
ppm = PyPassManager()
|
||||
ppm.gen_new_parameter(pattern)
|
||||
|
||||
|
||||
def cancel_new_parameter(pattern):
|
||||
"""
|
||||
Use with gen_new_parameter to unregister gen_new_parameter pass.
|
||||
|
||||
Args:
|
||||
pattern (NewParameter): NewParameter type, cancel the pass which would add new parameter as this pattern
|
||||
describes.
|
||||
NewParameter类型,取消将添加新参数的传递,如该模式所描述的。
|
||||
Examples:
|
||||
>>> from mindspore.graph_utils.graph_pattern import NewParameter
|
||||
>>> abc = NewParameter("abc")
|
||||
>>> gen_new_parameter(abs)
|
||||
>>> # some compilations
|
||||
>>> cancel_new_parameter(abc)
|
||||
"""
|
||||
if not isinstance(pattern, NewParameter):
|
||||
raise TypeError(f"Expect pattern to be a NewParameter Pattern, got {pattern}")
|
||||
ppm = PyPassManager()
|
||||
ppm.unregister(pattern.para_name)
|
||||
|
||||
|
||||
def set_renorm(should_renorm):
|
||||
"""
|
||||
Set whether or not to do renormalization after modified graph in python pass(es).
|
||||
|
||||
Args:
|
||||
should_renorm(bool): whether or not to do renormalization after modified graph in python pass(es).
|
||||
|
||||
NOTE:
|
||||
This interface is mainly intended for testing modifying graph without worrying about its validity. Turn off
|
||||
renormalization may BREAK the network.
|
||||
"""
|
||||
ppm = PyPassManager()
|
||||
ppm.set_renorm(should_renorm)
|
||||
|
||||
|
||||
def set_reopt(do_reopt):
|
||||
"""
|
||||
Set whether or not to do optimization after modified graph in python pass(es).
|
||||
|
||||
Args:
|
||||
do_reopt(bool): whether or not to do optimization after modified graph in python pass(es).
|
||||
在python的pass中修改图形后是否进行重新规范化。
|
||||
NOTE:
|
||||
This interface is mainly intended for testing modifying graph without worrying about its validity. Turn off
|
||||
renormalization may BREAK the network.
|
||||
该接口主要用于测试修改图,而不用担心修改图的有效性。关闭重整化可能会破坏网络。
|
||||
"""
|
||||
ppm = PyPassManager()
|
||||
ppm.set_reopt(do_reopt)
|
||||
Loading…
Reference in New Issue