mindspore2022/mindspore/ccsrc/transform-update/python_pass_register.py

266 lines
12 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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