(self, net, param_init_net, param_info)
| 428 | ) |
| 429 | |
| 430 | def _run(self, net, param_init_net, param_info): |
| 431 | param = param_info.blob |
| 432 | param_fp32 = ( |
| 433 | param_info.blob_copy[core.DataType.FLOAT] |
| 434 | if param_info.blob_copy is not None |
| 435 | else None |
| 436 | ) |
| 437 | |
| 438 | # If we have a straight fp32 parameter, run the base class |
| 439 | if param_fp32 is None: |
| 440 | return SgdOptimizer._run(self, net, param_init_net, param_info) |
| 441 | |
| 442 | grad = param_info.grad |
| 443 | if self.base_learning_rate == 0: |
| 444 | return |
| 445 | assert ( |
| 446 | self.base_learning_rate > 0 |
| 447 | ), "Expect positive base learning rate, got {}".format(self.base_learning_rate) |
| 448 | |
| 449 | lr, _ = self.build_lr( |
| 450 | net, |
| 451 | param_init_net, |
| 452 | base_learning_rate=-self.base_learning_rate, |
| 453 | policy=self.policy, |
| 454 | **(self.init_kwargs) |
| 455 | ) |
| 456 | |
| 457 | momentum_data = param_init_net.ConstantFill( |
| 458 | param_fp32, str(param) + "_momentum", value=0.0 |
| 459 | ) |
| 460 | self._aux_params.local.append(momentum_data) |
| 461 | |
| 462 | assert not isinstance( |
| 463 | grad, core.GradientSlice |
| 464 | ), "MultiPrecisionSgd does not support sparse gradients" |
| 465 | |
| 466 | # Copy gradient to fp32 |
| 467 | grad_fp32 = net.HalfToFloat(grad, grad + "_fp32") |
| 468 | |
| 469 | # update (fused) in fp32 |
| 470 | net.MomentumSGDUpdate( |
| 471 | [grad_fp32, momentum_data, lr, param_fp32], |
| 472 | [grad_fp32, momentum_data, param_fp32], |
| 473 | momentum=self.momentum, |
| 474 | nesterov=self.nesterov, |
| 475 | ) |
| 476 | |
| 477 | # Copy updated param back to fp16 |
| 478 | net.FloatToHalf(param_fp32, param) |
| 479 | |
| 480 | |
| 481 | class FP16SgdOptimizer(SgdOptimizer): |
nothing calls this directly
no test coverage detected