A decorator for registering gradient mappings.
(cls, op_type)
| 1110 | |
| 1111 | @classmethod |
| 1112 | def RegisterGradient(cls, op_type): |
| 1113 | """A decorator for registering gradient mappings.""" |
| 1114 | |
| 1115 | def Wrapper(func): |
| 1116 | cls.gradient_registry_[op_type] = func |
| 1117 | return func |
| 1118 | |
| 1119 | return Wrapper |
| 1120 | |
| 1121 | @classmethod |
| 1122 | def _GetGradientForOpCC(cls, op_def, g_output): |