Returns a map for param => grad. If params is not specified, all parameters will be considered.
(self, params=None)
| 347 | return param_to_grad |
| 348 | |
| 349 | def GetOptimizationParamInfo(self, params=None): |
| 350 | ''' |
| 351 | Returns a map for param => grad. |
| 352 | If params is not specified, all parameters will be considered. |
| 353 | ''' |
| 354 | if not self.gradient_ops_added: |
| 355 | raise RuntimeError("Need to call AddGradientOperators first") |
| 356 | |
| 357 | param_to_grad = self.param_to_grad |
| 358 | if params: |
| 359 | param_to_grad = self.get_param_to_grad(params) |
| 360 | |
| 361 | return [ |
| 362 | self.get_param_info(param) for param, grad in param_to_grad.items() |
| 363 | if ( |
| 364 | not self.skip_sparse_optim or |
| 365 | not isinstance(grad, core.GradientSlice) |
| 366 | ) |
| 367 | ] |
| 368 | |
| 369 | def _Validate(self): |
| 370 | ''' |