(grad, value)
| 78 | |
| 79 | @get_square_sum.register("Tensor", "Number") |
| 80 | def _get_square_sum(grad, value): |
| 81 | norm = P.ReduceSum(False)(F.square(grad), ()) / value |
| 82 | norm = F.expand_dims(F.cast(norm, mstype.float32), 0) |
| 83 | return norm |
| 84 | |
| 85 | |
| 86 | apply_global_norm = C.MultitypeFuncGraph("apply_global_norm") |
nothing calls this directly
no outgoing calls
no test coverage detected