MCPcopy Create free account
hub / github.com/MegEngine/MegEngine / __call__

Method __call__

imperative/python/megengine/core/autodiff/grad.py:180–215  ·  view source on GitHub ↗
(self, *args)

Source from the content-addressed store, hash-verified

178 return self._default_rule(*args), self.backward
179
180 def __call__(self, *args):
181 from ...tensor import Tensor
182
183 for arg in args:
184 if not isinstance(arg, Tensor):
185 raise TypeError(
186 "op Function expect type Tensor as inputs, got {}".format(type(arg))
187 )
188
189 grad_key = core2.get_grad_key(args)
190 if grad_key is None:
191 return self._default_rule(*args)
192
193 grad = Grad.key2grad[grad_key]
194 group = [ref() for ref in grad._group]
195
196 origin_args = [Tensor(arg) for arg in args]
197
198 for grad in group:
199 grad.suppress()
200 outputs, backward = self._grad_rule(*args)
201 for grad in reversed(group):
202 grad.resume()
203
204 def normalized_backward(*output_grads):
205 input_grads = backward(*output_grads)
206 if isinstance(input_grads, Tensor) or input_grads is None:
207 input_grads = (input_grads,)
208 return input_grads
209
210 if self.__single_output:
211 outputs = (outputs,)
212 outputs = core2.set_grad(normalized_backward, origin_args, outputs)
213 if self.__single_output:
214 (outputs,) = outputs
215 return outputs
216
217 def __getstate__(self):
218 return self.__dict__

Callers

nothing calls this directly

Calls 7

_default_ruleMethod · 0.95
_grad_ruleMethod · 0.95
refFunction · 0.50
TensorClass · 0.50
formatMethod · 0.45
suppressMethod · 0.45
resumeMethod · 0.45

Tested by

no test coverage detected