mindspore代码评注-等风也等你 #28

Open
darrenlu wants to merge 8 commits from darrenlu/mindspore2022:master into master
1 changed files with 36 additions and 12 deletions
Showing only changes of commit 7c1c7c665f - Show all commits

View File

@ -434,6 +434,15 @@ class _Grad(GradOperation_):
def __init__(self, get_by_list=False, sens_param=False, get_by_position=False):
"""Initialize _Grad."""
"""
参数:
- get_by_list: 是否按列表获取梯度默认为 False
- sens_param: 是否使用灵敏度参数默认为 False
- get_by_position: 是否按位置获取梯度默认为 False
- has_aux: 是否包含辅助输出默认为 False
- get_value: 是否获取值默认为 False
- return_ids: 是否返回标识符默认为 False
"""
if not isinstance(get_by_position, bool):
raise TypeError(f"For '_Grad', the 'get_by_position' should be bool, "
f"but got {type(get_by_position).__name__}")
@ -453,6 +462,14 @@ class _Grad(GradOperation_):
self.grad_position = None
def __call__(self, fn, weights=None, grad_position=0):
"""
调用 _Grad 类实例以生成梯度函数
参数
- fn: 输入的函数
- weights: 权重参数默认为 None
- grad_position: 梯度位置默认为 0
"""
if self.grad_fn is not None and self.fn == fn and self.grad_position == grad_position:
return self.grad_fn
grad_ = _Grad(self.get_by_list, self.sens_param, self.get_by_position)
@ -507,25 +524,32 @@ class _Grad(GradOperation_):
def _pynative_forward_run(self, grad, args, kwargs, fn):
""" Pynative forward runs to build grad graph. """
new_kwargs = kwargs
new_kwargs = kwargs # 将传入的关键字参数 kwargs 赋值给新变量 new_kwargs
if self.sens_param:
# 如果 self.sens_param 为 True表示函数需要灵敏度参数
if 'sens' in kwargs.keys():
new_kwargs = kwargs.copy()
new_kwargs.pop('sens')
# 如果 'sens' 存在于 kwargs 的键中
new_kwargs = kwargs.copy() # 复制 kwargs 到新的字典 new_kwargs
new_kwargs.pop('sens') # 从 new_kwargs 中删除 'sens' 键
else:
args = args[:-1]
args = args[:-1] # 否则,移除参数列表 args 的最后一个参数
if isinstance(fn, FunctionType):
# 如果 fn 是 FunctionType 类型的对象(通常表示普通函数或方法)
if not _pynative_executor.check_run(grad, fn, *args, **new_kwargs):
_pynative_executor.set_grad_flag(True)
_pynative_executor.new_graph(fn, *args, **new_kwargs)
outputs = fn(*args, **new_kwargs)
_pynative_executor.end_graph(fn, outputs, *args, **new_kwargs)
# 如果还没有运行过与 grad 关联的 fn 函数
_pynative_executor.set_grad_flag(True) # 设置梯度标志为 True表示进入梯度计算模式
_pynative_executor.new_graph(fn, *args, **new_kwargs) # 创建新的计算图
outputs = fn(*args, **new_kwargs) # 执行 fn 函数并获取其输出
_pynative_executor.end_graph(fn, outputs, *args, **new_kwargs) # 结束计算图的构建
else:
# Check if fn has run already.
# 如果 fn 不是 FunctionType 类型的对象(可能是 nn.Cell 或其他类型的对象)
# 检查 fn 是否已经运行过,如果没有,则设置梯度标志为 True运行 fn然后将梯度标志设置为 False
if not _pynative_executor.check_run(grad, fn, *args, **new_kwargs):
fn.set_grad()
fn(*args, **new_kwargs)
fn.set_grad(False)
fn.set_grad() # 设置 fn 为梯度计算模式
fn(*args, **new_kwargs) # 运行 fn 函数
fn.set_grad(False) # 将 fn 设置为非梯度计算模式
class _Vmap(VmapOperation_):