mindspore代码评注-等风也等你 #28
|
|
@ -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_):
|
||||
|
|
|
|||
Loading…
Reference in New Issue